diff --git a/README.md b/README.md index ae0153e..5d11bca 100644 --- a/README.md +++ b/README.md @@ -228,7 +228,8 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的

语音识别

- - 支持100种语言 + - 支持100种语言
+ - 支持实时语音识别
@@ -343,7 +344,7 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的 cn.smartjavaai smartjavaai-all - 1.0.23 + 1.0.24 ``` ### 3、完整示例代码 @@ -592,15 +593,20 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的 ## 献代码的步骤 1、在Gitee或者Github/Gitcode上fork项目到自己的repo + 2、把fork过去的项目也就是你的项目clone到你的本地 + 3、修改代码(记得一定要修改dev分支) + 4、commit后push到自己的库(dev分支) + 5、登录Gitee或Github/Gitcode在你首页可以看到一个 pull request 按钮,点击它,填写一些说明信息,然后提交即可。 + 6、等待维护者合并 ## 近期更新日志 -## [v1.0.23] - 2025-08-09 +## [v1.0.24] - 2025-08-09 - 新增 语音识别模块,集成 OpenAI 开源的 Whisper 和 Vosk - 修复 质量评估模型的 Bug - 修复 OCR 模块 recognizeAndDraw 方法的 Bug diff --git a/smartjavaai-all/pom.xml b/all/pom.xml similarity index 94% rename from smartjavaai-all/pom.xml rename to all/pom.xml index 75bb859..0e06207 100644 --- a/smartjavaai-all/pom.xml +++ b/all/pom.xml @@ -6,11 +6,11 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - smartjavaai-all - 1.0.23 + all + 1.0.24 ${project.artifactId} SmartJavaAI https://github.com/geekwenjie/SmartJavaAI @@ -33,31 +33,31 @@ cn.smartjavaai - smartjavaai-face + face ${project.version} cn.smartjavaai - smartjavaai-objectdetection + vision ${project.version} cn.smartjavaai - smartjavaai-ocr + ocr ${project.version} cn.smartjavaai - smartjavaai-translate + translate ${project.version} cn.smartjavaai - smartjavaai-speech + speech ${project.version} diff --git a/all/src/test/java/Test.java b/all/src/test/java/Test.java new file mode 100644 index 0000000..245980a --- /dev/null +++ b/all/src/test/java/Test.java @@ -0,0 +1,47 @@ +import ai.djl.Application; +import ai.djl.Model; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ModelZoo; +import cn.smartjavaai.ocr.config.OcrDetModelConfig; +import cn.smartjavaai.ocr.entity.OcrBox; +import cn.smartjavaai.ocr.enums.CommonDetModelEnum; +import cn.smartjavaai.ocr.factory.OcrModelFactory; +import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel; +import lombok.extern.slf4j.Slf4j; +import org.bytedeco.javacv.FFmpegFrameGrabber; +import org.bytedeco.javacv.Frame; +import org.bytedeco.javacv.Java2DFrameUtils; + +import java.awt.image.BufferedImage; +import java.io.IOException; +import java.util.List; + +/** + * @author dwj + * @date 2025/4/24 + */ +@Slf4j +public class Test { + + + public static String savePath = "/Users/wenjie/Downloads/"; + //public static String image1Path = "/Users/wenjie/Documents/idea_workplace/SmartJavaAI-Demo/src/main/resources/5.jpg"; + public static String image1Path = "/Users/wenjie/Documents/idea_workplace/SmartJavaAI-Demo/src/main/resources/MJ_20250226_172200.png"; + + public static String image2Path = "/Users/wenjie/Documents/idea_workplace/SmartJavaAI-Demo/src/main/resources/MJ_20250226_172222.png"; + + public static void main(String[] args) throws IOException { +// // 加载模型 +// Model model = ModelZoo.loadModel(Criteria.builder() +// .optApplication(Application.NLP.ANY) +// .optEngine("PyTorch") +// .optModelName("Llama 3") +// .optTranslatorFactory(new Llama3TranslatorFactory()) +// .optTranslatorProvider(() -> new Llama3Translator()) +// .build()); + + + } + + +} diff --git a/smartjavaai-bom/pom.xml b/bom/pom.xml similarity index 92% rename from smartjavaai-bom/pom.xml rename to bom/pom.xml index acf3279..c73693c 100644 --- a/smartjavaai-bom/pom.xml +++ b/bom/pom.xml @@ -6,12 +6,12 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - 1.0.23 - smartjavaai-bom - smartjavaai-bom + 1.0.24 + bom + sbom 统一版本管理的 BOM 包,同时支持 import 和全量依赖 @@ -25,27 +25,27 @@ cn.smartjavaai - smartjavaai-face + face ${project.parent.version} cn.smartjavaai - smartjavaai-objectdetection + vision ${project.parent.version} cn.smartjavaai - smartjavaai-ocr + ocr ${project.parent.version} cn.smartjavaai - smartjavaai-translate + translate ${project.parent.version} cn.smartjavaai - smartjavaai-speech + speech ${project.parent.version} diff --git a/build/output/0.png b/build/output/0.png new file mode 100644 index 0000000..7092e20 Binary files /dev/null and b/build/output/0.png differ diff --git a/build/output/0crop.png b/build/output/0crop.png new file mode 100644 index 0000000..b7e2cca Binary files /dev/null and b/build/output/0crop.png differ diff --git a/build/output/1.png b/build/output/1.png new file mode 100644 index 0000000..7fc2031 Binary files /dev/null and b/build/output/1.png differ diff --git a/build/output/10.png b/build/output/10.png new file mode 100644 index 0000000..017c967 Binary files /dev/null and b/build/output/10.png differ diff --git a/build/output/10crop.png b/build/output/10crop.png new file mode 100644 index 0000000..5afe051 Binary files /dev/null and b/build/output/10crop.png differ diff --git a/build/output/11.png b/build/output/11.png new file mode 100644 index 0000000..2c1b426 Binary files /dev/null and b/build/output/11.png differ diff --git a/build/output/11crop.png b/build/output/11crop.png new file mode 100644 index 0000000..63b5afd Binary files /dev/null and b/build/output/11crop.png differ diff --git a/build/output/12.png b/build/output/12.png new file mode 100644 index 0000000..1d429e8 Binary files /dev/null and b/build/output/12.png differ diff --git a/build/output/12crop.png b/build/output/12crop.png new file mode 100644 index 0000000..2ae7697 Binary files /dev/null and b/build/output/12crop.png differ diff --git a/build/output/13.png b/build/output/13.png new file mode 100644 index 0000000..f677beb Binary files /dev/null and b/build/output/13.png differ diff --git a/build/output/13crop.png b/build/output/13crop.png new file mode 100644 index 0000000..5fd78d0 Binary files /dev/null and b/build/output/13crop.png differ diff --git a/build/output/14.png b/build/output/14.png new file mode 100644 index 0000000..e7b9346 Binary files /dev/null and b/build/output/14.png differ diff --git a/build/output/14crop.png b/build/output/14crop.png new file mode 100644 index 0000000..af9584b Binary files /dev/null and b/build/output/14crop.png differ diff --git a/build/output/15.png b/build/output/15.png new file mode 100644 index 0000000..9ac9735 Binary files /dev/null and b/build/output/15.png differ diff --git a/build/output/15crop.png b/build/output/15crop.png new file mode 100644 index 0000000..03a69de Binary files /dev/null and b/build/output/15crop.png differ diff --git a/build/output/1crop.png b/build/output/1crop.png new file mode 100644 index 0000000..32bb534 Binary files /dev/null and b/build/output/1crop.png differ diff --git a/build/output/2.png b/build/output/2.png new file mode 100644 index 0000000..9e32185 Binary files /dev/null and b/build/output/2.png differ diff --git a/build/output/2crop.png b/build/output/2crop.png new file mode 100644 index 0000000..b62e6a1 Binary files /dev/null and b/build/output/2crop.png differ diff --git a/build/output/3.png b/build/output/3.png new file mode 100644 index 0000000..b90bc03 Binary files /dev/null and b/build/output/3.png differ diff --git a/build/output/3crop.png b/build/output/3crop.png new file mode 100644 index 0000000..2584aaa Binary files /dev/null and b/build/output/3crop.png differ diff --git a/build/output/4.png b/build/output/4.png new file mode 100644 index 0000000..62cb0ab Binary files /dev/null and b/build/output/4.png differ diff --git a/build/output/4crop.png b/build/output/4crop.png new file mode 100644 index 0000000..22c5211 Binary files /dev/null and b/build/output/4crop.png differ diff --git a/build/output/5.png b/build/output/5.png new file mode 100644 index 0000000..768affd Binary files /dev/null and b/build/output/5.png differ diff --git a/build/output/5crop.png b/build/output/5crop.png new file mode 100644 index 0000000..f9fba23 Binary files /dev/null and b/build/output/5crop.png differ diff --git a/build/output/6.png b/build/output/6.png new file mode 100644 index 0000000..e4602ba Binary files /dev/null and b/build/output/6.png differ diff --git a/build/output/6crop.png b/build/output/6crop.png new file mode 100644 index 0000000..75f7bbc Binary files /dev/null and b/build/output/6crop.png differ diff --git a/build/output/7.png b/build/output/7.png new file mode 100644 index 0000000..a56e06e Binary files /dev/null and b/build/output/7.png differ diff --git a/build/output/7crop.png b/build/output/7crop.png new file mode 100644 index 0000000..7d45920 Binary files /dev/null and b/build/output/7crop.png differ diff --git a/build/output/8.png b/build/output/8.png new file mode 100644 index 0000000..b9d0fdf Binary files /dev/null and b/build/output/8.png differ diff --git a/build/output/8crop.png b/build/output/8crop.png new file mode 100644 index 0000000..17b127e Binary files /dev/null and b/build/output/8crop.png differ diff --git a/build/output/9.png b/build/output/9.png new file mode 100644 index 0000000..3e769d4 Binary files /dev/null and b/build/output/9.png differ diff --git a/build/output/9crop.png b/build/output/9crop.png new file mode 100644 index 0000000..5e1f927 Binary files /dev/null and b/build/output/9crop.png differ diff --git a/build/output/9rotate.png b/build/output/9rotate.png new file mode 100644 index 0000000..59812e1 Binary files /dev/null and b/build/output/9rotate.png differ diff --git a/build/output/cn_layout_detect_result.png b/build/output/cn_layout_detect_result.png new file mode 100644 index 0000000..793a3bc Binary files /dev/null and b/build/output/cn_layout_detect_result.png differ diff --git a/build/output/ocr_1_detected.jpg b/build/output/ocr_1_detected.jpg new file mode 100644 index 0000000..f0e2c82 Binary files /dev/null and b/build/output/ocr_1_detected.jpg differ diff --git a/build/output/table.jpg b/build/output/table.jpg new file mode 100644 index 0000000..b1ba41f Binary files /dev/null and b/build/output/table.jpg differ diff --git a/build/output/yolo_detected.png b/build/output/yolo_detected.png new file mode 100644 index 0000000..0f6494b Binary files /dev/null and b/build/output/yolo_detected.png differ diff --git a/smartjavaai-common/pom.xml b/common/pom.xml similarity index 97% rename from smartjavaai-common/pom.xml rename to common/pom.xml index 749bec1..24a57e4 100644 --- a/smartjavaai-common/pom.xml +++ b/common/pom.xml @@ -6,11 +6,11 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - smartjavaai-common - smartjavaai-common + common + common SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/config/Config.java b/common/src/main/java/cn/smartjavaai/common/config/Config.java similarity index 83% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/config/Config.java rename to common/src/main/java/cn/smartjavaai/common/config/Config.java index 9e1c9d9..cf6ad84 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/config/Config.java +++ b/common/src/main/java/cn/smartjavaai/common/config/Config.java @@ -8,6 +8,7 @@ import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; import java.io.File; +import java.nio.file.Paths; /** * 全局配置 @@ -68,16 +69,25 @@ public class Config { String osName = SystemUtil.getOsInfo().getName(); log.info("当前操作系统:{}", osName); if(osName.toLowerCase().contains("windows")){ - cachePath = SystemUtil.getUserInfo().getHomeDir() + CACHE_DIR; + cachePath = Paths.get( + SystemUtil.getUserInfo().getHomeDir(), + "smartjavaai_cache" + ).toString(); FileUtil.mkdir(cachePath); }else if(osName.toLowerCase().contains("linux")){ cachePath = "/root/" + CACHE_DIR; FileUtil.mkdir(cachePath); }else if(osName.toLowerCase().contains("mac")){ - cachePath = SystemUtil.getUserInfo().getHomeDir() + CACHE_DIR; + cachePath = Paths.get( + SystemUtil.getUserInfo().getHomeDir(), + "smartjavaai_cache" + ).toString(); FileUtil.mkdir(cachePath); }else{ - cachePath = SystemUtil.getUserInfo().getHomeDir() + CACHE_DIR; + cachePath = Paths.get( + SystemUtil.getUserInfo().getHomeDir(), + "smartjavaai_cache" + ).toString(); FileUtil.mkdir(cachePath); } } diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/config/ModelConfig.java b/common/src/main/java/cn/smartjavaai/common/config/ModelConfig.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/config/ModelConfig.java rename to common/src/main/java/cn/smartjavaai/common/config/ModelConfig.java diff --git a/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java b/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java new file mode 100644 index 0000000..9f456a1 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java @@ -0,0 +1,65 @@ +package cn.smartjavaai.common.cv; + +import ai.djl.modality.cv.BufferedImageFactory; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.ImageFactory; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.util.NDImageUtils; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.opencv.OpenCVImageFactory; +import ai.djl.util.Utils; +import cn.smartjavaai.common.utils.Base64ImageUtils; +import cn.smartjavaai.common.utils.OpenCVUtils; +import nu.pattern.OpenCV; +import org.opencv.core.CvType; +import org.opencv.core.Mat; +import org.opencv.core.MatOfByte; +import org.opencv.imgcodecs.Imgcodecs; +import org.opencv.imgproc.Imgproc; + +import java.awt.image.BufferedImage; +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.nio.ByteBuffer; +import java.nio.IntBuffer; +import java.nio.file.Path; + +/** + * 图片处理工厂类 + * @author dwj + */ +public class SmartImageFactory extends BufferedImageFactory { + + private static volatile SmartImageFactory instance; + + public static SmartImageFactory newInstance() { + if (instance == null) { + synchronized (SmartImageFactory.class) { + if (instance == null) { + instance = new SmartImageFactory(); + } + } + } + return instance; + } + + public static SmartImageFactory getInstance(){ + return newInstance(); + } + + public Image fromBufferedImage(BufferedImage sourceImage){ + return fromImage(OpenCVUtils.image2Mat(sourceImage)); + } + + public Image fromBase64(String base64Image) throws IOException { + return fromUrl(base64Image); + } + + public Image fromBytes(byte[] imageData){ + return fromImage(new ByteArrayInputStream(imageData)); + } + +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionInfo.java b/common/src/main/java/cn/smartjavaai/common/entity/DetectionInfo.java similarity index 86% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionInfo.java rename to common/src/main/java/cn/smartjavaai/common/entity/DetectionInfo.java index 4142637..247adca 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionInfo.java +++ b/common/src/main/java/cn/smartjavaai/common/entity/DetectionInfo.java @@ -32,6 +32,16 @@ public class DetectionInfo { */ private ObjectDetInfo objectDetInfo; + /** + * 目标分割信息 + */ + private InstanceSegInfo instanceSegInfo; + + /** + * 旋转框信息 + */ + private ObbDetInfo obbDetInfo; + public DetectionInfo() { diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionRectangle.java b/common/src/main/java/cn/smartjavaai/common/entity/DetectionRectangle.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionRectangle.java rename to common/src/main/java/cn/smartjavaai/common/entity/DetectionRectangle.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionResponse.java b/common/src/main/java/cn/smartjavaai/common/entity/DetectionResponse.java similarity index 86% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionResponse.java rename to common/src/main/java/cn/smartjavaai/common/entity/DetectionResponse.java index 394f3ea..a0df61e 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/DetectionResponse.java +++ b/common/src/main/java/cn/smartjavaai/common/entity/DetectionResponse.java @@ -1,5 +1,6 @@ package cn.smartjavaai.common.entity; +import ai.djl.modality.cv.Image; import lombok.Data; import java.util.List; @@ -14,6 +15,8 @@ public class DetectionResponse { private List detectionInfoList; + private Image drawnImage; + public DetectionResponse() { } diff --git a/common/src/main/java/cn/smartjavaai/common/entity/InstanceSegInfo.java b/common/src/main/java/cn/smartjavaai/common/entity/InstanceSegInfo.java new file mode 100644 index 0000000..6d42298 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/entity/InstanceSegInfo.java @@ -0,0 +1,29 @@ +package cn.smartjavaai.common.entity; + +import lombok.Data; + +/** + * 目标分割信息 + * @author dwj + */ +@Data +public class InstanceSegInfo { + + /** + * 类别名称 + */ + private String className; + + /** + * 遮罩 + */ + private float[][] mask; + + public InstanceSegInfo() { + } + + public InstanceSegInfo(String className, float[][] mask) { + this.className = className; + this.mask = mask; + } +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/Language.java b/common/src/main/java/cn/smartjavaai/common/entity/Language.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/Language.java rename to common/src/main/java/cn/smartjavaai/common/entity/Language.java diff --git a/common/src/main/java/cn/smartjavaai/common/entity/ObbDetInfo.java b/common/src/main/java/cn/smartjavaai/common/entity/ObbDetInfo.java new file mode 100644 index 0000000..7490df3 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/entity/ObbDetInfo.java @@ -0,0 +1,29 @@ +package cn.smartjavaai.common.entity; + +import java.util.List; + +/** + * 定向边界框 检测结果 + * @author dwj + */ +public class ObbDetInfo { + + /** + * 类别名称 + */ + private String className; + + /** + * 检测框坐标 + */ + private RotatedBox rotatedBox; + + public ObbDetInfo() { + } + + + public ObbDetInfo(String className, RotatedBox rotatedBox) { + this.className = className; + this.rotatedBox = rotatedBox; + } +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/ObjectDetInfo.java b/common/src/main/java/cn/smartjavaai/common/entity/ObjectDetInfo.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/ObjectDetInfo.java rename to common/src/main/java/cn/smartjavaai/common/entity/ObjectDetInfo.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/Point.java b/common/src/main/java/cn/smartjavaai/common/entity/Point.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/Point.java rename to common/src/main/java/cn/smartjavaai/common/entity/Point.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/R.java b/common/src/main/java/cn/smartjavaai/common/entity/R.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/R.java rename to common/src/main/java/cn/smartjavaai/common/entity/R.java diff --git a/common/src/main/java/cn/smartjavaai/common/entity/RotatedBox.java b/common/src/main/java/cn/smartjavaai/common/entity/RotatedBox.java new file mode 100644 index 0000000..bf76541 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/entity/RotatedBox.java @@ -0,0 +1,38 @@ +package cn.smartjavaai.common.entity; + +/** + * 旋转框 + * @author dwj + */ +public class RotatedBox { + + /** + * 左上角 + */ + private Point topLeft; + + /** + * 右上角 + */ + private Point topRight; + + /** + * 右下角 + */ + private Point bottomRight; + + /** + * 左下角 + */ + private Point bottomLeft; + + public RotatedBox(Point topLeft, Point topRight, Point bottomRight, Point bottomLeft) { + this.topLeft = topLeft; + this.topRight = topRight; + this.bottomRight = bottomRight; + this.bottomLeft = bottomLeft; + } + + public RotatedBox() { + } +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/ExpressionResult.java b/common/src/main/java/cn/smartjavaai/common/entity/face/ExpressionResult.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/ExpressionResult.java rename to common/src/main/java/cn/smartjavaai/common/entity/face/ExpressionResult.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/FaceAttribute.java b/common/src/main/java/cn/smartjavaai/common/entity/face/FaceAttribute.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/FaceAttribute.java rename to common/src/main/java/cn/smartjavaai/common/entity/face/FaceAttribute.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/FaceInfo.java b/common/src/main/java/cn/smartjavaai/common/entity/face/FaceInfo.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/FaceInfo.java rename to common/src/main/java/cn/smartjavaai/common/entity/face/FaceInfo.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/FaceSearchResult.java b/common/src/main/java/cn/smartjavaai/common/entity/face/FaceSearchResult.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/FaceSearchResult.java rename to common/src/main/java/cn/smartjavaai/common/entity/face/FaceSearchResult.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/HeadPose.java b/common/src/main/java/cn/smartjavaai/common/entity/face/HeadPose.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/HeadPose.java rename to common/src/main/java/cn/smartjavaai/common/entity/face/HeadPose.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/LivenessResult.java b/common/src/main/java/cn/smartjavaai/common/entity/face/LivenessResult.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/face/LivenessResult.java rename to common/src/main/java/cn/smartjavaai/common/entity/face/LivenessResult.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/ocr/TableStructure.java b/common/src/main/java/cn/smartjavaai/common/entity/ocr/TableStructure.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/ocr/TableStructure.java rename to common/src/main/java/cn/smartjavaai/common/entity/ocr/TableStructure.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/DeviceEnum.java b/common/src/main/java/cn/smartjavaai/common/enums/DeviceEnum.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/DeviceEnum.java rename to common/src/main/java/cn/smartjavaai/common/enums/DeviceEnum.java diff --git a/common/src/main/java/cn/smartjavaai/common/enums/VideoSourceType.java b/common/src/main/java/cn/smartjavaai/common/enums/VideoSourceType.java new file mode 100644 index 0000000..522c501 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/enums/VideoSourceType.java @@ -0,0 +1,12 @@ +package cn.smartjavaai.common.enums; + +/** + * 视频源类型枚举 + * @author dwj + * @date 2025/8/27 + */ +public enum VideoSourceType { + STREAM, // RTSP 或 HTTP 流 + FILE, // 本地视频文件 + CAMERA; // 本地摄像头 +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/EyeStatus.java b/common/src/main/java/cn/smartjavaai/common/enums/face/EyeStatus.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/EyeStatus.java rename to common/src/main/java/cn/smartjavaai/common/enums/face/EyeStatus.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/FacialExpression.java b/common/src/main/java/cn/smartjavaai/common/enums/face/FacialExpression.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/FacialExpression.java rename to common/src/main/java/cn/smartjavaai/common/enums/face/FacialExpression.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/GenderType.java b/common/src/main/java/cn/smartjavaai/common/enums/face/GenderType.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/GenderType.java rename to common/src/main/java/cn/smartjavaai/common/enums/face/GenderType.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/LivenessStatus.java b/common/src/main/java/cn/smartjavaai/common/enums/face/LivenessStatus.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/enums/face/LivenessStatus.java rename to common/src/main/java/cn/smartjavaai/common/enums/face/LivenessStatus.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java b/common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java rename to common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/ModelPredictorPoolManager.java b/common/src/main/java/cn/smartjavaai/common/pool/ModelPredictorPoolManager.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/ModelPredictorPoolManager.java rename to common/src/main/java/cn/smartjavaai/common/pool/ModelPredictorPoolManager.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java b/common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java rename to common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java b/common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java rename to common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/ArrayUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/ArrayUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/ArrayUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/ArrayUtils.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/Base64ImageUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/Base64ImageUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/Base64ImageUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/Base64ImageUtils.java diff --git a/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java new file mode 100644 index 0000000..0f1ccb1 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java @@ -0,0 +1,32 @@ +package cn.smartjavaai.common.utils; + +import java.nio.file.Files; +import java.nio.file.Path; + +/** + * @author dwj + */ +public class DJLCommonUtils { + + /** + * 检查模型目录中是否存在 "serving.properties" 文件 + * + * @param modelPath 模型目录路径 + * @return true 表示存在,false 表示不存在 + */ + public static boolean isServingPropertiesExists(Path modelPath) { + if (modelPath == null || !Files.exists(modelPath)) { + return false; + } + // 确定目录路径 + Path dirPath = Files.isDirectory(modelPath) ? modelPath : modelPath.getParent(); + if (dirPath == null) { + return false; // 可能是根目录的文件 + } + + // 判断目录下的 serving.properties 是否存在 + Path servingFile = dirPath.resolve("serving.properties"); + return Files.exists(servingFile); + } + +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/FileUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/FileUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/FileUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/FileUtils.java diff --git a/common/src/main/java/cn/smartjavaai/common/utils/FrameConverterUtil.java b/common/src/main/java/cn/smartjavaai/common/utils/FrameConverterUtil.java new file mode 100644 index 0000000..b2ac439 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/utils/FrameConverterUtil.java @@ -0,0 +1,66 @@ +package cn.smartjavaai.common.utils; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.ImageFactory; +import org.bytedeco.javacpp.BytePointer; +import org.bytedeco.javacv.OpenCVFrameConverter; +import org.bytedeco.opencv.opencv_core.CvMat; +import org.bytedeco.opencv.opencv_core.Mat; +import org.opencv.core.CvType; + +import java.awt.image.BufferedImage; + +/** + * @author dwj + * @date 2025/8/27 + */ +public class FrameConverterUtil { + + /** + * 将 Bytedeco Mat 转为 DJL Image + * 支持 1/3/4 通道 + */ + public static Image matToDJLImage(Mat cvMat) { + if (cvMat == null || cvMat.empty()) { + return null; + } + + int width = cvMat.cols(); + int height = cvMat.rows(); + int channels = cvMat.channels(); + + int[] pixels = new int[width * height]; + + if (channels == 1) { // 灰度图 + byte[] data = new byte[width * height]; + cvMat.data().get(data); + for (int i = 0; i < width * height; i++) { + int gray = data[i] & 0xFF; + pixels[i] = (0xFF << 24) | (gray << 16) | (gray << 8) | gray; + } + } else if (channels == 3) { // BGR + byte[] data = new byte[width * height * 3]; + cvMat.data().get(data); + for (int i = 0; i < width * height; i++) { + int b = data[i * 3] & 0xFF; + int g = data[i * 3 + 1] & 0xFF; + int r = data[i * 3 + 2] & 0xFF; + pixels[i] = (0xFF << 24) | (r << 16) | (g << 8) | b; + } + } else if (channels == 4) { // BGRA + byte[] data = new byte[width * height * 4]; + cvMat.data().get(data); + for (int i = 0; i < width * height; i++) { + int b = data[i * 4] & 0xFF; + int g = data[i * 4 + 1] & 0xFF; + int r = data[i * 4 + 2] & 0xFF; + int a = data[i * 4 + 3] & 0xFF; + pixels[i] = (a << 24) | (r << 16) | (g << 8) | b; + } + } else { + throw new IllegalArgumentException("只支持 1/3/4 通道图像"); + } + + return ImageFactory.getInstance().fromPixels(pixels, width, height); + } +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java similarity index 99% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java index e71e281..e0b6344 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java +++ b/common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java @@ -139,7 +139,7 @@ public class ImageUtils { * @param width * @param height */ - public static void drawImageRect(Image image, int x, int y, int width, int height) { + public static void drawBufferedImageRect(Image image, int x, int y, int width, int height) { // 将绘制图像转换为Graphics2D BufferedImage bufferedImage = (BufferedImage)image.getWrappedImage(); Graphics2D g = (Graphics2D) bufferedImage.getGraphics(); @@ -501,4 +501,7 @@ public class ImageUtils { + + + } diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/LetterBoxUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/LetterBoxUtils.java similarity index 80% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/LetterBoxUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/LetterBoxUtils.java index 678251b..c2eca79 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/LetterBoxUtils.java +++ b/common/src/main/java/cn/smartjavaai/common/utils/LetterBoxUtils.java @@ -1,5 +1,6 @@ package cn.smartjavaai.common.utils; +import ai.djl.modality.cv.output.Rectangle; import ai.djl.modality.cv.util.NDImageUtils; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDManager; @@ -121,4 +122,27 @@ public class LetterBoxUtils { return boxes; } + /** + * 恢复缩放后的 box(左上角坐标) + * @param rectangle + * @param scale + * @param origImageWidth + * @param origImageHeight + */ + public static Rectangle restoreBox(Rectangle rectangle, float scale, int origImageWidth, int origImageHeight, int inputWidth, int inputHeight){ + double paddingWidth = (inputWidth - origImageWidth * scale) / 2; + double paddingHeight = (inputHeight - origImageHeight * scale) / 2; + + // 去掉 padding + double x_noPad = rectangle.getX() - paddingWidth; + double y_noPad = rectangle.getY() - paddingHeight; + + //模型输出就是原图坐标 + double x1 = x_noPad / scale / origImageWidth; + double y1 = y_noPad / scale / origImageHeight; + double boxW = rectangle.getWidth() / scale / origImageWidth ; + double boxH = rectangle.getHeight() / scale / origImageHeight; + return new Rectangle(x1, y1, boxW, boxH); + } + } diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/NMSUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/NMSUtils.java similarity index 53% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/NMSUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/NMSUtils.java index ef2799a..390f892 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/NMSUtils.java +++ b/common/src/main/java/cn/smartjavaai/common/utils/NMSUtils.java @@ -1,6 +1,10 @@ package cn.smartjavaai.common.utils; import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; import java.util.ArrayList; import java.util.List; @@ -64,4 +68,43 @@ public class NMSUtils { return keep.stream().mapToInt(i -> i).toArray(); } + + + /** + * 批量执行 NMS,输入 NDArray 形式的 boxes、scores 和 idxs,返回保留的索引列表 + * + * @param boxes NDArray 形状为 (N, 4),格式为 [x1, y1, x2, y2] + * @param scores NDArray 形状为 (N,) 或 (N,1),每个 box 的置信度 + * @param idxs NDArray 形状为 (N,),每个 box 对应的 batch id + * @param iouThreshold IOU 阈值,超过该阈值则认为有 + * @return 批量保留框的索引列表 + * + */ + public static NDArray batchedNms(NDArray boxes, NDArray scores, NDArray idxs, float iouThreshold, NDManager manager) { + List keepList = new ArrayList<>(); + // 获取唯一 batch id + NDArray uniqueIdxs = idxs.unique().get(0); + for (long batchId : uniqueIdxs.toLongArray()) { + // 找出当前 batch 的框 + NDArray mask = idxs.eq(batchId); + NDArray batchBoxes = boxes.get(mask); + NDArray batchScores = scores.get(mask); + // 执行单 batch NMS + int[] keepIndices = nms(batchBoxes, batchScores, iouThreshold); + if (keepIndices.length > 0) { + // 将局部索引映射回全局索引 + NDArray globalIndices = manager.arange(boxes.getShape().get(0)) + .get(mask) + .toType(DataType.INT64, false) + .get(manager.create(keepIndices)); + + keepList.add(globalIndices); + } + } + if (keepList.isEmpty()) { + return manager.create(new long[0]); + } + return NDArrays.concat(new NDList(keepList)); + } + } diff --git a/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java new file mode 100644 index 0000000..ea782d1 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java @@ -0,0 +1,265 @@ +package cn.smartjavaai.common.utils; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.ImageFactory; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDManager; +import ai.djl.util.RandomUtils; +import cn.smartjavaai.common.entity.DetectionInfo; +import org.apache.commons.collections.CollectionUtils; +import org.apache.commons.lang3.StringUtils; +import org.opencv.core.*; +import org.opencv.core.Point; +import org.opencv.imgproc.Imgproc; + +import java.awt.*; +import java.awt.image.BufferedImage; +import java.awt.image.DataBufferByte; +import java.util.List; +import java.util.Objects; + +/** + * OpenCV 工具类 + */ +public class OpenCVUtils { + /** + * canny算法,边缘检测 + * + * @param src + * @return + */ + public static Mat canny(Mat src) { + Mat mat = src.clone(); + Imgproc.Canny(src, mat, 100, 200); + return mat; + } + + /** + * 画线 + * + * @param mat + * @param point1 + * @param point2 + */ + public static void line(Mat mat, Point point1, Point point2) { + Imgproc.line(mat, point1, point2, new Scalar(255, 255, 255), 1); + } + + /** + * NDArray to opencv_core.Mat + * + * @param manager + * @param srcPoints + * @param dstPoints + * @return + */ + public static Mat toOpenCVMat(NDManager manager, NDArray srcPoints, NDArray dstPoints) { + NDArray svdMat = SVDUtils.transformationFromPoints(manager, srcPoints, dstPoints); + double[] doubleArray = svdMat.toDoubleArray(); + Mat newSvdMat = new Mat(2, 3, CvType.CV_64F); + for (int i = 0; i < 2; i++) { + for (int j = 0; j < 3; j++) { + newSvdMat.put(i, j, doubleArray[i * 3 + j]); + } + } + return newSvdMat; + } + + /** + * double[][] points array to Mat + * @param points + * @return + */ + public static Mat toOpenCVMat(double[][] points) { + Mat mat = new Mat(5, 2, CvType.CV_64F); + for (int i = 0; i < 5; i++) { + for (int j = 0; j < 2; j++) { + mat.put(i, j, points[i * 5 + j]); + } + } + return mat; + } + + /** + * 变换矩阵的逆矩阵 + * + * @param src + * @return + */ + public static Mat invertAffineTransform(Mat src) { + Mat dst = src.clone(); + Imgproc.invertAffineTransform(src, dst); + return dst; + } + + /** + * Mat to BufferedImage + * + * @param mat + * @return + */ + public static BufferedImage mat2Image(Mat mat) { + int width = mat.width(); + int height = mat.height(); + byte[] data = new byte[width * height * (int) mat.elemSize()]; + Imgproc.cvtColor(mat, mat, 4); + mat.get(0, 0, data); + BufferedImage ret = new BufferedImage(width, height, 5); + ret.getRaster().setDataElements(0, 0, width, height, data); + return ret; + } + + /** + * BufferedImage to Mat + * + * @param img + * @return + */ + public static Mat image2Mat(BufferedImage img) { + int width = img.getWidth(); + int height = img.getHeight(); + + // 强制转换为 TYPE_3BYTE_BGR,自动去除透明通道 + BufferedImage convertedImg = new BufferedImage(width, height, BufferedImage.TYPE_3BYTE_BGR); + Graphics2D g2d = convertedImg.createGraphics(); + g2d.drawImage(img, 0, 0, null); + g2d.dispose(); + + byte[] data = ((DataBufferByte) convertedImg.getRaster().getDataBuffer()).getData(); + Mat mat = new Mat(height, width, CvType.CV_8UC3); + mat.put(0, 0, data); + return mat; + } + + /** + * 透视变换 + * + * @param src + * @param srcPoints + * @param dstPoints + * @return + */ + public static Mat perspectiveTransform(Mat src, Mat srcPoints, Mat dstPoints) { + Mat dst = src.clone(); + Mat warp_mat = Imgproc.getPerspectiveTransform(srcPoints, dstPoints); + Imgproc.warpPerspective(src, dst, warp_mat, dst.size()); + warp_mat.release(); + return dst; + } + + /** + * 绘制矩形框和文字 + * + * @param image + * @param detectionInfoList + */ + public static void drawRectAndText(Image image, List detectionInfoList) { + if(CollectionUtils.isEmpty(detectionInfoList)) + return; + for(DetectionInfo detectionInfo : detectionInfoList){ + drawRectAndText(image, detectionInfo); + } + } + + + /** + * 绘制矩形框和文字 + * + * @param image + * @param detectionInfo + */ + public static void drawRectAndText(Image image, DetectionInfo detectionInfo) { + + + Mat mat = (Mat)image.getWrappedImage(); + if (image == null) return; + int x = detectionInfo.getDetectionRectangle().getX(); + int y = detectionInfo.getDetectionRectangle().getY(); + int width = detectionInfo.getDetectionRectangle().getWidth(); + int height = detectionInfo.getDetectionRectangle().getHeight(); + + Scalar rectangleColor = new Scalar((double)RandomUtils.nextInt(178), (double)RandomUtils.nextInt(178), (double)RandomUtils.nextInt(178)); + + // 绘制矩形框 + Point pt1 = new Point(x, y); + Point pt2 = new Point(x + width, y + height); + Imgproc.rectangle(mat, pt1, pt2, rectangleColor, 2); + + // 绘制文字 + if (Objects.nonNull(detectionInfo.getObjectDetInfo()) && StringUtils.isNotBlank(detectionInfo.getObjectDetInfo().getClassName())) { + String className = detectionInfo.getObjectDetInfo().getClassName(); + Size size = Imgproc.getTextSize(className, 1, 1.3, 1, (int[])null); + Point br = new Point((double)x + size.width + 4.0, (double)y + size.height + 4.0); + Imgproc.rectangle(mat, pt1, br, rectangleColor, -1); + Point point = new Point((double)x, (double)y + size.height + 2.0); + Scalar color = new Scalar(255.0, 255.0, 255.0); + Imgproc.putText(mat, className, point, 1, 1.3, color, 1); + } + image = ImageFactory.getInstance().fromImage(mat); + } + + /** + * 在Mat上绘制矩形框和文字 + * + * @param mat 待绘制的Mat + * @param x 矩形左上角X + * @param y 矩形左上角Y + * @param width 矩形宽度 + * @param height 矩形高度 + * @param color 框的颜色,例如 new Scalar(0, 255, 0) 绿色 + * @param thickness 框线宽度 + * @param text 需要绘制的文字,可以为null或空 + * @param fontScale 文字缩放比例 + * @param textColor 文字颜色 + */ + public static void drawRectAndText(Mat mat, + int x, int y, int width, int height, + Scalar color, int thickness, + String text, double fontScale, Scalar textColor) { + + if (mat == null || mat.empty()) return; + + // 绘制矩形框 + Point pt1 = new Point(x, y); + Point pt2 = new Point(x + width, y + height); + Imgproc.rectangle(mat, pt1, pt2, color, thickness); + + // 绘制文字 + if (text != null && !text.isEmpty()) { + int baseline[] = new int[1]; + Size textSize = Imgproc.getTextSize(text, Imgproc.FONT_HERSHEY_SIMPLEX, fontScale, thickness, baseline); + // 保证文字不超出矩形 + Point textOrg = new Point(x, y - 5 < 0 ? y + textSize.height + 5 : y - 5); + Imgproc.putText(mat, text, textOrg, Imgproc.FONT_HERSHEY_SIMPLEX, fontScale, textColor, thickness); + } + } + + /** + * 将 Bytedeco 的 Mat 转换为 OpenCV 官方的 Mat + * @param src Bytedeco Mat (BGR 或 BGRA) + * @return OpenCV Mat (BGR 或 BGRA) + */ + public static org.opencv.core.Mat convertToOpenCVMat(org.bytedeco.opencv.opencv_core.Mat bMat) { + + + try { + int width = bMat.cols(); + int height = bMat.rows(); + int channels = bMat.channels(); + + // 创建 OpenCV Mat + org.opencv.core.Mat cvMat = new org.opencv.core.Mat(height, width, channels == 3 ? CvType.CV_8UC3 : CvType.CV_8UC1); + + // 从 bytedeco Mat 获取像素数据 + byte[] data = new byte[width * height * channels]; + bMat.data().get(data); + + // 填充到 OpenCV Mat + cvMat.put(0, 0, data); + return cvMat; + } catch (Throwable e) { + e.printStackTrace(); + } + return null; + } +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/PointUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/PointUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/PointUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/PointUtils.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/PoolUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/PoolUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/PoolUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/PoolUtils.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/SVDUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/SVDUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/SVDUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/SVDUtils.java diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/VideoUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/VideoUtils.java similarity index 100% rename from smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/VideoUtils.java rename to common/src/main/java/cn/smartjavaai/common/utils/VideoUtils.java diff --git a/examples/face-example/pom.xml b/examples/face-example/pom.xml index cd2e3c9..d10c33d 100644 --- a/examples/face-example/pom.xml +++ b/examples/face-example/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.23 + 1.0.24 smartai.examples.face.facedet.FaceDetDemo @@ -255,6 +255,14 @@ runtime + + ai.djl.pytorch + pytorch-native-cpu + linux-aarch64 + runtime + 2.5.1 + + diff --git a/examples/face-example/src/main/java/smartai/examples/face/PythonTranslator.java b/examples/face-example/src/main/java/smartai/examples/face/PythonTranslator.java new file mode 100644 index 0000000..542118c --- /dev/null +++ b/examples/face-example/src/main/java/smartai/examples/face/PythonTranslator.java @@ -0,0 +1,118 @@ +/* + * Copyright 2023 Amazon.com, Inc. or its affiliates. All Rights Reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file except in compliance + * with the License. A copy of the License is located at + * + * http://aws.amazon.com/apache2.0/ + * + * or in the "license" file accompanying this file. This file is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES + * OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions + * and limitations under the License. + */ +package smartai.examples.face; + +import ai.djl.ModelException; +import ai.djl.inference.Predictor; +import ai.djl.modality.Classifications; +import ai.djl.modality.Input; +import ai.djl.modality.Output; +import ai.djl.modality.cv.Image; +import ai.djl.ndarray.NDList; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ZooModel; +import ai.djl.translate.NoBatchifyTranslator; +import ai.djl.translate.TranslateException; +import ai.djl.translate.TranslatorContext; +import ai.djl.util.JsonUtils; +import ai.djl.util.Utils; + +import com.google.gson.reflect.TypeToken; + +import java.io.IOException; +import java.io.InputStream; +import java.lang.reflect.Type; +import java.net.URL; +import java.nio.file.Paths; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +public class PythonTranslator implements NoBatchifyTranslator { + + private ZooModel model; + private Predictor predictor; + + @Override + public void prepare(TranslatorContext ctx) throws ModelException, IOException { + if (predictor == null) { + Criteria criteria = + Criteria.builder() + .setTypes(Input.class, Output.class) + .optModelPath(Paths.get("src/test/python")) + .optEngine("Python") + .build(); + model = criteria.loadModel(); + predictor = model.newPredictor(); + } + } + +// @Override +// public NDList processInput(TranslatorContext ctx, String url) +// throws IOException, TranslateException { +// Input input = new Input(); +// try (InputStream is = new URL(url).openStream()) { +// input.add("data", Utils.toByteArray(is)); +// } +// input.addProperty("Content-Type", "image/jpeg"); +// // calling preprocess() function in model.py +// input.addProperty("handler", "preprocess"); +// Output output = predictor.predict(input); +// if (output.getCode() != 200) { +// throw new TranslateException("Python preprocess() failed: " + output.getMessage()); +// } +// +// return output.getDataAsNDList(ctx.getNDManager()); +// } + + @Override + public NDList processInput(TranslatorContext ctx, byte[] image) + throws IOException, TranslateException { + Input input = new Input(); + input.add("data", image); + input.addProperty("Content-Type", "image/jpeg"); + // calling preprocess() function in model.py + input.addProperty("handler", "preprocess"); + Output output = predictor.predict(input); + if (output.getCode() != 200) { + throw new TranslateException("Python preprocess() failed: " + output.getMessage()); + } + return output.getDataAsNDList(ctx.getNDManager()); + } + + @Override + public Classifications processOutput(TranslatorContext ctx, NDList list) + throws TranslateException { + Input input = new Input(); + input.add("data", list); + // calling postprocess() function in processing.py + input.addProperty("handler", "postprocess"); + Output output = predictor.predict(input); + if (output.getCode() != 200) { + throw new TranslateException("Python postprocess() failed: " + output.getMessage()); + } + + String json = output.getData().getAsString(); + System.out.println("json:" + json); + return null; + } + + public void close() { + if (predictor != null) { + predictor.close(); + model.close(); + predictor = null; + model = null; + } + } +} diff --git a/examples/face-example/src/main/java/smartai/examples/face/Test.java b/examples/face-example/src/main/java/smartai/examples/face/Test.java new file mode 100644 index 0000000..cbcbcd7 --- /dev/null +++ b/examples/face-example/src/main/java/smartai/examples/face/Test.java @@ -0,0 +1,99 @@ +package smartai.examples.face; + +import ai.djl.Application; +import ai.djl.Device; +import ai.djl.MalformedModelException; +import ai.djl.inference.Predictor; +import ai.djl.modality.Classifications; +import ai.djl.modality.audio.Audio; +import ai.djl.modality.audio.AudioFactory; +import ai.djl.modality.audio.translator.SpeechRecognitionTranslatorFactory; +import ai.djl.repository.Artifact; +import ai.djl.repository.MRL; +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.translate.TranslateException; +import lombok.extern.slf4j.Slf4j; + +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.Map; + +/** + * @author dwj + * @date 2025/7/29 + */ +@Slf4j +public class Test { + + public static void main(String[] args) throws ModelNotFoundException, MalformedModelException, IOException, TranslateException { +// PythonTranslator translator = new PythonTranslator(); +// Criteria criteria = +// Criteria.builder() +// .setTypes(byte[].class, Classifications.class) +// .optModelPath(Paths.get("/Users/wenjie/Documents/develop/model/arcfaceresnet100-11-int8.onnx")) +// .optEngine("OnnxRuntime") +// .optTranslator(translator) +// .build(); +// String path = "/Users/wenjie/Downloads/facetest/jsy.jpg"; +// try (ZooModel model = criteria.loadModel(); +// Predictor predictor = model.newPredictor()) { +// byte[] data = Files.readAllBytes(Paths.get(path)); +// Classifications ret = predictor.predict(data); +// System.out.println(ret); +// } +// +// // unload python model +// translator.close(); + + + // Load model. + // Wav2Vec2 model is a speech model that accepts a float array corresponding to the raw + // waveform of the speech signal. + +// String url = "/Users/wenjie/Downloads/20210601_u2++_conformer_exp/final.pt"; +// Criteria criteria = +// Criteria.builder() +// .setTypes(Audio.class, String.class) +//// .optModelUrls(url) +// .optModelPath(Paths.get(url)) +// .optDevice(Device.cpu()) // torchscript model only support CPU +// .optTranslatorFactory(new SpeechRecognitionTranslatorFactory()) +//// .optModelName("data.pkl") +// .optEngine("PyTorch") +// .build(); +// +// // Read in audio file +// String wave = "https://resources.djl.ai/audios/speech.wav"; +// Audio audio = AudioFactory.newInstance().fromUrl(wave); +// try (ZooModel model = criteria.loadModel(); +// Predictor predictor = model.newPredictor()) { +// String result = predictor.predict(audio); +// log.info("Result: {}", result); +// } + + boolean withArtifacts = + args.length > 0 && ("--artifact".equals(args[0]) || "-a".equals(args[0])); + if (!withArtifacts) { + log.info("============================================================"); + log.info("user ./gradlew listModel --args='-a' to show artifact detail"); + log.info("============================================================"); + } + Map> models = ModelZoo.listModels(); + for (Map.Entry> entry : models.entrySet()) { + String appName = entry.getKey().toString(); + for (Artifact artifact : entry.getValue()) { + if (withArtifacts) { + log.info("{} djl://{}", appName, artifact); + } else { + log.info("{} {}", appName, artifact); + } + } + } + } +} diff --git a/examples/face-example/src/main/java/smartai/examples/face/facedet/FaceDetDemo.java b/examples/face-example/src/main/java/smartai/examples/face/facedet/FaceDetDemo.java index e084c20..52dc205 100644 --- a/examples/face-example/src/main/java/smartai/examples/face/facedet/FaceDetDemo.java +++ b/examples/face-example/src/main/java/smartai/examples/face/facedet/FaceDetDemo.java @@ -67,7 +67,7 @@ public class FaceDetDemo { //高精度模型,速度慢 config.setModelEnum(FaceDetModelEnum.RETINA_FACE);//人脸检测模型 //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/xxx/Documents/develop/model/retinaface.pt"); + config.setModelPath("/Users/wenjie/Documents/develop/face_model/retinaface.pt"); config.setConfidenceThreshold(FaceDetectConstant.DEFAULT_CONFIDENCE_THRESHOLD);//只返回相似度大于该值的人脸 config.setNmsThresh(FaceDetectConstant.NMS_THRESHOLD);//用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 return FaceDetModelFactory.getInstance().getModel(config); @@ -95,12 +95,21 @@ public class FaceDetDemo { @Test public void testFaceDetect(){ try { - FaceDetModel faceModel = FaceDetModelFactory.getInstance().getModel(); + FaceDetModel faceModel = getFaceDetModel(); R detectedResult = faceModel.detect(imgPath); - if(detectedResult.isSuccess()){ - log.info("人脸检测结果:{}", JSONObject.toJSONString(detectedResult.getData())); +// if(detectedResult.isSuccess()){ +// log.info("人脸检测结果:{}", JSONObject.toJSONString(detectedResult.getData())); +// }else{ +// log.info("人脸检测失败:{}", detectedResult.getMessage()); +// } + + long start = System.currentTimeMillis(); + R detectedResult2 = faceModel.detect("/Users/wenjie/Downloads/facetest/surprise.png"); + log.info("耗时:{}", System.currentTimeMillis() - start); + if(detectedResult2.isSuccess()){ + log.info("人脸检测结果:{}", JSONObject.toJSONString(detectedResult2.getData())); }else{ - log.info("人脸检测失败:{}", detectedResult.getMessage()); + log.info("人脸检测失败:{}", detectedResult2.getMessage()); } } catch (Exception e) { e.printStackTrace(); diff --git a/examples/face-example/src/main/java/smartai/examples/face/facerec/FaceRecDemo.java b/examples/face-example/src/main/java/smartai/examples/face/facerec/FaceRecDemo.java index 1ed3740..38f0495 100644 --- a/examples/face-example/src/main/java/smartai/examples/face/facerec/FaceRecDemo.java +++ b/examples/face-example/src/main/java/smartai/examples/face/facerec/FaceRecDemo.java @@ -63,7 +63,7 @@ public class FaceRecDemo { //高精度模型,速度慢 config.setModelEnum(FaceDetModelEnum.RETINA_FACE);//人脸检测模型 //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/xxx/Documents/develop/model/retinaface.pt"); +// config.setModelPath("/Users/wenjie/Documents/develop/model/retinaface.pt"); config.setConfidenceThreshold(FaceDetectConstant.DEFAULT_CONFIDENCE_THRESHOLD);//只返回相似度大于该值的人脸 config.setNmsThresh(FaceDetectConstant.NMS_THRESHOLD);//用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 config.setDevice(device); @@ -170,7 +170,7 @@ public class FaceRecDemo { FaceRecConfig config = new FaceRecConfig(); //高精度模型,速度慢, 追求速度请更换高速模型,具体其他模型参数可以查看文档:http://doc.smartjavaai.cn/face.html config.setModelEnum(FaceRecModelEnum.ELASTIC_FACE_MODEL);//人脸检测模型 - config.setModelPath("/Users/xxx/Documents/develop/model/elasticface.pt"); + config.setModelPath("/Users/wenjie/Documents/develop/model/elasticface.pt"); //裁剪人脸:如果图片已经是裁剪过的,则请将此参数设置为false config.setCropFace(true); //开启人脸对齐:适用于人脸不正的场景,开启将提升人脸特征准确度,关闭可以提升性能 diff --git a/examples/face-example/src/test/python/model.py b/examples/face-example/src/test/python/model.py new file mode 100644 index 0000000..a3fd8a1 --- /dev/null +++ b/examples/face-example/src/test/python/model.py @@ -0,0 +1,126 @@ +#!/usr/bin/env python +# +# Copyright 2023 Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file +# except in compliance with the License. A copy of the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "LICENSE.txt" file accompanying this file. This file is distributed on an "AS IS" +# BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, express or implied. See the License for +# the specific language governing permissions and limitations under the License. +""" +PyTorch resnet18 pre/post processing example. +""" + +import json +import logging +import os +from typing import Optional, Any +import sklearn + +import torch +import torch.nn.functional as F +from torchvision import transforms + +from djl_python import Input +from djl_python import Output + + +class Processing(object): + + def __init__(self): + self.topK = 5 + self.image_processing = None + self.mapping = None + self.initialized = False + + def initialize(self, properties: dict): + """ + Initialize model. + """ + self.image_processing = transforms.Compose([ + transforms.Resize(112), + transforms.CenterCrop(112), + transforms.ToTensor(), + transforms.Normalize(mean=[0.485, 0.456, 0.406], + std=[0.229, 0.224, 0.225]) + ]) + #self.mapping = self.load_label_mapping("index_to_name.json") + self.initialized = True + + def preprocess(self, inputs: Input) -> Output: + outputs = Output() + try: + batch = inputs.get_batches() + images = [] + for i, item in enumerate(batch): + image = self.image_processing(item.get_as_image()) + images.append(image) + images = torch.stack(images) + outputs.add_as_numpy(images.detach().numpy()) + outputs.add_property("content-type", "tensor/ndlist") + except Exception as e: + logging.exception("pre-process failed") + # error handling + outputs = Output().error(str(e)) + + return outputs + + def postprocess(self, inputs: Input) -> Output: + outputs = Output() + try: + data = inputs.get_as_numpy(0)[0] + item = torch.from_numpy(data) + print("data shape:", item.shape) + embedding = sklearn.preprocessing.normalize(item).flatten() + outputs.add(embedding) + except Exception as e: + logging.exception("post-process failed") + # error handling + outputs = Output().error(str(e)) + + return outputs + + @staticmethod + def load_label_mapping(mapping_file_path: Any) -> dict: + if not os.path.isfile(mapping_file_path): + raise Exception('mapping file not found: ' + mapping_file_path) + + with open(mapping_file_path) as f: + mapping = json.load(f) + if not isinstance(mapping, dict): + raise Exception('mapping file should be in "class":"label" format') + + for key, value in mapping.items(): + new_value = value + if isinstance(new_value, list): + new_value = value[-1] + if not isinstance(new_value, str): + raise Exception( + 'labels in mapping must be either str or [str]') + mapping[key] = new_value + return mapping + + +_service = Processing() + + +def preprocess(inputs: Input) -> Output: + return _service.preprocess(inputs) + + +def postprocess(inputs: Input) -> Output: + return _service.postprocess(inputs) + + +def handle(inputs: Input) -> Optional[Output]: + """ + Default handler function + """ + if not _service.initialized: + # stateful model + _service.initialize(inputs.get_properties()) + + return None diff --git a/examples/objectdetection-example/output/object_detection_detected.png b/examples/objectdetection-example/output/object_detection_detected.png new file mode 100644 index 0000000..bb2edcc Binary files /dev/null and b/examples/objectdetection-example/output/object_detection_detected.png differ diff --git a/examples/objectdetection-example/pom.xml b/examples/objectdetection-example/pom.xml index 79d8365..bac5870 100644 --- a/examples/objectdetection-example/pom.xml +++ b/examples/objectdetection-example/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.23 + 1.0.24 smartai.examples.objectdetection.ObjectDetection @@ -34,7 +34,7 @@ cn.smartjavaai - smartjavaai-bom + bom ${smartjavaai.version} pom @@ -94,7 +94,7 @@ cn.smartjavaai - smartjavaai-objectdetection + vision diff --git a/examples/objectdetection-example/src/main/java/smartai/examples/objectdetection/ObjectDetection.java b/examples/objectdetection-example/src/main/java/smartai/examples/objectdetection/ObjectDetection.java index 20663f8..330559c 100644 --- a/examples/objectdetection-example/src/main/java/smartai/examples/objectdetection/ObjectDetection.java +++ b/examples/objectdetection-example/src/main/java/smartai/examples/objectdetection/ObjectDetection.java @@ -2,6 +2,9 @@ package smartai.examples.objectdetection; import ai.djl.Application; import ai.djl.MalformedModelException; +import ai.djl.ModelException; +import ai.djl.inference.Predictor; +import ai.djl.modality.Classifications; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.ImageFactory; import ai.djl.modality.cv.output.*; @@ -11,6 +14,7 @@ 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 ai.djl.translate.TranslateException; import cn.smartjavaai.common.config.Config; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionRectangle; @@ -42,6 +46,7 @@ import java.awt.*; import java.awt.image.BufferedImage; import java.io.File; import java.io.IOException; +import java.net.URL; import java.nio.file.Paths; import java.util.*; import java.util.List; @@ -64,6 +69,32 @@ public class ObjectDetection { //设备类型 public static DeviceEnum device = DeviceEnum.CPU; + public static void main(String[] args) throws ModelException, TranslateException, IOException { + Classifications classification = predict(); + log.info("{}", classification); + } + + + public static Classifications predict() throws IOException, ModelException, TranslateException { + + Config.setCachePath("/Users/wenjie/smartjavaai_cache"); + URL url = new URL("https://resources.djl.ai/images/action_dance.jpg"); + // Use DJL PyTorch model zoo model + Criteria criteria = + Criteria.builder() + .setTypes(URL.class, Classifications.class) + .optModelUrls( + "djl://ai.djl.mxnet/action_recognition") + .optEngine("MXNet") + .optProgress(new ProgressBar()) + .build(); + + try (ZooModel inception = criteria.loadModel(); + Predictor action = inception.newPredictor()) { + return action.predict(url); + } + } + @BeforeClass public static void beforeAll() throws IOException { //修改缓存路径 @@ -95,6 +126,8 @@ public class ObjectDetection { try { DetectorModelConfig config = new DetectorModelConfig(); config.setModelEnum(DetectorModelEnum.SSD_300_RESNET50);//检测模型,目前支持19种预置模型 + config.setModelEnum(DetectorModelEnum.YOLOV12_OFFICIAL); + config.setModelPath("yolov11s"); // 指定允许的类别 // config.setAllowedClasses(Arrays.asList("person")); //指定返回检测数量 @@ -205,6 +238,32 @@ public class ObjectDetection { } } + /** + * tensorflow目标检测 + */ + @Test + public void objectDetection3(){ + try { + DetectorModelConfig config = new DetectorModelConfig(); + config.setModelEnum(DetectorModelEnum.TENSORFLOW2_OFFICIAL); + config.setModelPath("/Users/wenjie/Documents/develop/model/tensorflow/ssd_mobilenet_v2_320x320_coco17_tpu-8"); +// config.putCustomParam("synsetUrl", "https://raw.githubusercontent.com/tensorflow/models/master/research/object_detection/data/mscoco_label_map.pbtxt"); +// config.putCustomParam("synsetPath", "/Users/wenjie/Downloads/mscoco_label_map.pbtxt.txt"); + config.putCustomParam("synsetFileName", "mscoco.pbtxt"); + // 指定允许的类别 +// config.setAllowedClasses(Arrays.asList("person")); + //指定返回检测数量 + config.setTopK(100); + config.setDevice(device); + DetectorModel detectorModel = ObjectDetectionModelFactory.getInstance().getModel(config); + DetectionResponse detectionResponse = detectorModel.detect("src/main/resources/dog_bike_car.jpg"); + detectorModel.detectAndDraw("src/main/resources/dog_bike_car.jpg", "output/dog_bike_car_detect.jpg"); + log.info("目标检测结果:{}", JSONObject.toJSONString(detectionResponse)); + } catch (Exception e) { + e.printStackTrace(); + } + } + /** * 摄像头目标检测 diff --git a/examples/ocr-examples/output/ocr_4_recognized.jpg b/examples/ocr-examples/output/ocr_4_recognized.jpg new file mode 100644 index 0000000..9d5dfa1 Binary files /dev/null and b/examples/ocr-examples/output/ocr_4_recognized.jpg differ diff --git a/examples/ocr-examples/output/plate_recognized.jpg b/examples/ocr-examples/output/plate_recognized.jpg new file mode 100644 index 0000000..36afe53 Binary files /dev/null and b/examples/ocr-examples/output/plate_recognized.jpg differ diff --git a/examples/ocr-examples/output/plate_recognized2.jpg b/examples/ocr-examples/output/plate_recognized2.jpg new file mode 100644 index 0000000..c69dec0 Binary files /dev/null and b/examples/ocr-examples/output/plate_recognized2.jpg differ diff --git a/examples/ocr-examples/output/table_ch2_result.html b/examples/ocr-examples/output/table_ch2_result.html new file mode 100644 index 0000000..28da28b --- /dev/null +++ b/examples/ocr-examples/output/table_ch2_result.html @@ -0,0 +1,5 @@ + +
主要财务比率202020212022E2023E2024E
成长能力
营业收入97.08%33.28%65.00%42.10%21.00%
营业利润165.21%22.38%31.65%64.55%36.68%
归属於母公司净利润164.75%24.17%39.44%64.13%38.63%
获利能力
毛利率25.45%23.01%16.80%17.00%18.00%
净利率13.98%13.03%11.01%12.72%14.57%
ROE19.29%19.25%20.77%47.11%35.24%
ROIC44.53%41.55%44.21%32.59%62.14%
偿债能力
资产负债率48.28%54.90%57.79%65.62%58.84%
净负债率-39.12%-36.03%6.62%8.70%5.28%
流动比率1.771.741.601.411.65
速动比率1.261.070.850.620.81
营运能力
应收账款周转率5.164.594.115.245.24
存货周转率3.482.892.552.772.63
总资产周转率0.800.780.931.211.22
每股指标(元)
每股收益0.841.041.452.383.30
每股经营现金流0.030.04-2.544.28-1.13
每股净资产4.345.406.975.059.35
估值比率
市盈率41.3033.2623.8514.5310.48
市净率7.976.404.956.853.69
EV/EBITDA5.0822.7223.6514.4010.60
EV/EBIT5.3324.1925.4515.0510.95
\ No newline at end of file diff --git a/examples/ocr-examples/output/table_ch2_result.jpg b/examples/ocr-examples/output/table_ch2_result.jpg new file mode 100644 index 0000000..fe84cb2 Binary files /dev/null and b/examples/ocr-examples/output/table_ch2_result.jpg differ diff --git a/examples/ocr-examples/output/table_ch2_result.xls b/examples/ocr-examples/output/table_ch2_result.xls new file mode 100644 index 0000000..3d5f319 Binary files /dev/null and b/examples/ocr-examples/output/table_ch2_result.xls differ diff --git a/examples/ocr-examples/output/table_ch2_result2.xls b/examples/ocr-examples/output/table_ch2_result2.xls new file mode 100644 index 0000000..3d5f319 Binary files /dev/null and b/examples/ocr-examples/output/table_ch2_result2.xls differ diff --git a/examples/ocr-examples/pom.xml b/examples/ocr-examples/pom.xml index e9b4155..7d329ff 100644 --- a/examples/ocr-examples/pom.xml +++ b/examples/ocr-examples/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.23 + 1.0.24 smartai.examples.ocr.common.OcrRecognizeDemo diff --git a/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrRecognizeDemo.java b/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrRecognizeDemo.java index 583d08e..1f2ca9f 100644 --- a/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrRecognizeDemo.java +++ b/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrRecognizeDemo.java @@ -61,7 +61,7 @@ public class OcrRecognizeDemo { //指定文本识别模型 recModelConfig.setRecModelEnum(CommonRecModelEnum.PP_OCR_V5_MOBILE_REC_MODEL); //指定识别模型位置,需要更改为自己的模型路径(下载地址请查看文档) - recModelConfig.setRecModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-OCRv5_mobile_rec_infer/PP-OCRv5_mobile_rec_infer.onnx"); + recModelConfig.setRecModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-OCRv5_server_rec_infer/PP-OCRv5_server_rec.onnx"); recModelConfig.setDevice(device); recModelConfig.setTextDetModel(getDetectionModel()); return OcrModelFactory.getInstance().getRecModel(recModelConfig); @@ -76,7 +76,7 @@ public class OcrRecognizeDemo { //指定检测模型 config.setModelEnum(CommonDetModelEnum.PP_OCR_V5_MOBILE_DET_MODEL); //指定模型位置,需要更改为自己的模型路径(下载地址请查看文档) - config.setDetModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-OCRv5_mobile_det_infer/PP-OCRv5_mobile_det_infer.onnx"); + config.setDetModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-OCRv5_server_det_infer/PP-OCRv5_server_det.onnx"); config.setDevice(device); return OcrModelFactory.getInstance().getDetModel(config); } @@ -127,7 +127,7 @@ public class OcrRecognizeDemo { OcrCommonRecModel recModel = getRecModel(); //不带方向矫正,分行返回文本 OcrRecOptions options = new OcrRecOptions(false, true); - OcrInfo ocrInfo = recModel.recognize("src/main/resources/ocr_2.jpg",options); + OcrInfo ocrInfo = recModel.recognize("/Users/wenjie/Downloads/49421755855753_.pic_hd.jpg",options); log.info("OCR识别结果:{}", JSONObject.toJSONString(ocrInfo)); } catch (Exception e) { e.printStackTrace(); diff --git a/examples/speech-examples/pom.xml b/examples/speech-examples/pom.xml index 3cf90ba..cf980e1 100644 --- a/examples/speech-examples/pom.xml +++ b/examples/speech-examples/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.23 + 1.0.24 smartai.examples.speech.asr.common.OcrRecognizeDemo diff --git a/examples/translation-example/pom.xml b/examples/translation-example/pom.xml index 7516c65..d0bafb7 100644 --- a/examples/translation-example/pom.xml +++ b/examples/translation-example/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.23 + 1.0.24 smartai.examples.nlp.translation.TranslationDemo diff --git a/smartjavaai-face/pom.xml b/face/pom.xml similarity index 96% rename from smartjavaai-face/pom.xml rename to face/pom.xml index 4e74ce4..fac0167 100644 --- a/smartjavaai-face/pom.xml +++ b/face/pom.xml @@ -6,12 +6,12 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - smartjavaai-face - 1.0.23 - smartjavaai-face + face + 1.0.24 + face SmartJavaAI https://github.com/geekwenjie/SmartJavaAI @@ -34,7 +34,7 @@ cn.smartjavaai - smartjavaai-common + common ${project.version} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceAttributeConfig.java b/face/src/main/java/cn/smartjavaai/face/config/FaceAttributeConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceAttributeConfig.java rename to face/src/main/java/cn/smartjavaai/face/config/FaceAttributeConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceDetConfig.java b/face/src/main/java/cn/smartjavaai/face/config/FaceDetConfig.java similarity index 92% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceDetConfig.java rename to face/src/main/java/cn/smartjavaai/face/config/FaceDetConfig.java index 4a23fc6..4f08060 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceDetConfig.java +++ b/face/src/main/java/cn/smartjavaai/face/config/FaceDetConfig.java @@ -24,7 +24,7 @@ public class FaceDetConfig extends ModelConfig { /** * 置信度阈值 */ - private double confidenceThreshold = FaceDetectConstant.DEFAULT_CONFIDENCE_THRESHOLD; + private double confidenceThreshold; /** diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceExpressionConfig.java b/face/src/main/java/cn/smartjavaai/face/config/FaceExpressionConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceExpressionConfig.java rename to face/src/main/java/cn/smartjavaai/face/config/FaceExpressionConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceRecConfig.java b/face/src/main/java/cn/smartjavaai/face/config/FaceRecConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/config/FaceRecConfig.java rename to face/src/main/java/cn/smartjavaai/face/config/FaceRecConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/LivenessConfig.java b/face/src/main/java/cn/smartjavaai/face/config/LivenessConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/config/LivenessConfig.java rename to face/src/main/java/cn/smartjavaai/face/config/LivenessConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/config/QualityConfig.java b/face/src/main/java/cn/smartjavaai/face/config/QualityConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/config/QualityConfig.java rename to face/src/main/java/cn/smartjavaai/face/config/QualityConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/FaceDetectConstant.java b/face/src/main/java/cn/smartjavaai/face/constant/FaceDetectConstant.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/FaceDetectConstant.java rename to face/src/main/java/cn/smartjavaai/face/constant/FaceDetectConstant.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/FaceNetConstant.java b/face/src/main/java/cn/smartjavaai/face/constant/FaceNetConstant.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/FaceNetConstant.java rename to face/src/main/java/cn/smartjavaai/face/constant/FaceNetConstant.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/LivenessConstant.java b/face/src/main/java/cn/smartjavaai/face/constant/LivenessConstant.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/LivenessConstant.java rename to face/src/main/java/cn/smartjavaai/face/constant/LivenessConstant.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/MiniVisionConstant.java b/face/src/main/java/cn/smartjavaai/face/constant/MiniVisionConstant.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/MiniVisionConstant.java rename to face/src/main/java/cn/smartjavaai/face/constant/MiniVisionConstant.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/RetinaFaceConstant.java b/face/src/main/java/cn/smartjavaai/face/constant/RetinaFaceConstant.java similarity index 77% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/RetinaFaceConstant.java rename to face/src/main/java/cn/smartjavaai/face/constant/RetinaFaceConstant.java index de02817..2a23f43 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/RetinaFaceConstant.java +++ b/face/src/main/java/cn/smartjavaai/face/constant/RetinaFaceConstant.java @@ -20,9 +20,5 @@ public class RetinaFaceConstant { */ public static final double[] variance = {0.1f, 0.2f}; - /** - * 模型下载地址 - */ - public static final String MODEL_URL = "https://resources.djl.ai/test-models/pytorch/retinaface.zip"; } diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/UltraLightFastGenericFaceConstant.java b/face/src/main/java/cn/smartjavaai/face/constant/UltraLightFastGenericFaceConstant.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/constant/UltraLightFastGenericFaceConstant.java rename to face/src/main/java/cn/smartjavaai/face/constant/UltraLightFastGenericFaceConstant.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/context/PredictorContext.java b/face/src/main/java/cn/smartjavaai/face/context/PredictorContext.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/context/PredictorContext.java rename to face/src/main/java/cn/smartjavaai/face/context/PredictorContext.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/dao/FaceDao.java b/face/src/main/java/cn/smartjavaai/face/dao/FaceDao.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/dao/FaceDao.java rename to face/src/main/java/cn/smartjavaai/face/dao/FaceDao.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceQualityResult.java b/face/src/main/java/cn/smartjavaai/face/entity/FaceQualityResult.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceQualityResult.java rename to face/src/main/java/cn/smartjavaai/face/entity/FaceQualityResult.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceQualitySummary.java b/face/src/main/java/cn/smartjavaai/face/entity/FaceQualitySummary.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceQualitySummary.java rename to face/src/main/java/cn/smartjavaai/face/entity/FaceQualitySummary.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceRegisterInfo.java b/face/src/main/java/cn/smartjavaai/face/entity/FaceRegisterInfo.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceRegisterInfo.java rename to face/src/main/java/cn/smartjavaai/face/entity/FaceRegisterInfo.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceResult.java b/face/src/main/java/cn/smartjavaai/face/entity/FaceResult.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceResult.java rename to face/src/main/java/cn/smartjavaai/face/entity/FaceResult.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceSearchParams.java b/face/src/main/java/cn/smartjavaai/face/entity/FaceSearchParams.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/entity/FaceSearchParams.java rename to face/src/main/java/cn/smartjavaai/face/entity/FaceSearchParams.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/ExpressionModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/ExpressionModelEnum.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/ExpressionModelEnum.java rename to face/src/main/java/cn/smartjavaai/face/enums/ExpressionModelEnum.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceAttributeModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/FaceAttributeModelEnum.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceAttributeModelEnum.java rename to face/src/main/java/cn/smartjavaai/face/enums/FaceAttributeModelEnum.java diff --git a/face/src/main/java/cn/smartjavaai/face/enums/FaceDetModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/FaceDetModelEnum.java new file mode 100644 index 0000000..c717736 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/enums/FaceDetModelEnum.java @@ -0,0 +1,93 @@ +package cn.smartjavaai.face.enums; + +import lombok.Data; + +/** + * 人脸检测模型枚举 + * @author dwj + */ +public enum FaceDetModelEnum { + + RETINA_FACE("PyTorch",0,0, "https://resources.djl.ai/test-models/pytorch/retinaface.zip"), + RETINA_FACE_ONNX("OnnxRuntime",0,0, null), + RETINA_FACE_640_ONNX("OnnxRuntime",640,640, null), + RETINA_FACE_320_ONNX("OnnxRuntime",320,320, null), + RETINA_FACE_720_1280_ONNX("OnnxRuntime",720,1280, null), + RETINA_FACE_MOBILE_ONNX("OnnxRuntime",0,0, null), + RETINA_FACE_MOBILE_320_ONNX("OnnxRuntime",320,320, null), + RETINA_FACE_MOBILE_640_ONNX("OnnxRuntime",640,640, null), + RETINA_FACE_MOBILE_720_1280_ONNX("OnnxRuntime",720,1080, null), + ULTRA_LIGHT_FAST_GENERIC_FACE("PyTorch",0,0, "https://resources.djl.ai/test-models/pytorch/ultranet.zip"), + SEETA_FACE6_MODEL(null,0,0, null), + YOLOV8_FACE("OnnxRuntime",640,640, null), + YOLOV5_FACE_640("OnnxRuntime", 640,640, null), + YOLOV5_FACE_320("OnnxRuntime",320,320, null), + SCRFD_160("OnnxRuntime",160,160, null), + SCRFD_320("OnnxRuntime",320,320, null), + SCRFD_640("OnnxRuntime",640,640, null), + SCRFD_1280("OnnxRuntime",1280,1280, null), + + MTCNN("OnnxRuntime",1280,1280, null); + + + + + /** + * 模型输入尺寸:宽 + */ + private final int inputWidth; + + /** + * 模型输入尺寸:高 + */ + private final int inputHeight; + + /** + * 模型地址 + */ + private final String modelUrl; + + /** + * 模型引擎 + */ + private final String engine; + + FaceDetModelEnum(String engine, int inputWidth, int inputHeight, String modelUrl) { + this.inputWidth = inputWidth; + this.inputHeight = inputHeight; + this.modelUrl = modelUrl; + this.engine = engine; + } + + public int getInputWidth() { + return inputWidth; + } + + public int getInputHeight() { + return inputHeight; + } + + public String getModelUrl() { + return modelUrl; + } + + public String getEngine() { + return engine; + } + + + /** + * 根据名称获取枚举 (忽略大小写和下划线变体) + */ + public static FaceDetModelEnum fromName(String name) { + String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); + for (FaceDetModelEnum model : values()) { + if (model.name().replaceAll("_", "").equals(formatted)) { + return model; + } + } + throw new IllegalArgumentException("未知模型名称: " + name); + } + + +} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java similarity index 95% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java rename to face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java index 20f446f..562d1fc 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java +++ b/face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java @@ -8,6 +8,7 @@ public enum FaceRecModelEnum { FACENET_MODEL("FaceNetModel"), SEETA_FACE6_MODEL("SeetaFace6Model"), + SEETA_FACE6_LIGHT_MODEL("SeetaFace6Model"), INSIGHT_FACE_IRSE50_MODEL("InsightFaceIRSE50Model"), INSIGHT_FACE_MOBILE_FACENET_MODEL("InsightFaceMobilefacenetModel"), ELASTIC_FACE_MODEL("ElasticFaceModel"); diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/IdStrategy.java b/face/src/main/java/cn/smartjavaai/face/enums/IdStrategy.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/IdStrategy.java rename to face/src/main/java/cn/smartjavaai/face/enums/IdStrategy.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/LivenessModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/LivenessModelEnum.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/LivenessModelEnum.java rename to face/src/main/java/cn/smartjavaai/face/enums/LivenessModelEnum.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/QualityGrade.java b/face/src/main/java/cn/smartjavaai/face/enums/QualityGrade.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/QualityGrade.java rename to face/src/main/java/cn/smartjavaai/face/enums/QualityGrade.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/QualityModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/QualityModelEnum.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/QualityModelEnum.java rename to face/src/main/java/cn/smartjavaai/face/enums/QualityModelEnum.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/SimilarityType.java b/face/src/main/java/cn/smartjavaai/face/enums/SimilarityType.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/SimilarityType.java rename to face/src/main/java/cn/smartjavaai/face/enums/SimilarityType.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/VectorDBType.java b/face/src/main/java/cn/smartjavaai/face/enums/VectorDBType.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/VectorDBType.java rename to face/src/main/java/cn/smartjavaai/face/enums/VectorDBType.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/exception/FaceException.java b/face/src/main/java/cn/smartjavaai/face/exception/FaceException.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/exception/FaceException.java rename to face/src/main/java/cn/smartjavaai/face/exception/FaceException.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java rename to face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java rename to face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java similarity index 76% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java rename to face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java index f12f77f..4f09aac 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java @@ -25,12 +25,12 @@ public class FaceDetModelFactory { // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 private static volatile FaceDetModelFactory instance; - private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); + private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); /** * 模型注册表 */ - private static final Map> registry = + private static final Map> registry = new ConcurrentHashMap<>(); @@ -49,11 +49,11 @@ public class FaceDetModelFactory { /** * 注册模型 - * @param name + * @param faceDetModelEnum * @param clazz */ - private static void registerAlgorithm(String name, Class clazz) { - registry.put(name.toLowerCase(), clazz); + private static void registerAlgorithm(FaceDetModelEnum faceDetModelEnum, Class clazz) { + registry.put(faceDetModelEnum, clazz); } @@ -66,7 +66,7 @@ public class FaceDetModelFactory { if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){ throw new FaceException("未配置人脸模型"); } - return modelMap.computeIfAbsent(config.getModelEnum().name(), k -> { + return modelMap.computeIfAbsent(config.getModelEnum(), k -> { return createFaceDetModel(config); }); } @@ -90,7 +90,7 @@ public class FaceDetModelFactory { * @return */ private FaceDetModel createFaceDetModel(FaceDetConfig config) { - Class clazz = registry.get(config.getModelEnum().getModelClassName().toLowerCase()); + Class clazz = registry.get(config.getModelEnum()); if(clazz == null){ throw new FaceException("Unsupported algorithm"); } @@ -121,9 +121,12 @@ public class FaceDetModelFactory { // 初始化默认算法 static { - registerAlgorithm("retinafacemodel", CommonFaceDetModel.class); - registerAlgorithm("ultralightfastgenericfacemodel", CommonFaceDetModel.class); - registerAlgorithm("seetaface6model", SeetaFace6FaceDetModel.class); + registerAlgorithm(FaceDetModelEnum.RETINA_FACE, CommonFaceDetModel.class); + registerAlgorithm(FaceDetModelEnum.ULTRA_LIGHT_FAST_GENERIC_FACE, CommonFaceDetModel.class); + registerAlgorithm(FaceDetModelEnum.SEETA_FACE6_MODEL, SeetaFace6FaceDetModel.class); + registerAlgorithm(FaceDetModelEnum.YOLOV8_FACE, CommonFaceDetModel.class); + registerAlgorithm(FaceDetModelEnum.YOLOV5_FACE_640, CommonFaceDetModel.class); + registerAlgorithm(FaceDetModelEnum.YOLOV5_FACE_320, CommonFaceDetModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java rename to face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java similarity index 96% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java rename to face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java index 502cb70..25cf08c 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java @@ -97,6 +97,7 @@ public class FaceRecModelFactory { registerAlgorithm(FaceRecModelEnum.INSIGHT_FACE_MOBILE_FACENET_MODEL, CommonFaceRecModel.class); registerAlgorithm(FaceRecModelEnum.ELASTIC_FACE_MODEL, CommonFaceRecModel.class); registerAlgorithm(FaceRecModelEnum.SEETA_FACE6_MODEL, SeetaFace6FaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.SEETA_FACE6_LIGHT_MODEL, SeetaFace6FaceRecModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java rename to face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java b/face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java rename to face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java b/face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java rename to face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java b/face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java rename to face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java b/face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java rename to face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/criterial/EmotionCriteriaFactory.java b/face/src/main/java/cn/smartjavaai/face/model/expression/criterial/EmotionCriteriaFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/criterial/EmotionCriteriaFactory.java rename to face/src/main/java/cn/smartjavaai/face/model/expression/criterial/EmotionCriteriaFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/translator/DenseNetEmotionTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/expression/translator/DenseNetEmotionTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/translator/DenseNetEmotionTranslator.java rename to face/src/main/java/cn/smartjavaai/face/model/expression/translator/DenseNetEmotionTranslator.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/translator/FrEmotionTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/expression/translator/FrEmotionTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/expression/translator/FrEmotionTranslator.java rename to face/src/main/java/cn/smartjavaai/face/model/expression/translator/FrEmotionTranslator.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java rename to face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java similarity index 99% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java rename to face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java index 475cd7d..c10aab2 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java @@ -93,5 +93,4 @@ public interface FaceDetModel extends AutoCloseable{ } - } diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/MtcnnFaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/MtcnnFaceDetModel.java new file mode 100644 index 0000000..673594f --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/MtcnnFaceDetModel.java @@ -0,0 +1,399 @@ +package cn.smartjavaai.face.model.facedect; + +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.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.ndarray.NDArray; +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.ZooModel; +import ai.djl.training.util.ProgressBar; +import ai.djl.translate.NoopTranslator; +import cn.smartjavaai.common.entity.*; +import cn.smartjavaai.common.entity.Point; +import cn.smartjavaai.common.entity.face.FaceInfo; +import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.Base64ImageUtils; +import cn.smartjavaai.common.utils.FileUtils; +import cn.smartjavaai.common.utils.ImageUtils; +import cn.smartjavaai.common.utils.OpenCVUtils; +import cn.smartjavaai.face.config.FaceDetConfig; +import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.model.facedect.criterial.FaceDetCriteriaFactory; +import cn.smartjavaai.face.model.facedect.mtcnn.*; +import cn.smartjavaai.face.utils.FaceUtils; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.collections.CollectionUtils; +import org.apache.commons.lang3.StringUtils; +import org.apache.commons.pool2.impl.GenericObjectPool; +import org.opencv.core.Mat; + +import javax.imageio.ImageIO; +import java.awt.*; +import java.awt.image.BufferedImage; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +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.Iterator; +import java.util.List; +import java.util.Objects; + +/** + * MTCNN 人脸检测模型实现 + * @author dwj + */ +@Slf4j +public class MtcnnFaceDetModel implements FaceDetModel{ + + + public ZooModel pNetModel; + + public ZooModel rNetModel; + + public ZooModel oNetModel; + private GenericObjectPool> pnetPredictorPool; + private GenericObjectPool> rnetPredictorPool; + private GenericObjectPool> onetPredictorPool; + + + /** + * 加载模型 + * @param config + */ + @Override + public void loadModel(FaceDetConfig config){ + if(StringUtils.isBlank(config.getModelPath())){ + throw new FaceException("modelPath is null"); + } + Path modelPath = Paths.get(config.getModelPath()); + if(!Files.isDirectory(modelPath)){ + throw new FaceException("MTCNN 模型需要指定存放模型文件的目录路径"); + } + try { + Path pnetPath = modelPath.resolve("pnet_script.pt"); + Path rnetPath = modelPath.resolve("rnet_script.pt"); + Path onetPath = modelPath.resolve("onet_script.pt"); + pNetModel = getModel(pnetPath); + rNetModel = getModel(pnetPath); + oNetModel = getModel(pnetPath); + + this.pnetPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(pNetModel)); + this.rnetPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(rNetModel)); + this.onetPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(oNetModel)); + int predictorPoolSize = config.getPredictorPoolSize(); + if(config.getPredictorPoolSize() <= 0){ + predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数 + } + pnetPredictorPool.setMaxTotal(predictorPoolSize); + rnetPredictorPool.setMaxTotal(predictorPoolSize); + onetPredictorPool.setMaxTotal(predictorPoolSize); + log.debug("当前设备: " + pNetModel.getNDManager().getDevice()); + log.debug("当前引擎: " + Engine.getInstance().getEngineName()); + log.debug("模型推理器线程池最大数量: " + predictorPoolSize); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new FaceException("mtcnn人脸检测模型加载失败", e); + } + } + + /** + * 加载模型 + * @param modelPath + * @throws ModelNotFoundException + * @throws MalformedModelException + * @throws IOException + */ + public ZooModel getModel(Path modelPath) throws ModelNotFoundException, MalformedModelException, IOException { + Criteria criteria = + Criteria.builder() + .setTypes(NDList.class, NDList.class) + .optTranslator(new NoopTranslator()) + .optEngine("PyTorch") + .optModelPath(modelPath) + .optProgress(new ProgressBar()) + .build(); + return criteria.loadModel(); + } + + + + /** + * 检测人脸 + * @param imagePath 图片路径 + * @return + * @throws Exception + */ + @Override + public R 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 FaceException("无效的图片", e); + } finally { + if (img != null) { + ((Mat)img.getWrappedImage()).release(); + } + } + + } + + /** + * 检测人脸 + * @param imageInputStream 图片流 + * @return + * @throws Exception + */ + @Override + public R detect(InputStream imageInputStream){ + if(Objects.isNull(imageInputStream)){ + return R.fail(R.Status.INVALID_IMAGE); + } + Image img = null; + try { + img = ImageFactory.getInstance().fromInputStream(imageInputStream); + return detect(img); + } catch (IOException e) { + throw new FaceException("无效图片输入流", e); + } finally { + if (img != null) { + ((Mat)img.getWrappedImage()).release(); + } + } + } + + @Override + public R 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 FaceException(e); + } finally { + if (img != null) { + ((Mat)img.getWrappedImage()).release(); + } + } + + } + + @Override + public R detect(byte[] imageData) { + if(Objects.isNull(imageData)){ + return R.fail(R.Status.INVALID_IMAGE); + } + return detect(new ByteArrayInputStream(imageData)); + } + + @Override + public R detectBase64(String base64Image) { + if(StringUtils.isBlank(base64Image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); + return detect(imageData); + } + + @Override + public R detectAndDraw(String imagePath, String outputPath) { + if(!FileUtils.isFileExists(imagePath)){ + return R.fail(R.Status.FILE_NOT_FOUND); + } + Image img = null; + try { + img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + R detectionResponseR = detect(img); + if(!detectionResponseR.isSuccess()){ + return R.fail(detectionResponseR.getCode(), detectionResponseR.getMessage()); + } + if(Objects.isNull(detectionResponseR.getData()) || + CollectionUtils.isEmpty(detectionResponseR.getData().getDetectionInfoList())){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + BufferedImage sourceImage = OpenCVUtils.mat2Image((Mat)img.getWrappedImage()); + FaceUtils.drawBoundingBoxes(sourceImage, detectionResponseR.getData(), outputPath); + return R.ok(); + } catch (IOException e) { + throw new FaceException(e); + } finally { + if (img != null){ + ((Mat)img.getWrappedImage()).release(); + } + } + } + + @Override + public R detectAndDraw(BufferedImage sourceImage) { + if(!ImageUtils.isImageValid(sourceImage)){ + return R.fail(R.Status.INVALID_IMAGE); + } + try { + R detectionResponseR = detect(sourceImage); + if(!detectionResponseR.isSuccess()){ + return R.fail(detectionResponseR.getCode(), detectionResponseR.getMessage()); + } + if(Objects.isNull(detectionResponseR.getData()) || + CollectionUtils.isEmpty(detectionResponseR.getData().getDetectionInfoList())){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + return R.ok(FaceUtils.drawBoundingBoxes(sourceImage, detectionResponseR.getData())); + } catch (IOException e) { + throw new FaceException("导出图片失败", e); + } + } + + /** + * 人脸检测 + * @param image + * @return + */ + public R detect(Image image){ + try (NDManager manager = NDManager.newBaseManager(pNetModel.getNDManager().getDevice())){ + List scales = MtcnnProcess.generateScales(image); + NDArray imgs = MtcnnProcess.processInput(manager, image); + int h = image.getHeight(); + int w = image.getWidth(); + NDList outputPnet = PNetModel.firstStage(manager, pnetPredictorPool.borrowObject(), imgs, scales, w, h); + NDArray boxes = outputPnet.get(0); + NDArray image_inds = outputPnet.get(1); + NDList pad = MtcnnUtils.pad(boxes, w, h); + NDList outputRnet = RNetModel.secondStage(manager, rnetPredictorPool.borrowObject(), imgs,boxes,pad, image_inds); + NDArray image_indsFiltered = outputRnet.get(0); + NDArray scoresFiltered = outputRnet.get(1); + MtcnnBatchResult oNetResult = ONetModel.thirdStage(manager, onetPredictorPool.borrowObject(), imgs,boxes, w, h, scoresFiltered, image_indsFiltered); + return R.ok(convertToDetectionResponse(oNetResult)); + } catch (Exception e) { + throw new RuntimeException(e); + } + } + + + + + /** + * 转换为FaceDetectedResult + * @param mtcnnBatchResult + * @return + */ + public static DetectionResponse convertToDetectionResponse(MtcnnBatchResult mtcnnBatchResult){ + if(Objects.isNull(mtcnnBatchResult) || CollectionUtils.isEmpty(mtcnnBatchResult.boxes) + || CollectionUtils.isEmpty(mtcnnBatchResult.points) + || CollectionUtils.isEmpty(mtcnnBatchResult.probs)){ + return null; + } + DetectionResponse detectionResponse = new DetectionResponse(); + List detectionInfoList = new ArrayList(); + + NDArray boxes = mtcnnBatchResult.boxes.get(0); + NDArray probs = mtcnnBatchResult.probs.get(0); + NDArray points = mtcnnBatchResult.points.get(0); + + long numBoxes = boxes.getShape().get(0); + for (int i = 0; i < numBoxes; i++) { + float[] boxCoords = boxes.get(i).toFloatArray(); // [x1, y1, x2, y2] + float score = probs.getFloat(i); + NDArray pointND = points.get(i); // shape [5,2] + float[] flatPoints = pointND.toFloatArray(); // 一维长度 10 + List keyPoints = new ArrayList(); + for (int p = 0; p < 5; p++) { + keyPoints.add(new Point(flatPoints[p * 2], flatPoints[p * 2 + 1])); + } + int x = Math.round(boxCoords[0]); + int y = Math.round(boxCoords[1]); + int w = Math.round(boxCoords[2] - boxCoords[0]); + int h = Math.round(boxCoords[3] - boxCoords[1]); + + DetectionRectangle rectangle = new DetectionRectangle(x, y, w, h); + FaceInfo faceInfo = new FaceInfo(keyPoints); + DetectionInfo detectionInfo = new DetectionInfo(rectangle, score, faceInfo); + detectionInfoList.add(detectionInfo); + } + detectionResponse.setDetectionInfoList(detectionInfoList); + return detectionResponse; + } + + + + public GenericObjectPool> getPnetPredictorPool() { + return pnetPredictorPool; + } + + public GenericObjectPool> getRnetPredictorPool() { + return rnetPredictorPool; + } + + public GenericObjectPool> getOnetPredictorPool() { + return onetPredictorPool; + } + + @Override + public void close() { + try { + if (pnetPredictorPool != null) { + pnetPredictorPool.close(); + } + } catch (Exception e) { + log.warn("关闭 predictorPool 失败", e); + } + try { + if (rnetPredictorPool != null) { + rnetPredictorPool.close(); + } + } catch (Exception e) { + log.warn("关闭 predictorPool 失败", e); + } + try { + if (onetPredictorPool != null) { + onetPredictorPool.close(); + } + } catch (Exception e) { + log.warn("关闭 predictorPool 失败", e); + } + try { + if (pNetModel != null) { + pNetModel.close(); + } + } catch (Exception e) { + log.warn("关闭 model 失败", e); + } + try { + if (pNetModel != null) { + pNetModel.close(); + } + } catch (Exception e) { + log.warn("关闭 model 失败", e); + } + try { + if (rNetModel != null) { + rNetModel.close(); + } + } catch (Exception e) { + log.warn("关闭 model 失败", e); + } + try { + if (oNetModel != null) { + oNetModel.close(); + } + } catch (Exception e) { + log.warn("关闭 model 失败", e); + } + } +} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java similarity index 97% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java rename to face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java index ecba841..bdde257 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java @@ -36,6 +36,11 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ private FaceDetectorPool faceDetectorPool; private FaceLandmarkerPool faceLandmarkerPool; + /** + * 阈值 + */ + private static final double THRESHOLD = 0.9d; + @Override public void loadModel(FaceDetConfig config) { @@ -116,6 +121,7 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ FaceLandmarker faceLandmarker = null; try { predictor = faceDetectorPool.borrowObject(); + predictor.set(FaceDetector.Property.PROPERTY_THRESHOLD, config.getConfidenceThreshold() > 0 ? config.getConfidenceThreshold() : THRESHOLD); faceLandmarker = faceLandmarkerPool.borrowObject(); SeetaRect[] seetaResult = predictor.Detect(imageData); List seetaPointFSList = new ArrayList(); diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/criterial/FaceDetCriteriaFactory.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/criterial/FaceDetCriteriaFactory.java new file mode 100644 index 0000000..a8b0bd9 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/criterial/FaceDetCriteriaFactory.java @@ -0,0 +1,107 @@ +package cn.smartjavaai.face.model.facedect.criterial; + +import ai.djl.Device; +import ai.djl.modality.Classifications; +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 ai.djl.translate.Translator; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.face.config.FaceDetConfig; +import cn.smartjavaai.face.config.FaceExpressionConfig; +import cn.smartjavaai.face.constant.FaceDetectConstant; +import cn.smartjavaai.face.constant.RetinaFaceConstant; +import cn.smartjavaai.face.constant.UltraLightFastGenericFaceConstant; +import cn.smartjavaai.face.enums.ExpressionModelEnum; +import cn.smartjavaai.face.enums.FaceDetModelEnum; +import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.model.expression.translator.DenseNetEmotionTranslator; +import cn.smartjavaai.face.model.expression.translator.FrEmotionTranslator; +import cn.smartjavaai.face.translator.FaceDetectionTranslator; +import cn.smartjavaai.face.translator.SCRFDFaceTranslator; +import cn.smartjavaai.face.translator.YoloV5FaceTranslator; +import cn.smartjavaai.face.translator.YoloV8FaceTranslator; +import org.apache.commons.lang3.StringUtils; + +import java.nio.file.Paths; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; + +/** + * 人脸检测 Criteria构建工厂 + * @author dwj + */ +public class FaceDetCriteriaFactory { + + /** + * 创建人脸检测Criteria + * @param config + * @return + */ + public static Criteria createCriteria(FaceDetConfig config) { + Device device = null; + if(!Objects.isNull(config.getDevice())){ + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + Translator translator = getTranslator(config); + if(StringUtils.isBlank(config.getModelEnum().getModelUrl())){ + //检查模型路径 + if (StringUtils.isBlank(config.getModelPath())){ + throw new FaceException("请指定模型路径"); + } + } + Criteria criteria = + Criteria.builder() + .setTypes(Image.class, DetectedObjects.class) + .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : config.getModelEnum().getModelUrl()) + .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) + .optTranslator(translator) + .optDevice(device) + .optProgress(new ProgressBar()) + .optEngine(config.getModelEnum().getEngine()) + .build(); + return criteria; + } + + + /** + * 获取人脸检测Translator + * @param config + * @return + */ + public static Translator getTranslator(FaceDetConfig config) { + Translator translator = null; + if(config.getModelEnum() == FaceDetModelEnum.RETINA_FACE){ + translator = + new FaceDetectionTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), RetinaFaceConstant.variance, FaceDetectConstant.MAX_FACE_LIMIT, RetinaFaceConstant.scales, RetinaFaceConstant.steps); + }else if (config.getModelEnum() == FaceDetModelEnum.ULTRA_LIGHT_FAST_GENERIC_FACE){ + translator = + new FaceDetectionTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), UltraLightFastGenericFaceConstant.variance, FaceDetectConstant.MAX_FACE_LIMIT, UltraLightFastGenericFaceConstant.scales, UltraLightFastGenericFaceConstant.steps); + }else if (config.getModelEnum() == FaceDetModelEnum.YOLOV8_FACE){ + Map arguments = new HashMap<>(); + arguments.put("width", 640); + arguments.put("height", 640); + arguments.put("resizeShort", 640); + arguments.put("centerFit", true); + translator = YoloV8FaceTranslator.builder(arguments).build(); + }else if (config.getModelEnum() == FaceDetModelEnum.YOLOV5_FACE_640 + || config.getModelEnum() == FaceDetModelEnum.YOLOV5_FACE_320){ + Map arguments = new HashMap<>(); + arguments.put("width", config.getModelEnum().getInputWidth()); + arguments.put("height", config.getModelEnum().getInputHeight()); + arguments.put("resizeShort", true); + arguments.put("centerFit", true); + translator = YoloV5FaceTranslator.builder(arguments).build(); + }else if (config.getModelEnum() == FaceDetModelEnum.SCRFD_160 + || config.getModelEnum() == FaceDetModelEnum.SCRFD_320 + || config.getModelEnum() == FaceDetModelEnum.SCRFD_640 + || config.getModelEnum() == FaceDetModelEnum.SCRFD_1280){ + translator = + new SCRFDFaceTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), 100, new int[]{8, 16, 32}); + } + return translator; + } + +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnBatchResult.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnBatchResult.java new file mode 100644 index 0000000..f9e41e4 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnBatchResult.java @@ -0,0 +1,24 @@ +package cn.smartjavaai.face.model.facedect.mtcnn; + +import ai.djl.ndarray.NDArray; +import lombok.Data; + +import java.util.List; + +/** + * @author dwj + */ +@Data +public class MtcnnBatchResult { + + public List boxes; + public List probs; + public List points; + + public MtcnnBatchResult(List boxes, List probs, List points) { + this.boxes = boxes; + this.probs = probs; + this.points = points; + } + +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnProcess.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnProcess.java new file mode 100644 index 0000000..1b06fda --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnProcess.java @@ -0,0 +1,64 @@ +package cn.smartjavaai.face.model.facedect.mtcnn; + +import ai.djl.modality.cv.Image; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; + +import java.util.ArrayList; +import java.util.List; + +/** + * @author dwj + */ +public class MtcnnProcess { + + + /** + * 预处理 + * @param image + * @return + */ + public static NDArray processInput(NDManager manager, Image image){ + // (N, C, H, W) + // Image -> NDArray (H, W, C) + NDArray array = image.toNDArray(manager, Image.Flag.COLOR); + // 增加 batch 维度 (1, H, W, C) ---- + array = array.expandDims(0); + // 交换维度 (N, C, H, W) + array = array.transpose(0, 3, 1, 2); + // 转成模型的数据类型 + if (!array.getDataType().equals(DataType.FLOAT32)) { + array = array.toType(DataType.FLOAT32, false); + } + return array; + } + + /** + * 生成金字塔缩放比例列表 + * @param image + * @return + */ + public static List generateScales(Image image){ + long h = image.getHeight(); + long w = image.getWidth(); + // 计算最小缩放比例 + double minsize = 20; + double m = 12.0 / minsize; + double minl = Math.min(h, w) * m; + + // 创建金字塔缩放比例列表 + double factor = 0.709; // 你原代码的 factor + List scales = new ArrayList<>(); + double scale_i = m; + while (minl >= 12) { + scales.add(scale_i); + scale_i *= factor; + minl *= factor; + } + return scales; + } + + +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnUtils.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnUtils.java new file mode 100644 index 0000000..ec2e95e --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/MtcnnUtils.java @@ -0,0 +1,112 @@ +package cn.smartjavaai.face.model.facedect.mtcnn; + +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.index.NDIndex; +import ai.djl.ndarray.types.DataType; + +/** + * @author dwj + */ +public class MtcnnUtils { + + /** + * 矩形框重新映射 + * @param bboxA + * @return + */ + public static NDArray rerec(NDArray bboxA) { + // h = y2 - y1 + NDArray h = bboxA.get(":, 3").sub(bboxA.get(":, 1")); + // w = x2 - x1 + NDArray w = bboxA.get(":, 2").sub(bboxA.get(":, 0")); + + // l = max(w, h) + NDArray l = w.maximum(h); + + // x1 = x1 + w*0.5 - l*0.5 + NDArray x1 = bboxA.get(":, 0").add(w.mul(0.5)).sub(l.mul(0.5)); + // y1 = y1 + h*0.5 - l*0.5 + NDArray y1 = bboxA.get(":, 1").add(h.mul(0.5)).sub(l.mul(0.5)); + + // x2, y2 + NDArray x2 = x1.add(l); + NDArray y2 = y1.add(l); + + // 坐标变成 [N,1] 方便拼接 + NDArray coords = NDArrays.concat( + new NDList( + x1.expandDims(1), + y1.expandDims(1), + x2.expandDims(1), + y2.expandDims(1) + ), + 1 + ); + + // 保留原来的 score(bboxA[:, 4:]) + if (bboxA.getShape().get(1) > 4) { + NDArray rest = bboxA.get(":, 4:"); + return NDArrays.concat(new NDList(coords, rest), 1); + } else { + return coords; + } + } + + /** + * 限制范围 + * @param boxes + * @param w + * @param h + * @return + */ + public static NDList pad(NDArray boxes, int w, int h) { + // 去小数 -> 转 int + boxes = boxes.floor().toType(DataType.INT32, false); + NDArray x = boxes.get(":, 0"); + NDArray y = boxes.get(":, 1"); + NDArray ex = boxes.get(":, 2"); + NDArray ey = boxes.get(":, 3"); + // 限制范围 + x = x.maximum(1); + y = y.maximum(1); + ex = ex.minimum(w); + ey = ey.minimum(h); + return new NDList(y, ey, x, ex); + } + + /** + * bbox regression + * @param boundingbox + * @param reg + * @return + */ + public static NDArray bbreg(NDArray boundingbox, NDArray reg) { + + // 如果 reg 是形状 [N,1,H,W],重塑为 [H,W] 或 [N,H] 这里假设 NCHW + if (reg.getShape().get(1) == 1) { + reg = reg.reshape(reg.getShape().get(2), reg.getShape().get(3)); + } + + // 确保 float32 + boundingbox = boundingbox.toType(DataType.FLOAT32, false); + reg = reg.toType(DataType.FLOAT32, false); + + // 计算宽高 + NDArray w = boundingbox.get(":, 2").sub(boundingbox.get(":, 0")).add(1); + NDArray h = boundingbox.get(":, 3").sub(boundingbox.get(":, 1")).add(1); + + NDArray b1 = boundingbox.get(":, 0").add(reg.get(":, 0").mul(w)); + NDArray b2 = boundingbox.get(":, 1").add(reg.get(":, 1").mul(h)); + NDArray b3 = boundingbox.get(":, 2").add(reg.get(":, 2").mul(w)); + NDArray b4 = boundingbox.get(":, 3").add(reg.get(":, 3").mul(h)); + + // stack + transpose 对应 Python stack + permute + NDArray newBox = NDArrays.stack(new NDList(b1, b2, b3, b4), 0).transpose(); + + // 更新 boundingbox[:, :4] + boundingbox.set(new NDIndex(":, 0:4"), newBox); + return boundingbox; + } +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/ONetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/ONetModel.java new file mode 100644 index 0000000..fcf94fd --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/ONetModel.java @@ -0,0 +1,202 @@ +package cn.smartjavaai.face.model.facedect.mtcnn; + +import ai.djl.inference.Predictor; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.util.NDImageUtils; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.translate.TranslateException; +import cn.smartjavaai.common.utils.NMSUtils; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; + +/** + * @author dwj + */ +public class ONetModel { + + + public static MtcnnBatchResult thirdStage(NDManager manager, Predictor onetPredictor, NDArray imgs, NDArray boxes, int w, int h, NDArray scoresFiltered,NDArray image_indsFiltered) throws TranslateException { + // Third stage + NDArray points = manager.zeros(new Shape(0, 5, 2)); + NDList pad = MtcnnUtils.pad(boxes, (int)w, (int)h); + NDArray y = pad.get(0); + NDArray ey = pad.get(1); + NDArray x = pad.get(2); + NDArray ex = pad.get(3); + List crops = new ArrayList<>(); + long numFaces = y.size(0); + for (long k = 0; k < numFaces; k++) { + // 检查坐标合法性 + if (ey.getInt(k) > (y.getInt(k) - 1) && + ex.getInt(k) > (x.getInt(k) - 1)) { + + // 裁剪 (imageInd, :, y1:ey, x1:ex) + NDArray imgK = imgs.get( + image_indsFiltered.getLong(k) + ", :" + + ", " + (y.getInt(k) - 1) + ":" + ey.getInt(k) + + ", " + (x.getInt(k) - 1) + ":" + ex.getInt(k) + ).expandDims(0); // 加 batch 维 + + // 缩放到 (24, 24) + + // (N, H, W, C) + NDArray transposed = imgK.transpose(0, 2, 3, 1); + transposed = NDImageUtils.resize(transposed, 48, 48, Image.Interpolation.AREA); + // (N, C, H, W) + transposed = transposed.transpose(0, 3, 1, 2); + crops.add(transposed); + } + } + + // 合并成一个 batch + NDArray im_data = NDArrays.concat(new NDList(crops), 0); + // 归一化 + im_data = im_data.sub(127.5).mul(0.0078125); + // 74 48 48 + NDList out = onetPredictor.predict(new NDList(im_data)); + + NDArray out0 = out.get(0).transpose(1, 0); // permute(1,0) + NDArray out1 = out.get(1).transpose(1, 0); + NDArray out2 = out.get(2).transpose(1, 0); + NDArray score = out1.get(1); // out1[1, :] + points = out1.duplicate(); + NDArray ipass = score.gt(0.7); // score > threshold[1] + // ipass 为布尔/0-1张量,长度应等于 points 的第 1 维(这里是 7) + NDArray ipassBool = ipass.toType(DataType.BOOLEAN, false); + long[] colIdx = ipassBool.nonzero().toLongArray(); // 取 True 的列索引 + // 把第 1 维换到第 0 维:(10, 7) -> (7, 10) + NDArray moved = points.swapAxes(0, 1); + // 现在按第一维取行即可,相当于选中列 + NDArray selected = moved.get(points.getManager().create(colIdx)); // (sel, 10) + // 换回原来的轴顺序:(sel, 10) -> (10, sel) + points = selected.swapAxes(0, 1); + // 筛选 boxes 和 scores + // 先获取布尔索引为 true 的行索引 + long[] validIndices = ipass.nonzero().toLongArray(); + // 筛选 boxes 对应行 + NDArray boxesSelected = boxes.get(manager.create(validIndices)); // 行筛选 + // 取前 4 列 + boxesSelected = boxesSelected.get(":, 0:4"); // 只保留前 4 列 + scoresFiltered = scoresFiltered.get(ipass).reshape(-1, 1); // score[ipass].unsqueeze(1) + boxes = NDArrays.concat(new NDList(boxesSelected, scoresFiltered), 1); // 拼接成 (N,5) + + // 筛选 image_inds + image_indsFiltered = image_indsFiltered.get(ipass); + NDArray mv = out0.transpose() // (N, 4) + .get(ipass); // 1-D 花式索引在第 0 维,得到 (k, 4) + + System.out.println("----------"); + + // w_i = boxes[:, 2] - boxes[:, 0] + 1 + NDArray w_i = boxes.get(":,2").sub(boxes.get(":,0")).add(1); + + // h_i = boxes[:, 3] - boxes[:, 1] + 1 + NDArray h_i = boxes.get(":,3").sub(boxes.get(":,1")).add(1); + + // points_x = w_i.repeat(5, 1) * points[:5, :] + boxes[:, 0].repeat(5, 1) - 1 + NDArray w_repeat = w_i.expandDims(0).repeat(0, 5); // shape: [5, N] + NDArray p_x = points.get("0:5,:").mul(w_repeat) + .add(boxes.get(":,0").expandDims(0).repeat(0, 5)) + .sub(1); + + // points_y = h_i.repeat(5, 1) * points[5:10, :] + boxes[:, 1].repeat(5, 1) - 1 + NDArray h_repeat = h_i.expandDims(0).repeat(0, 5); // shape: [5, N] + NDArray p_y = points.get("5:10,:").mul(h_repeat) + .add(boxes.get(":,1").expandDims(0).repeat(0, 5)) + .sub(1); + + // points = torch.stack((points_x, points_y)).permute(2, 1, 0) + NDArray pointsStacked = NDArrays.stack(new NDList(p_x, p_y)); // shape: [2, 5, N] + points = pointsStacked.transpose(2, 1, 0); // permute(2, 1, 0) => shape [N, 5, 2] + + // boxes = bbreg(boxes, mv) + boxes = MtcnnUtils.bbreg(boxes, mv); + + NDArray pick = NMSUtils.batchedNms(boxes.get(":, :4"), boxes.get(":, 4"), image_indsFiltered, 0.7f, manager); + boxes = boxes.get(pick); + image_indsFiltered = image_indsFiltered.get(pick); + points = points.get(pick); + + List batchBoxes = new ArrayList<>(); + List batchPoints = new ArrayList<>(); + + for (int b_i = 0; b_i < 1; b_i++) { + // mask: image_inds == b_i + NDArray mask = image_indsFiltered.eq(b_i); + + // 只保留当前 batch 的 boxes 和 points + NDArray batchBox = boxes.get(mask); + NDArray batchPoint = points.get(mask); + + batchBoxes.add(batchBox); + batchPoints.add(batchPoint); + } + return processBatchBoxes(batchBoxes, batchPoints,true, manager); + } + + public static MtcnnBatchResult processBatchBoxes( + List batchBoxes, + List batchPoints, + boolean selectLargest, + NDManager manager) { + + List boxesOut = new ArrayList<>(); + List probsOut = new ArrayList<>(); + List pointsOut = new ArrayList<>(); + + for (int i = 0; i < batchBoxes.size(); i++) { + NDArray box = batchBoxes.get(i); // shape [num_boxes, ?] 或空 NDArray + NDArray point = batchPoints.get(i); // shape [num_boxes, 5, 2] 或空 NDArray + + if (box == null || box.isEmpty()) { + boxesOut.add(null); + probsOut.add(null); + pointsOut.add(null); + continue; + } + + NDArray boxesSelected; + NDArray probsSelected; + NDArray pointsSelected; + + if (selectLargest) { + // 计算面积 (x2 - x1) * (y2 - y1) + NDArray w = box.get(":,2").sub(box.get(":,0")); + NDArray h = box.get(":,3").sub(box.get(":,1")); + NDArray areas = w.mul(h); + + // 按面积降序排序 + NDArray order = areas.argSort().flip(0); + boxesSelected = box.get(order); + pointsSelected = point.get(order); + } else { + boxesSelected = box; + pointsSelected = point; + } + + // boxes[:, :4] + boxesSelected = boxesSelected.get(":,0:4"); + + // probs = box[:, 4] + probsSelected = box.get(":,4"); + + boxesOut.add(boxesSelected); + probsOut.add(probsSelected); + pointsOut.add(pointsSelected); + } + + MtcnnBatchResult result = new MtcnnBatchResult(boxesOut, probsOut, pointsOut); + result.boxes = boxesOut; + result.probs = probsOut; + result.points = pointsOut; + return result; + } + +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/PNetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/PNetModel.java new file mode 100644 index 0000000..456e9ad --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/PNetModel.java @@ -0,0 +1,186 @@ +package cn.smartjavaai.face.model.facedect.mtcnn; + +import ai.djl.inference.Predictor; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.util.NDImageUtils; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.translate.TranslateException; +import ai.djl.translate.TranslatorContext; +import cn.smartjavaai.common.utils.NMSUtils; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; + +/** + * @author dwj + */ +public class PNetModel { + + + + /** + * 输入图片 + * @param imgs + * @param w + * @param h + * @param scale + * @return + */ + public static NDList processInput(NDArray imgs, int w, int h, double scale){ + int newH = (int) (h * scale + 1); + int newW = (int) (w * scale + 1); + // (N, H, W, C) + NDArray transposed = imgs.transpose(0, 2, 3, 1); + transposed = NDImageUtils.resize(transposed, newW, newH, Image.Interpolation.AREA); + // (N, C, H, W) + transposed = transposed.transpose(0, 3, 1, 2); + // 归一化 + transposed = transposed.sub(127.5).mul(0.0078125f); + System.out.println("inputPnet: " + transposed.getShape()); + return new NDList(transposed); + } + + public static NDList processOutput(NDList outputPnet, double scale, NDManager manager){ + NDArray reg = outputPnet.get(0); // [B, 4, H, W] + NDArray probs = outputPnet.get(1); // [B, 2, H, W] + List boundingBox = generateBoundingBox(reg, probs.get(":, 1"), (float)scale, 0.6f); + NDArray boxes_scale = boundingBox.get(0); // [N,9] + NDArray imgIndND = boundingBox.get(1); // [N] + NDArray pick = NMSUtils.batchedNms(boxes_scale.get(":,:4"), boxes_scale.get(":,4"), imgIndND, 0.5f, manager); + return new NDList(boxes_scale, imgIndND, pick); + } + + + public static NDList firstStage(NDManager manager, Predictor pnetPredictor, NDArray imgs, List scales, int width, int height) throws TranslateException { + // 第一阶段 + NDList boxes_list = new NDList(); + NDList image_inds_list = new NDList(); + NDList scale_picks_list = new NDList(); + int offset = 0; + for (double scale : scales) { + NDList inputPnet = processInput(imgs, width, height, scale); + NDList outputPnet = pnetPredictor.predict(inputPnet); + NDList output = processOutput(outputPnet, scale, manager); + NDArray boxes_scale = output.get(0); + NDArray imgIndND = output.get(1); + NDArray pick = output.get(2); + boxes_list.add(boxes_scale); + image_inds_list.add(imgIndND); + scale_picks_list.add(pick.add(offset)); + offset += boxes_scale.getShape().get(0); + } + // 270 9 + NDArray boxes = NDArrays.concat(boxes_list, 0); + NDArray image_inds = NDArrays.concat(image_inds_list, 0); + NDArray scale_picks = NDArrays.concat(scale_picks_list, 0); + + // NMS within each scale + image + boxes = boxes.get(scale_picks); // scalePicksAll 是 NDArrays.concat 后的 INT64 NDArray + image_inds = image_inds.get(scale_picks); // 同样索引 + + // NMS within each image + NDArray pick = NMSUtils.batchedNms( + boxes.get(":, :4"), // 坐标 + boxes.get(":, 4"), // score + image_inds, // 每个框对应的图片编号 + 0.7f, // IoU 阈值 + manager + ); + // 8 9 + boxes = boxes.get(pick); + image_inds = image_inds.get(pick); + + System.out.println(Arrays.toString(boxes.get(0).toFloatArray())); + + NDArray regw = boxes.get(":, 2").sub(boxes.get(":, 0")); + NDArray regh = boxes.get(":, 3").sub(boxes.get(":, 1")); + + NDArray qq1 = boxes.get(":, 0").add(boxes.get(":, 5").mul(regw)); + NDArray qq2 = boxes.get(":, 1").add(boxes.get(":, 6").mul(regh)); + NDArray qq3 = boxes.get(":, 2").add(boxes.get(":, 7").mul(regw)); + NDArray qq4 = boxes.get(":, 3").add(boxes.get(":, 8").mul(regh)); + + boxes = NDArrays.stack(new NDList(qq1, qq2, qq3, qq4, boxes.get(":, 4")), 1); + boxes = MtcnnUtils.rerec(boxes); + return new NDList(boxes, image_inds); + } + + /** + * 生成候选框,等价于 Python 版 generateBoundingBox + * + * @param reg NDArray [B,4,H,W],回归偏移量 + * @param probs NDArray [B,H,W],人脸概率 + * @param scale 当前金字塔缩放比例 + * @param threshold 阈值 + * @return 一个包含两个元素的 List: + * 0 -> NDArray bounding boxes [N,9] (x1,y1,x2,y2,score,dx1,dy1,dx2,dy2) + * 1 -> NDArray image_inds [N] + */ + public static List generateBoundingBox( + NDArray reg, NDArray probs, float scale, float threshold) { + + float stride = 2f; + float cellSize = 12f; + + // mask = probs >= thresh -> [B,H,W] + NDArray mask = probs.gte(threshold); + + System.out.println("reg: " + reg.getShape()); + System.out.println("probs shape: " + probs.getShape()); + System.out.println("scale: " + scale); + System.out.println("mask: " + mask.getShape()); + // mask_inds = mask.nonzero() -> [N,3] 每行: (batch, y, x) + NDArray maskInds = mask.nonzero(); + System.out.println("maskInds: " + maskInds.getShape()); + + // image_inds = mask_inds[:, 0] + NDArray imageInds = maskInds.get(":,0"); + + // yx = mask_inds[:, 1:] [N,2] -> (y, x) + NDArray yx = maskInds.get(":,1:"); + + // bb = mask_inds[:, 1:].flip(1) Python 是 (y,x) -> (x,y) + NDArray bb = yx.flip(1); // [N,2] (x, y) + + // 左上角 (x1, y1) 坐标 q1 = ((stride * bb + 1) / scale).floor() + NDArray q1 = bb.mul(stride).add(1).div(scale).floor(); + + // 右下角 (x2, y2) 坐标 q2 = ((stride * bb + cellsize) / scale).floor() + NDArray q2 = bb.mul(stride).add(cellSize).div(scale).floor(); + + Shape probShape = probs.getShape(); // [B,H,W] + long H = probShape.get(1); + long W = probShape.get(2); + + NDArray linearIndex = maskInds.get(":,0").mul(H * W) + .add(maskInds.get(":,1").mul(W)) + .add(maskInds.get(":,2")); + + NDArray scores = probs.reshape(-1).gather(linearIndex, 0); + + NDArray regPerm = reg.transpose(1, 0, 2, 3); // [4,B,H,W] + NDArray regFlat = regPerm.reshape(4, -1); // [4, total] + + NDArray linearIndexForGather = linearIndex.expandDims(0).repeat(0, 4); // [4, N] + NDArray regPicked = regFlat.gather(linearIndexForGather, 1).transpose(); // [N,4] + + NDArray x1 = q1.get(":, 0").expandDims(1); + NDArray y1 = q1.get(":, 1").expandDims(1); + NDArray x2 = q2.get(":, 0").expandDims(1); + NDArray y2 = q2.get(":, 1").expandDims(1); + + NDArray boundingBoxes = NDArrays.concat( + new NDList(x1, y1, x2, y2, scores.expandDims(1), regPicked), 1 + ); + + return Arrays.asList(boundingBoxes, imageInds); + } + + +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/RNetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/RNetModel.java new file mode 100644 index 0000000..2f4aa2a --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/mtcnn/RNetModel.java @@ -0,0 +1,107 @@ +package cn.smartjavaai.face.model.facedect.mtcnn; + +import ai.djl.inference.Predictor; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.util.NDImageUtils; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.translate.TranslateException; +import cn.smartjavaai.common.utils.NMSUtils; +import lombok.extern.slf4j.Slf4j; + +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; + +/** + * @author dwj + */ +@Slf4j +public class RNetModel { + + public static NDList secondStage(NDManager manager, Predictor rnetPredictor,NDArray imgs, NDArray boxes, NDList pad, NDArray image_inds) throws TranslateException { + NDArray y = pad.get(0); + NDArray ey = pad.get(1); + NDArray x = pad.get(2); + NDArray ex = pad.get(3); + List crops = new ArrayList<>(); + long numFaces = y.size(0); + for (long k = 0; k < numFaces; k++) { + // 检查坐标合法性 + if (ey.getInt(k) > (y.getInt(k) - 1) && + ex.getInt(k) > (x.getInt(k) - 1)) { + + // 裁剪 (imageInd, :, y1:ey, x1:ex) + NDArray imgK = imgs.get( + image_inds.getLong(k) + ", :" + + ", " + (y.getInt(k) - 1) + ":" + ey.getInt(k) + + ", " + (x.getInt(k) - 1) + ":" + ex.getInt(k) + ).expandDims(0); // 加 batch 维 + + // 缩放到 (24, 24) + + // (N, H, W, C) + NDArray transposed = imgK.transpose(0, 2, 3, 1); + transposed = NDImageUtils.resize(transposed, 24, 24, Image.Interpolation.AREA); + // (N, C, H, W) + transposed = transposed.transpose(0, 3, 1, 2); + crops.add(transposed); + } + } + if (crops.isEmpty()) { + log.debug("No face detected."); + return null; + } + + // 合并成一个 batch + NDArray im_data = NDArrays.concat(new NDList(crops), 0); + // 归一化 + im_data = im_data.sub(127.5).mul(0.0078125); + NDList out = rnetPredictor.predict(new NDList(im_data)); + // 假设 out 是 NDList,threshold 是 float[],NMSUtils2.batchedNms 已经有了 + NDArray out0 = out.get(0).transpose(1, 0); // permute(1,0) + NDArray out1 = out.get(1).transpose(1, 0); + NDArray score = out1.get(1); // out1[1, :] + NDArray ipass = score.gt(0.7); // score > threshold[1] + + // 筛选 boxes 和 scores + // 先获取布尔索引为 true 的行索引 + long[] validIndices = ipass.nonzero().toLongArray(); + // 筛选 boxes 对应行 + NDArray boxesSelected = boxes.get(manager.create(validIndices)); // 行筛选 + // 取前 4 列 + boxesSelected = boxesSelected.get(":, 0:4"); // 只保留前 4 列 + NDArray scoresFiltered = score.get(ipass).reshape(-1, 1); // score[ipass].unsqueeze(1) + boxes = NDArrays.concat(new NDList(boxesSelected, scoresFiltered), 1); // 拼接成 (N,5) + + // 筛选 image_inds + NDArray image_indsFiltered = image_inds.get(ipass); + + // out0: (4, N) + NDArray mv = out0.transpose() // (N, 4) + .get(ipass); // 1-D 花式索引在第 0 维,得到 (k, 4) +// .transpose(); // 如需要保持 (k, 4) 可省略;如想与 Python 顺序一致可再转置 + + + // NMS + NDArray pick = NMSUtils.batchedNms(boxes.get(":, :4"), boxes.get(":, 4"), image_indsFiltered, 0.7f, manager); + + // 最终筛选 + boxes = boxes.get(pick); + image_indsFiltered = image_indsFiltered.get(pick); + mv = mv.get(pick); + + // 框回归和方形化 + boxes = MtcnnUtils.bbreg(boxes, mv); + boxes = MtcnnUtils.rerec(boxes); + + if(boxes.size(0) == 0){ + log.debug("No face detected."); + return null; + } + return new NDList(image_indsFiltered, scoresFiltered); + } + +} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java rename to face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java rename to face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java similarity index 98% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java rename to face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java index bcb884d..0a195fe 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java @@ -13,6 +13,7 @@ import cn.smartjavaai.face.constant.FaceDetectConstant; import cn.smartjavaai.face.entity.FaceRegisterInfo; import cn.smartjavaai.face.entity.FaceResult; import cn.smartjavaai.face.entity.FaceSearchParams; +import cn.smartjavaai.face.enums.FaceRecModelEnum; import cn.smartjavaai.face.enums.SimilarityType; import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.utils.FaceUtils; @@ -88,6 +89,10 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ log.debug("Loading seetaFace6 library successfully."); String[] faceDetectorModelPath = {config.getModelPath() + File.separator + "face_detector.csta"}; String[] faceRecognizerModelPath = {config.getModelPath() + File.separator + "face_recognizer.csta"}; + //轻量模型 + if(config.getModelEnum() == FaceRecModelEnum.SEETA_FACE6_LIGHT_MODEL){ + faceRecognizerModelPath = new String[] { config.getModelPath() + File.separator + "face_recognizer_light.csta" }; + } String[] faceLandmarkerModelPath = {config.getModelPath() + File.separator + "face_landmarker_pts5.csta"}; SeetaDevice device = SeetaDevice.SEETA_DEVICE_AUTO; int gpuId = config.getGpuId(); @@ -871,7 +876,16 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public void upsertFace(FaceRegisterInfo faceRegisterInfo, byte[] imageData) { - FaceRecModel.super.upsertFace(faceRegisterInfo, imageData); + if(Objects.isNull(imageData)){ + throw new FaceException("图像无效"); + } + BufferedImage bufferedImage = null; + try { + bufferedImage = ImageIO.read(new ByteArrayInputStream(imageData)); + } catch (IOException e) { + throw new FaceException(e); + } + upsertFace(faceRegisterInfo, bufferedImage); } @Override diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/criteria/FaceRecCriteriaFactory.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/criteria/FaceRecCriteriaFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/criteria/FaceRecCriteriaFactory.java rename to face/src/main/java/cn/smartjavaai/face/model/facerec/criteria/FaceRecCriteriaFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceFeatureTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceFeatureTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceFeatureTranslator.java rename to face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceFeatureTranslator.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceNetRecTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceNetRecTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceNetRecTranslator.java rename to face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceNetRecTranslator.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/criterial/LivenessCriteriaFactory.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/criterial/LivenessCriteriaFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/criterial/LivenessCriteriaFactory.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/criterial/LivenessCriteriaFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/translator/IicFrTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/translator/IicFrTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/translator/IicFrTranslator.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/translator/IicFrTranslator.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/translator/MiniVisionTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/translator/MiniVisionTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/liveness/translator/MiniVisionTranslator.java rename to face/src/main/java/cn/smartjavaai/face/model/liveness/translator/MiniVisionTranslator.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java b/face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java rename to face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java b/face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java rename to face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java b/face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java rename to face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/seetaface/ClarityDLResult.java b/face/src/main/java/cn/smartjavaai/face/seetaface/ClarityDLResult.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/seetaface/ClarityDLResult.java rename to face/src/main/java/cn/smartjavaai/face/seetaface/ClarityDLResult.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/seetaface/NativeLoader.java b/face/src/main/java/cn/smartjavaai/face/seetaface/NativeLoader.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/seetaface/NativeLoader.java rename to face/src/main/java/cn/smartjavaai/face/seetaface/NativeLoader.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/sqllite/ResultSetExtractor.java b/face/src/main/java/cn/smartjavaai/face/sqllite/ResultSetExtractor.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/sqllite/ResultSetExtractor.java rename to face/src/main/java/cn/smartjavaai/face/sqllite/ResultSetExtractor.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/sqllite/RowMapper.java b/face/src/main/java/cn/smartjavaai/face/sqllite/RowMapper.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/sqllite/RowMapper.java rename to face/src/main/java/cn/smartjavaai/face/sqllite/RowMapper.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/sqllite/SqliteHelper.java b/face/src/main/java/cn/smartjavaai/face/sqllite/SqliteHelper.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/sqllite/SqliteHelper.java rename to face/src/main/java/cn/smartjavaai/face/sqllite/SqliteHelper.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/translator/FaceDetectionTranslator.java b/face/src/main/java/cn/smartjavaai/face/translator/FaceDetectionTranslator.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/translator/FaceDetectionTranslator.java rename to face/src/main/java/cn/smartjavaai/face/translator/FaceDetectionTranslator.java diff --git a/face/src/main/java/cn/smartjavaai/face/translator/SCRFDFaceTranslator.java b/face/src/main/java/cn/smartjavaai/face/translator/SCRFDFaceTranslator.java new file mode 100644 index 0000000..ca0997e --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/translator/SCRFDFaceTranslator.java @@ -0,0 +1,242 @@ +package cn.smartjavaai.face.translator; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.output.Landmark; +import ai.djl.modality.cv.output.Point; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; +import ai.djl.translate.Translator; +import ai.djl.translate.TranslatorContext; +import cn.smartjavaai.common.utils.LetterBoxUtils; +import cn.smartjavaai.common.utils.NMSUtils; + +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +/** + * SCRFD Translator + */ +public class SCRFDFaceTranslator implements Translator { + + private double confThresh; + private double nmsThresh; + private int topK; + private int[] steps; + private int inputWidth = 640; + private int inputHeight = 640; + + public SCRFDFaceTranslator( + double confThresh, + double nmsThresh, + int topK, + int[] steps) { + this.confThresh = confThresh; + this.nmsThresh = nmsThresh; + this.topK = topK; + this.steps = steps; + } + + /** {@inheritDoc} */ + @Override + public NDList processInput(TranslatorContext ctx, Image input) { + + ctx.setAttachment("width", input.getWidth()); + ctx.setAttachment("height", input.getHeight()); + + NDManager manager = ctx.getNDManager(); + NDArray array = input.toNDArray(manager, Image.Flag.COLOR); + + //Letter box resize 640x640 with padding (保持比例,补边缘) + LetterBoxUtils.ResizeResult letterBoxResult = LetterBoxUtils.letterbox(manager, array, inputWidth, inputHeight, 0f, LetterBoxUtils.PaddingPosition.LEFT_TOP); + ctx.setAttachment("scale", letterBoxResult.r); + array = letterBoxResult.image; + array = array.transpose(2, 0, 1).flip(0); // HWC -> CHW RGB -> BGR + // The network by default takes float32 + if (!array.getDataType().equals(DataType.FLOAT32)) { + array = array.toType(DataType.FLOAT32, false); + } + // 归一化 + array = array.sub(127.5).mul(0.0078125); + return new NDList(array); + } + + /** {@inheritDoc} */ + @Override + public DetectedObjects processOutput(TranslatorContext ctx, NDList list) { + + NDManager manager = + NDManager.newBaseManager(ctx.getNDManager().getDevice(), "PyTorch"); + int sourceWidth = (int) ctx.getAttachment("width"); + int sourceHeight = (int) ctx.getAttachment("height"); + float detScale = (float) ctx.getAttachment("scale"); + + Map centerCache = new HashMap<>(); + List scores_list = new ArrayList<>(); + List bboxes_list = new ArrayList<>(); + List kpss_list = new ArrayList<>(); + // 步长长度 + int fmc = steps.length; + int numAnchors = 2; + // 多尺度后处理 + for (int idx = 0; idx < steps.length; idx++) { + int stride = steps[idx]; + NDArray scores, bboxPreds, kpsPreds = null; + scores = list.get(idx); + bboxPreds = list.get(idx + fmc).mul(stride); + kpsPreds = list.get(idx + fmc * 2).mul(stride); + int height = inputHeight / stride; + int width = inputWidth / stride; + String key = height + "_" + width + "_" + stride; + // anchor centers cache + NDArray anchorCenters; + if (centerCache.containsKey(key)) { + anchorCenters = centerCache.get(key); + } else { + NDArray yv = manager.arange((float) height).reshape(height, 1).repeat(1, width); + NDArray xv = manager.arange((float) width).reshape(1, width).repeat(0, height); + // stack x, y 到最后一维 + anchorCenters = xv.stack(yv, -1); // shape [height, width, 2] + // 乘 stride + anchorCenters = anchorCenters.mul(stride); + // 拉平成 [-1, 2],等价于 NumPy 的 reshape((-1, 2)) + long total = anchorCenters.getShape().get(0) * anchorCenters.getShape().get(1); + // anchorCenters 现在是 [height*width, 2] + anchorCenters = anchorCenters.reshape(total, 2); + // 在第一维重复 numAnchors 次,直接拉平成最终形状 + anchorCenters = anchorCenters.repeat(0, numAnchors); // shape [N*numAnchors, 2] + if (centerCache.size() < 100) { + centerCache.put(key, anchorCenters); + } + } + + NDArray pos_mask = scores.gte(confThresh); // scores >= thresh + NDArray pos_inds = pos_mask.nonzero(); // shape: [N, 2] + pos_inds = pos_inds.get(":, 0"); // 取第一列的索引 +// System.out.println(Arrays.toString(pos_inds.toLongArray())); + // 计算 bbox + NDArray bboxes = distance2bbox(anchorCenters, bboxPreds); // [num_anchors, 4] + + // 取出符合阈值的 + NDArray pos_scores = scores.get(pos_inds); + NDArray pos_bboxes = bboxes.get(pos_inds); + + scores_list.add(pos_scores); + bboxes_list.add(pos_bboxes); + + NDArray kpss = distance2kps(anchorCenters, kpsPreds); // [num_anchors, num_kps*2] + kpss = kpss.reshape(kpss.getShape().get(0), -1, 2); // reshape (N, -1, 2) + NDArray pos_kpss = kpss.get(pos_inds); + kpss_list.add(pos_kpss); + } + + // 1. 合并 scores + NDArray scores = NDArrays.concat(new NDList(scores_list), 0); + NDArray scoresRavel = scores.reshape(-1); + + // 2. 得到排序索引 + long[] orderLong = scoresRavel.argSort().flip(0).get(":" + topK).toLongArray(); + + // 3. 合并 bboxes + NDArray bboxes = NDArrays.concat(new NDList(bboxes_list), 0).div(detScale); + + NDArray kpss = NDArrays.concat(new NDList(kpss_list), 0).div(detScale); + // 4. 拼接 [x1,y1,x2,y2,score] + NDArray preDet = bboxes.concat(scores.reshape(-1,1), 1); + // 5. 按 order 排序 + preDet = preDet.get(manager.create(orderLong)); + // 6. NMS + int[] keep = NMSUtils.nms(preDet.get(":,0:4"), preDet.get(":,4"), (float)nmsThresh); + NDArray det = preDet.get(manager.create(keep)); +// System.out.println(Arrays.toString(det.toFloatArray())); + if (kpss != null) { + kpss = kpss.get(manager.create(orderLong)); + kpss = kpss.get(manager.create(keep)); + } + List retNames = new ArrayList<>(); + List retProbs = new ArrayList<>(); + List retBB = new ArrayList<>(); + long numDet = det.getShape().get(0); // N + long numCols = det.getShape().get(1); // 应该是 5: x1,y1,x2,y2,score + float[] flat = det.toFloatArray(); // 一维 + for (int i = 0; i < numDet; i++) { + int base = (int) (i * numCols); + float x1 = flat[base] / sourceWidth; + float y1 = flat[base + 1] / sourceHeight; + float x2 = flat[base + 2] / sourceWidth; + float y2 = flat[base + 3] / sourceHeight; + float score = flat[base + 4]; + + retNames.add("face"); // 类别 + retProbs.add((double) score); + + float width = x2 - x1; + float height = y2 - y1; + Landmark rect = new Landmark(x1, y1, width, height, decodeKps(kpss.get(i))); + retBB.add(rect); + } + return new DetectedObjects(retNames, retProbs, retBB); + } + + + public NDArray distance2bbox(NDArray points, NDArray distance) { + // points: [N, 2], distance: [N, 4] + NDArray x1 = points.get(":, 0").sub(distance.get(":, 0")); // x - left + NDArray y1 = points.get(":, 1").sub(distance.get(":, 1")); // y - top + NDArray x2 = points.get(":, 0").add(distance.get(":, 2")); // x + right + NDArray y2 = points.get(":, 1").add(distance.get(":, 3")); // y + bottom + // stack([x1, y1, x2, y2], axis=-1) + NDList list = new NDList(x1.expandDims(1), y1.expandDims(1), x2.expandDims(1), y2.expandDims(1)); + return NDArrays.concat(list, 1); // axis=1 表示最后一维 + } + + + public NDArray distance2kps(NDArray points, NDArray distance) { + // points: [N, 2], distance: [N, 2*num_kps] + int numKps = (int) distance.getShape().get(1) / 2; + List preds = new ArrayList<>(); + for (int i = 0; i < numKps * 2; i += 2) { + NDArray px = points.get(":, " + (i % 2)).add(distance.get(":, " + i)); + NDArray py = points.get(":, " + ((i % 2) + 1)).add(distance.get(":, " + (i + 1))); + preds.add(px); + preds.add(py); + } + // stack(preds, axis=-1) + NDList stackList = new NDList(); + for (NDArray arr : preds) { + stackList.add(arr.expandDims(1)); + } + return NDArrays.concat(stackList, 1); // shape [N, num_kps*2] + } + + + public List decodeKps(NDArray kpss) { + // 转成一维 float 数组 + float[] flat = kpss.toFloatArray(); + + // reshape 成二维 [5][2] + int numPoints = (int) kpss.getShape().get(0); // 5 + int dim = (int) kpss.getShape().get(1); // 2 + float[][] kpsArray = new float[numPoints][dim]; + for (int i = 0; i < numPoints; i++) { + for (int j = 0; j < dim; j++) { + kpsArray[i][j] = flat[i * dim + j]; + } + } + + // 转成 Point 数组 + List points = new ArrayList<>(); + for (int i = 0; i < numPoints; i++) { + points.add(new Point(Math.round(kpsArray[i][0]), Math.round(kpsArray[i][1]))); + } + return points; + } + + +} diff --git a/face/src/main/java/cn/smartjavaai/face/translator/YoloV5FaceTranslator.java b/face/src/main/java/cn/smartjavaai/face/translator/YoloV5FaceTranslator.java new file mode 100644 index 0000000..8733de6 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/translator/YoloV5FaceTranslator.java @@ -0,0 +1,462 @@ +package cn.smartjavaai.face.translator; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.*; +import ai.djl.modality.cv.transform.*; +import ai.djl.modality.cv.util.NDImageUtils; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.translate.*; + +import java.util.*; + + +/** + * YoloV5Translator + * @author dwj + */ +public class YoloV5FaceTranslator implements Translator { + + private int maxBoxes; + + private YoloOutputType yoloOutputLayerType; + private float nmsThreshold; + + protected float threshold; +// private BaseImageTranslator.SynsetLoader synsetLoader; + protected List classes; + protected boolean applyRatio; + protected boolean removePadding; + + protected Pipeline pipeline; + private Image.Flag flag; + private Batchifier batchifier; + protected int width; + protected int height; + + + /** + * Constructs an ImageTranslator with the provided builder. + * + * @param builder the data to build with + */ + protected YoloV5FaceTranslator(Builder builder) { + this.yoloOutputLayerType = builder.outputType; + this.nmsThreshold = builder.nmsThreshold; + maxBoxes = builder.maxBox; + this.threshold = builder.threshold; +// this.synsetLoader = builder.synsetLoader; + this.applyRatio = builder.applyRatio; + this.removePadding = builder.removePadding; + this.flag = builder.flag; + this.pipeline = builder.pipeline; + this.batchifier = builder.batchifier; + this.width = builder.width; + this.height = builder.height; + classes = Arrays.asList("face"); + } + + /** + * Creates a builder to build a {@code YoloV8Translator} with specified arguments. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Creates a builder to build a {@code YoloV8Translator} with specified arguments. + * + * @param arguments arguments to specify builder options + * @return a new builder + */ + public static Builder builder(Map arguments) { + Builder builder = new Builder(); + builder.configPreProcess(arguments); + builder.configPostProcess(arguments); + return builder; + } + + /** {@inheritDoc} */ + protected DetectedObjects processFromBoxOutput(int imageWidth, int imageHeight, NDList list) { + float[] flattened = list.get(0).toFloatArray(); + int sizeClasses = classes.size(); + int stride = 15 + sizeClasses; + int size = flattened.length / stride; + + ArrayList boxes = new ArrayList<>(); + ArrayList scores = new ArrayList<>(); + ArrayList classIds = new ArrayList<>(); + + for (int i = 0; i < size; i++) { + int indexBase = i * stride; + float maxClass = 0; + int maxIndex = 0; +// for (int c = 0; c < sizeClasses; c++) { +// if (flattened[indexBase + c + 5] > maxClass) { +// maxClass = flattened[indexBase + c + 5]; +// maxIndex = c; +// } +// } + float score = flattened[indexBase + 4]; + if (score > threshold) { + float xPos = flattened[indexBase]; + float yPos = flattened[indexBase + 1]; + float w = flattened[indexBase + 2]; + float h = flattened[indexBase + 3]; + List keypoints = new ArrayList<>(); + keypoints.add(new Point(flattened[indexBase + 5], flattened[indexBase + 6])); + keypoints.add(new Point(flattened[indexBase + 7], flattened[indexBase + 8])); + keypoints.add(new Point(flattened[indexBase + 9], flattened[indexBase + 10])); + keypoints.add(new Point(flattened[indexBase + 11], flattened[indexBase + 12])); + keypoints.add(new Point(flattened[indexBase + 13], flattened[indexBase + 14])); + Landmark rect = + new Landmark(Math.max(0, xPos - w / 2), Math.max(0, yPos - h / 2), w, h,keypoints); + boxes.add(rect); + scores.add(score); + classIds.add(maxIndex); + } + } + return nms(imageWidth, imageHeight, boxes, classIds, scores); + } + + private DetectedObjects processFromDetectOutput() { + throw new UnsupportedOperationException( + "detect layer output is not supported yet, check correct YoloV5 export format"); + } + + @Override + public DetectedObjects processOutput(TranslatorContext ctx, NDList list) throws Exception { + int imageWidth = (Integer) ctx.getAttachment("width"); + int imageHeight = (Integer) ctx.getAttachment("height"); + switch (yoloOutputLayerType) { + case DETECT: + return processFromDetectOutput(); + case AUTO: + if (list.get(0).getShape().dimension() > 2) { + return processFromDetectOutput(); + } else { + return processFromBoxOutput(imageWidth, imageHeight, list); + } + case BOX: + default: + return processFromBoxOutput(imageWidth, imageHeight, list); + } + } + + @Override + public NDList processInput(TranslatorContext ctx, Image input) throws Exception { + NDArray array = input.toNDArray(ctx.getNDManager(), flag); + NDList list = pipeline.transform(new NDList(array)); + Shape shape = list.get(0).getShape(); + int processedWidth; + int processedHeight; + long[] dim = shape.getShape(); + if (NDImageUtils.isCHW(shape)) { + processedWidth = (int) dim[dim.length - 1]; + processedHeight = (int) dim[dim.length - 2]; + } else { + processedWidth = (int) dim[dim.length - 2]; + processedHeight = (int) dim[dim.length - 3]; + } + ctx.setAttachment("width", input.getWidth()); + ctx.setAttachment("height", input.getHeight()); + ctx.setAttachment("processedWidth", processedWidth); + ctx.setAttachment("processedHeight", processedHeight); + return list; + } + + protected DetectedObjects nms( + int imageWidth, + int imageHeight, + List boxes, + List classIds, + List scores) { + List retClasses = new ArrayList<>(); + List retProbs = new ArrayList<>(); + List retBB = new ArrayList<>(); + + for (int classId = 0; classId < classes.size(); classId++) { + List r = new ArrayList<>(); + List s = new ArrayList<>(); + List map = new ArrayList<>(); + for (int j = 0; j < classIds.size(); ++j) { + if (classIds.get(j) == classId) { + r.add(boxes.get(j)); + s.add(scores.get(j).doubleValue()); + map.add(j); + } + } + if (r.isEmpty()) { + continue; + } + List nms = Rectangle.nms(r, s, nmsThreshold); + for (int index : nms) { + int pos = map.get(index); + int id = classIds.get(pos); + retClasses.add(classes.get(id)); + retProbs.add(scores.get(pos).doubleValue()); +// Rectangle rect = boxes.get(pos); + Landmark rect = boxes.get(pos); + List keypoints = new ArrayList<>(); + if (removePadding) { + int padW = (width - imageWidth) / 2; + int padH = (height - imageHeight) / 2; + rect.getPath().forEach(point -> { + keypoints.add(new Point(point.getX() - padW, point.getY() - padH)); + }); + rect = + new Landmark( + (rect.getX() - padW) / imageWidth, + (rect.getY() - padH) / imageHeight, + rect.getWidth() / imageWidth, + rect.getHeight() / imageHeight,keypoints); + } else if (applyRatio) { + rect.getPath().forEach(point -> { + keypoints.add(new Point(point.getX() / width, point.getY() / height)); + }); + rect = + new Landmark( + rect.getX() / width, + rect.getY() / height, + rect.getWidth() / width, + rect.getHeight() / height,keypoints); + } + retBB.add(rect); + } + } + return new DetectedObjects(retClasses, retProbs, retBB); + } + + public static class Builder { + + private int maxBox = 8400; + + YoloOutputType outputType; + float nmsThreshold; + + protected float threshold = 0.2F; + protected boolean applyRatio; + protected boolean removePadding; + + protected int width = 224; + protected int height = 224; + protected Image.Flag flag; + protected Pipeline pipeline; + protected Batchifier batchifier; + + public Builder() { + this.outputType = YoloOutputType.AUTO; + this.nmsThreshold = 0.4F; + } + + public Builder optOutputType(YoloOutputType outputType) { + this.outputType = outputType; + return this; + } + + public Builder optNmsThreshold(float nmsThreshold) { + this.nmsThreshold = nmsThreshold; + return this; + } + + /** + * Builds the translator. + * + * @return the new translator + */ + public YoloV5FaceTranslator build() { + if (pipeline == null) { + addTransform( + array -> array.transpose(2, 0, 1).toType(DataType.FLOAT32, false).div(255)); + } +// validate(); + return new YoloV5FaceTranslator(this); + } + + protected Builder self() { + return this; + } + + public Builder addTransform(Transform transform) { + if (this.pipeline == null) { + this.pipeline = new Pipeline(); + } + + this.pipeline.add(transform); + return this.self(); + } + + public Builder optApplyRatio(boolean value) { + this.applyRatio = value; + return this.self(); + } + + public Builder optFlag(Image.Flag flag) { + this.flag = flag; + return this.self(); + } + + public Builder setPipeline(Pipeline pipeline) { + this.pipeline = pipeline; + return this.self(); + } + + public Builder setImageSize(int width, int height) { + this.width = width; + this.height = height; + return this.self(); + } + + + public Builder optBatchifier(Batchifier batchifier) { + this.batchifier = batchifier; + return this.self(); + } + + public Builder optThreshold(float threshold) { + this.threshold = threshold; + return this.self(); + } + + /** {@inheritDoc} */ + protected void configPostProcess(Map arguments) { + if (ArgumentsUtil.booleanValue(arguments, "optApplyRatio") || ArgumentsUtil.booleanValue(arguments, "applyRatio")) { + this.optApplyRatio(true); + } + this.threshold = ArgumentsUtil.floatValue(arguments, "threshold", 0.2F); + String centerFit = ArgumentsUtil.stringValue(arguments, "centerFit", "false"); + this.removePadding = "true".equals(centerFit); + String type = ArgumentsUtil.stringValue(arguments, "outputType", "AUTO"); + this.outputType = YoloOutputType.valueOf(type.toUpperCase(Locale.ENGLISH)); + this.nmsThreshold = ArgumentsUtil.floatValue(arguments, "nmsThreshold", 0.4F); + maxBox = ArgumentsUtil.intValue(arguments, "maxBox", 8400); + } + + protected void configPreProcess(Map arguments) { + if (this.pipeline == null) { + this.pipeline = new Pipeline(); + } + + this.width = ArgumentsUtil.intValue(arguments, "width", 224); + this.height = ArgumentsUtil.intValue(arguments, "height", 224); + if (arguments.containsKey("flag")) { + this.flag = Image.Flag.valueOf(arguments.get("flag").toString()); + } + + String pad = ArgumentsUtil.stringValue(arguments, "pad", "false"); + if ("true".equals(pad)) { + this.addTransform(new Pad(0.0)); + } else if (!"false".equals(pad)) { + double padding = Double.parseDouble(pad); + this.addTransform(new Pad(padding)); + } + + String resize = ArgumentsUtil.stringValue(arguments, "resize", "false"); + int w; + int shortEdge; + if ("true".equals(resize)) { + this.addTransform(new Resize(this.width, this.height)); + } else if (!"false".equals(resize)) { + String[] tokens = resize.split("\\s*,\\s*"); + w = (int)Double.parseDouble(tokens[0]); + if (tokens.length > 1) { + shortEdge = (int)Double.parseDouble(tokens[1]); + } else { + shortEdge = w; + } + + Image.Interpolation interpolation; + if (tokens.length > 2) { + interpolation = Image.Interpolation.valueOf(tokens[2]); + } else { + interpolation = Image.Interpolation.BILINEAR; + } + + this.addTransform(new Resize(w, shortEdge, interpolation)); + } + + String resizeShort = ArgumentsUtil.stringValue(arguments, "resizeShort", "false"); + if ("true".equals(resizeShort)) { + w = Math.max(this.width, this.height); + this.addTransform(new ResizeShort(w)); + } else if (!"false".equals(resizeShort)) { + String[] tokens = resizeShort.split("\\s*,\\s*"); + shortEdge = (int)Double.parseDouble(tokens[0]); + int longEdge; + if (tokens.length > 1) { + longEdge = (int)Double.parseDouble(tokens[1]); + } else { + longEdge = -1; + } + + Image.Interpolation interpolation; + if (tokens.length > 2) { + interpolation = Image.Interpolation.valueOf(tokens[2]); + } else { + interpolation = Image.Interpolation.BILINEAR; + } + + this.addTransform(new ResizeShort(shortEdge, longEdge, interpolation)); + } + + if (ArgumentsUtil.booleanValue(arguments, "centerCrop", false)) { + this.addTransform(new CenterCrop(this.width, this.height)); + } + + if (ArgumentsUtil.booleanValue(arguments, "centerFit")) { + this.addTransform(new CenterFit(this.width, this.height)); + } + + if (ArgumentsUtil.booleanValue(arguments, "toTensor", true)) { + this.addTransform(new ToTensor()); + } + + String normalize = ArgumentsUtil.stringValue(arguments, "normalize", "false"); + if ("true".equals(normalize)) { + float[] MEAN = new float[]{0.485F, 0.456F, 0.406F}; + float[] STD = new float[]{0.229F, 0.224F, 0.225F}; + this.addTransform(new Normalize(MEAN, STD)); + } else if (!"false".equals(normalize)) { + String[] tokens = normalize.split("\\s*,\\s*"); + if (tokens.length != 6) { + throw new IllegalArgumentException("Invalid normalize value: " + normalize); + } + + float[] mean = new float[]{Float.parseFloat(tokens[0]), Float.parseFloat(tokens[1]), Float.parseFloat(tokens[2])}; + float[] std = new float[]{Float.parseFloat(tokens[3]), Float.parseFloat(tokens[4]), Float.parseFloat(tokens[5])}; + this.addTransform(new Normalize(mean, std)); + } + + String range = (String)arguments.get("range"); + if ("0,1".equals(range)) { + this.addTransform((a) -> { + return a.div(255.0F); + }); + } else if ("-1,1".equals(range)) { + this.addTransform((a) -> { + return a.div(128.0F).sub(1); + }); + } + + if (arguments.containsKey("batchifier")) { + this.batchifier = Batchifier.fromString((String)arguments.get("batchifier")); + } + + } + } + + public static enum YoloOutputType { + BOX, + DETECT, + AUTO; + + private YoloOutputType() { + } + } + + + +} diff --git a/face/src/main/java/cn/smartjavaai/face/translator/YoloV8FaceTranslator.java b/face/src/main/java/cn/smartjavaai/face/translator/YoloV8FaceTranslator.java new file mode 100644 index 0000000..95a3194 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/translator/YoloV8FaceTranslator.java @@ -0,0 +1,476 @@ +package cn.smartjavaai.face.translator; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.*; +import ai.djl.modality.cv.transform.*; +import ai.djl.modality.cv.util.NDImageUtils; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.translate.*; + +import java.util.*; + +/** + * YoloV8Translator + * @author dwj + */ +public class YoloV8FaceTranslator implements Translator { + + private int maxBoxes; + + private YoloOutputType yoloOutputLayerType; + private float nmsThreshold; + + protected float threshold; +// private BaseImageTranslator.SynsetLoader synsetLoader; + protected List classes; + protected boolean applyRatio; + protected boolean removePadding; + + protected Pipeline pipeline; + private Image.Flag flag; + private Batchifier batchifier; + protected int width; + protected int height; + + + /** + * Constructs an ImageTranslator with the provided builder. + * + * @param builder the data to build with + */ + protected YoloV8FaceTranslator(Builder builder) { + this.yoloOutputLayerType = builder.outputType; + this.nmsThreshold = builder.nmsThreshold; + maxBoxes = builder.maxBox; + this.threshold = builder.threshold; +// this.synsetLoader = builder.synsetLoader; + this.applyRatio = builder.applyRatio; + this.removePadding = builder.removePadding; + this.flag = builder.flag; + this.pipeline = builder.pipeline; + this.batchifier = builder.batchifier; + this.width = builder.width; + this.height = builder.height; + classes = Arrays.asList("face"); + } + + /** + * Creates a builder to build a {@code YoloV8Translator} with specified arguments. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Creates a builder to build a {@code YoloV8Translator} with specified arguments. + * + * @param arguments arguments to specify builder options + * @return a new builder + */ + public static Builder builder(Map arguments) { + Builder builder = new Builder(); + builder.configPreProcess(arguments); + builder.configPostProcess(arguments); + return builder; + } + + /** {@inheritDoc} */ + protected DetectedObjects processFromBoxOutput(int imageWidth, int imageHeight, NDList list) { + NDArray rawResult = list.get(0); + NDArray reshapedResult = rawResult.transpose(); + Shape shape = reshapedResult.getShape(); + float[] buf = reshapedResult.toFloatArray(); + int numberRows = Math.toIntExact(shape.get(0)); + int nClasses = Math.toIntExact(shape.get(1)); + int padding = nClasses - classes.size(); + System.out.println(Arrays.toString(reshapedResult.get(0).toFloatArray())); +// if (padding != 0 && padding != 4) { +// throw new IllegalStateException( +// "Expected classes: " + (nClasses - 4) + ", got " + classes.size()); +// } + + ArrayList boxes = new ArrayList<>(); + ArrayList scores = new ArrayList<>(); + ArrayList classIds = new ArrayList<>(); + + // reverse order search in heap; searches through #maxBoxes for optimization when set + for (int i = numberRows - 1; i > numberRows - maxBoxes; --i) { + int index = i * nClasses; + float maxClassProb = buf[index + 4]; +// int maxIndex = -1; +// for (int c = 4; c < nClasses; c++) { +// float classProb = buf[index + c]; +// if (classProb > maxClassProb) { +// maxClassProb = classProb; +// maxIndex = c; +// } +// } +// maxIndex -= padding; + + if (maxClassProb > threshold) { + float xPos = buf[index]; // center x + float yPos = buf[index + 1]; // center y + float w = buf[index + 2]; + float h = buf[index + 3]; + Rectangle rect = + new Rectangle(Math.max(0, xPos - w / 2), Math.max(0, yPos - h / 2), w, h); +// boxes.add(rect); + scores.add(maxClassProb); + classIds.add(0); + List keypoints = new ArrayList<>(); + keypoints.add(new Point(buf[index + 5], buf[index + 6])); + keypoints.add(new Point(buf[index + 8], buf[index + 9])); + keypoints.add(new Point(buf[index + 11], buf[index + 12])); + keypoints.add(new Point(buf[index + 14], buf[index + 15])); + keypoints.add(new Point(buf[index + 17], buf[index + 18])); + Landmark kps = new Landmark(Math.max(0, xPos - w / 2), Math.max(0, yPos - h / 2), w, h, keypoints); + boxes.add(kps); + } + } + + return nms(imageWidth, imageHeight, boxes, classIds, scores); + } + + private DetectedObjects processFromDetectOutput() { + throw new UnsupportedOperationException( + "detect layer output is not supported yet, check correct YoloV5 export format"); + } + + @Override + public DetectedObjects processOutput(TranslatorContext ctx, NDList list) throws Exception { + int imageWidth = (Integer) ctx.getAttachment("width"); + int imageHeight = (Integer) ctx.getAttachment("height"); + switch (yoloOutputLayerType) { + case DETECT: + return processFromDetectOutput(); + case AUTO: + if (list.get(0).getShape().dimension() > 2) { + return processFromDetectOutput(); + } else { + return processFromBoxOutput(imageWidth, imageHeight, list); + } + case BOX: + default: + return processFromBoxOutput(imageWidth, imageHeight, list); + } + } + + @Override + public NDList processInput(TranslatorContext ctx, Image input) throws Exception { + NDArray array = input.toNDArray(ctx.getNDManager(), flag); + NDList list = pipeline.transform(new NDList(array)); + Shape shape = list.get(0).getShape(); + int processedWidth; + int processedHeight; + long[] dim = shape.getShape(); + if (NDImageUtils.isCHW(shape)) { + processedWidth = (int) dim[dim.length - 1]; + processedHeight = (int) dim[dim.length - 2]; + } else { + processedWidth = (int) dim[dim.length - 2]; + processedHeight = (int) dim[dim.length - 3]; + } + ctx.setAttachment("width", input.getWidth()); + ctx.setAttachment("height", input.getHeight()); + ctx.setAttachment("processedWidth", processedWidth); + ctx.setAttachment("processedHeight", processedHeight); + return list; + } + + protected DetectedObjects nms( + int imageWidth, + int imageHeight, + List boxes, + List classIds, + List scores) { + List retClasses = new ArrayList<>(); + List retProbs = new ArrayList<>(); + List retBB = new ArrayList<>(); + + for (int classId = 0; classId < classes.size(); classId++) { + List r = new ArrayList<>(); + List s = new ArrayList<>(); + List map = new ArrayList<>(); + for (int j = 0; j < classIds.size(); ++j) { + if (classIds.get(j) == classId) { + r.add(boxes.get(j)); + s.add(scores.get(j).doubleValue()); + map.add(j); + } + } + if (r.isEmpty()) { + continue; + } + List nms = Rectangle.nms(r, s, nmsThreshold); + for (int index : nms) { + int pos = map.get(index); + int id = classIds.get(pos); + retClasses.add(classes.get(id)); + retProbs.add(scores.get(pos).doubleValue()); +// Rectangle rect = boxes.get(pos); + Landmark rect = boxes.get(pos); + List keypoints = new ArrayList<>(); + if (removePadding) { + int padW = (width - imageWidth) / 2; + int padH = (height - imageHeight) / 2; + rect.getPath().forEach(point -> { + keypoints.add(new Point(point.getX() - padW, point.getY() - padH)); + }); + rect = + new Landmark( + (rect.getX() - padW) / imageWidth, + (rect.getY() - padH) / imageHeight, + rect.getWidth() / imageWidth, + rect.getHeight() / imageHeight,keypoints); + } else if (applyRatio) { + rect.getPath().forEach(point -> { + keypoints.add(new Point(point.getX() / width, point.getY() / height)); + }); + rect = + new Landmark( + rect.getX() / width, + rect.getY() / height, + rect.getWidth() / width, + rect.getHeight() / height,keypoints); + } + retBB.add(rect); + } + } + return new DetectedObjects(retClasses, retProbs, retBB); + } + + public static class Builder { + + private int maxBox = 8400; + + YoloOutputType outputType; + float nmsThreshold; + + protected float threshold = 0.2F; + protected boolean applyRatio; + protected boolean removePadding; + + protected int width = 224; + protected int height = 224; + protected Image.Flag flag; + protected Pipeline pipeline; + protected Batchifier batchifier; + + public Builder() { + this.outputType = YoloOutputType.AUTO; + this.nmsThreshold = 0.4F; + } + + public Builder optOutputType(YoloOutputType outputType) { + this.outputType = outputType; + return this; + } + + public Builder optNmsThreshold(float nmsThreshold) { + this.nmsThreshold = nmsThreshold; + return this; + } + + /** + * Builds the translator. + * + * @return the new translator + */ + public YoloV8FaceTranslator build() { + if (pipeline == null) { + addTransform( + array -> array.transpose(2, 0, 1).toType(DataType.FLOAT32, false).div(255)); + } +// validate(); + return new YoloV8FaceTranslator(this); + } + + protected Builder self() { + return this; + } + + public Builder addTransform(Transform transform) { + if (this.pipeline == null) { + this.pipeline = new Pipeline(); + } + + this.pipeline.add(transform); + return this.self(); + } + + public Builder optApplyRatio(boolean value) { + this.applyRatio = value; + return this.self(); + } + + public Builder optFlag(Image.Flag flag) { + this.flag = flag; + return this.self(); + } + + public Builder setPipeline(Pipeline pipeline) { + this.pipeline = pipeline; + return this.self(); + } + + public Builder setImageSize(int width, int height) { + this.width = width; + this.height = height; + return this.self(); + } + + + public Builder optBatchifier(Batchifier batchifier) { + this.batchifier = batchifier; + return this.self(); + } + + public Builder optThreshold(float threshold) { + this.threshold = threshold; + return this.self(); + } + + /** {@inheritDoc} */ + protected void configPostProcess(Map arguments) { + if (ArgumentsUtil.booleanValue(arguments, "optApplyRatio") || ArgumentsUtil.booleanValue(arguments, "applyRatio")) { + this.optApplyRatio(true); + } + this.threshold = ArgumentsUtil.floatValue(arguments, "threshold", 0.2F); + String centerFit = ArgumentsUtil.stringValue(arguments, "centerFit", "false"); + this.removePadding = "true".equals(centerFit); + String type = ArgumentsUtil.stringValue(arguments, "outputType", "AUTO"); + this.outputType = YoloOutputType.valueOf(type.toUpperCase(Locale.ENGLISH)); + this.nmsThreshold = ArgumentsUtil.floatValue(arguments, "nmsThreshold", 0.4F); + maxBox = ArgumentsUtil.intValue(arguments, "maxBox", 8400); + } + + protected void configPreProcess(Map arguments) { + if (this.pipeline == null) { + this.pipeline = new Pipeline(); + } + + this.width = ArgumentsUtil.intValue(arguments, "width", 224); + this.height = ArgumentsUtil.intValue(arguments, "height", 224); + if (arguments.containsKey("flag")) { + this.flag = Image.Flag.valueOf(arguments.get("flag").toString()); + } + + String pad = ArgumentsUtil.stringValue(arguments, "pad", "false"); + if ("true".equals(pad)) { + this.addTransform(new Pad(0.0)); + } else if (!"false".equals(pad)) { + double padding = Double.parseDouble(pad); + this.addTransform(new Pad(padding)); + } + + String resize = ArgumentsUtil.stringValue(arguments, "resize", "false"); + int w; + int shortEdge; + if ("true".equals(resize)) { + this.addTransform(new Resize(this.width, this.height)); + } else if (!"false".equals(resize)) { + String[] tokens = resize.split("\\s*,\\s*"); + w = (int)Double.parseDouble(tokens[0]); + if (tokens.length > 1) { + shortEdge = (int)Double.parseDouble(tokens[1]); + } else { + shortEdge = w; + } + + Image.Interpolation interpolation; + if (tokens.length > 2) { + interpolation = Image.Interpolation.valueOf(tokens[2]); + } else { + interpolation = Image.Interpolation.BILINEAR; + } + + this.addTransform(new Resize(w, shortEdge, interpolation)); + } + + String resizeShort = ArgumentsUtil.stringValue(arguments, "resizeShort", "false"); + if ("true".equals(resizeShort)) { + w = Math.max(this.width, this.height); + this.addTransform(new ResizeShort(w)); + } else if (!"false".equals(resizeShort)) { + String[] tokens = resizeShort.split("\\s*,\\s*"); + shortEdge = (int)Double.parseDouble(tokens[0]); + int longEdge; + if (tokens.length > 1) { + longEdge = (int)Double.parseDouble(tokens[1]); + } else { + longEdge = -1; + } + + Image.Interpolation interpolation; + if (tokens.length > 2) { + interpolation = Image.Interpolation.valueOf(tokens[2]); + } else { + interpolation = Image.Interpolation.BILINEAR; + } + + this.addTransform(new ResizeShort(shortEdge, longEdge, interpolation)); + } + + if (ArgumentsUtil.booleanValue(arguments, "centerCrop", false)) { + this.addTransform(new CenterCrop(this.width, this.height)); + } + + if (ArgumentsUtil.booleanValue(arguments, "centerFit")) { + this.addTransform(new CenterFit(this.width, this.height)); + } + + if (ArgumentsUtil.booleanValue(arguments, "toTensor", true)) { + this.addTransform(new ToTensor()); + } + + String normalize = ArgumentsUtil.stringValue(arguments, "normalize", "false"); + if ("true".equals(normalize)) { + float[] MEAN = new float[]{0.485F, 0.456F, 0.406F}; + float[] STD = new float[]{0.229F, 0.224F, 0.225F}; + this.addTransform(new Normalize(MEAN, STD)); + } else if (!"false".equals(normalize)) { + String[] tokens = normalize.split("\\s*,\\s*"); + if (tokens.length != 6) { + throw new IllegalArgumentException("Invalid normalize value: " + normalize); + } + + float[] mean = new float[]{Float.parseFloat(tokens[0]), Float.parseFloat(tokens[1]), Float.parseFloat(tokens[2])}; + float[] std = new float[]{Float.parseFloat(tokens[3]), Float.parseFloat(tokens[4]), Float.parseFloat(tokens[5])}; + this.addTransform(new Normalize(mean, std)); + } + + String range = (String)arguments.get("range"); + if ("0,1".equals(range)) { + this.addTransform((a) -> { + return a.div(255.0F); + }); + } else if ("-1,1".equals(range)) { + this.addTransform((a) -> { + return a.div(128.0F).sub(1); + }); + } + + if (arguments.containsKey("batchifier")) { + this.batchifier = Batchifier.fromString((String)arguments.get("batchifier")); + } + + } + } + + public static enum YoloOutputType { + BOX, + DETECT, + AUTO; + + private YoloOutputType() { + } + } + + + +} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java b/face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java rename to face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java b/face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java similarity index 99% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java rename to face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java index 35d61df..d483b29 100644 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java +++ b/face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java @@ -190,7 +190,7 @@ public class FaceUtils { } } graphics.dispose(); - ImageIO.write(sourceImage, "jpg", new File(savePath)); + ImageIO.write(sourceImage, "png", new File(savePath)); } /** @@ -639,7 +639,7 @@ public class FaceUtils { } } graphics.dispose(); - ImageIO.write(sourceImage, "jpg", new File(savePath)); + ImageIO.write(sourceImage, "png", new File(savePath)); } private static void drawMultilineTextWithBackground(Graphics2D g, List lines, int x, int y) { diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/SVDUtils.java b/face/src/main/java/cn/smartjavaai/face/utils/SVDUtils.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/SVDUtils.java rename to face/src/main/java/cn/smartjavaai/face/utils/SVDUtils.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java b/face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java rename to face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/SimilarityUtil.java b/face/src/main/java/cn/smartjavaai/face/utils/SimilarityUtil.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/SimilarityUtil.java rename to face/src/main/java/cn/smartjavaai/face/utils/SimilarityUtil.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/VectorUtils.java b/face/src/main/java/cn/smartjavaai/face/utils/VectorUtils.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/utils/VectorUtils.java rename to face/src/main/java/cn/smartjavaai/face/utils/VectorUtils.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/config/MilvusConfig.java b/face/src/main/java/cn/smartjavaai/face/vector/config/MilvusConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/config/MilvusConfig.java rename to face/src/main/java/cn/smartjavaai/face/vector/config/MilvusConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/config/SQLiteConfig.java b/face/src/main/java/cn/smartjavaai/face/vector/config/SQLiteConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/config/SQLiteConfig.java rename to face/src/main/java/cn/smartjavaai/face/vector/config/SQLiteConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/config/VectorDBConfig.java b/face/src/main/java/cn/smartjavaai/face/vector/config/VectorDBConfig.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/config/VectorDBConfig.java rename to face/src/main/java/cn/smartjavaai/face/vector/config/VectorDBConfig.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/constant/VectorDBConstants.java b/face/src/main/java/cn/smartjavaai/face/vector/constant/VectorDBConstants.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/constant/VectorDBConstants.java rename to face/src/main/java/cn/smartjavaai/face/vector/constant/VectorDBConstants.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/MilvusClient.java b/face/src/main/java/cn/smartjavaai/face/vector/core/MilvusClient.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/MilvusClient.java rename to face/src/main/java/cn/smartjavaai/face/vector/core/MilvusClient.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java b/face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java rename to face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBClient.java b/face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBClient.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBClient.java rename to face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBClient.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBFactory.java b/face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBFactory.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBFactory.java rename to face/src/main/java/cn/smartjavaai/face/vector/core/VectorDBFactory.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/entity/FaceVector.java b/face/src/main/java/cn/smartjavaai/face/vector/entity/FaceVector.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/entity/FaceVector.java rename to face/src/main/java/cn/smartjavaai/face/vector/entity/FaceVector.java diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/exception/VectorDBException.java b/face/src/main/java/cn/smartjavaai/face/vector/exception/VectorDBException.java similarity index 100% rename from smartjavaai-face/src/main/java/cn/smartjavaai/face/vector/exception/VectorDBException.java rename to face/src/main/java/cn/smartjavaai/face/vector/exception/VectorDBException.java diff --git a/smartjavaai-face/src/main/resources/db/schema.sql b/face/src/main/resources/db/schema.sql similarity index 100% rename from smartjavaai-face/src/main/resources/db/schema.sql rename to face/src/main/resources/db/schema.sql diff --git a/face/src/test/java/Test.java b/face/src/test/java/Test.java new file mode 100644 index 0000000..8462133 --- /dev/null +++ b/face/src/test/java/Test.java @@ -0,0 +1,87 @@ +import ai.djl.Application; +import ai.djl.repository.Artifact; +import ai.djl.repository.MRL; +import ai.djl.repository.zoo.ModelNotFoundException; +import ai.djl.repository.zoo.ModelZoo; +import ai.djl.util.JsonUtils; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.face.config.FaceDetConfig; +import cn.smartjavaai.face.config.FaceRecConfig; +import cn.smartjavaai.face.constant.FaceDetectConstant; +import cn.smartjavaai.face.enums.FaceDetModelEnum; +import cn.smartjavaai.face.enums.FaceRecModelEnum; +import cn.smartjavaai.face.enums.SimilarityType; +import cn.smartjavaai.face.factory.FaceDetModelFactory; +import cn.smartjavaai.face.factory.FaceRecModelFactory; +import cn.smartjavaai.face.model.facedect.FaceDetModel; +import cn.smartjavaai.face.model.facerec.FaceRecModel; +import cn.smartjavaai.face.utils.SimilarityUtil; +import lombok.extern.slf4j.Slf4j; + +import java.io.IOException; +import java.util.List; +import java.util.Map; + +/** + * @author dwj + * @date 2025/7/25 + */ +@Slf4j +public class Test { + + /** + * 获取人脸检测模型 + * @return + */ + public static FaceDetModel getFaceDetModel(){ + FaceDetConfig config = new FaceDetConfig(); + config.setModelEnum(FaceDetModelEnum.RETINA_FACE);//人脸检测模型 + config.setConfidenceThreshold(FaceDetectConstant.DEFAULT_CONFIDENCE_THRESHOLD);//只返回相似度大于该值的人脸 + config.setNmsThresh(FaceDetectConstant.NMS_THRESHOLD);//用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 + return FaceDetModelFactory.getInstance().getModel(config); + } + + /** + * 获取人脸识别模型 + * @return + */ + public static FaceRecModel getFaceRecModel(){ + FaceRecConfig config = new FaceRecConfig(); + config.setModelEnum(FaceRecModelEnum.ELASTIC_FACE_MODEL); + config.setModelPath("/Users/wenjie/Documents/develop/model/arcfaceresnet100-11-int8.onnx"); +// config.setModelPath("/Users/xxx/Documents/develop/model/InsightFace/model_mobilefacenet.pt"); + //裁剪人脸:如果图片已经是裁剪过的,则请将此参数设置为false + config.setCropFace(true); + //开启人脸对齐:适用于人脸不正的场景,开启将提升人脸特征准确度,关闭可以提升性能 + config.setAlign(true); + //指定人脸检测模型 + config.setDetectModel(getFaceDetModel()); + return FaceRecModelFactory.getInstance().getModel(config); + } + + public static void main(String[] args) throws ModelNotFoundException, IOException { +// boolean withArtifacts = +// args.length > 0 && ("--artifact".equals(args[0]) || "-a".equals(args[0])); +// if (!withArtifacts) { +// logger.info("============================================================"); +// logger.info("user ./gradlew listModel --args='-a' to show artifact detail"); +// logger.info("============================================================"); +// } +// Map> models = ModelZoo.listModels(); +// for (Map.Entry> entry : models.entrySet()) { +// String appName = entry.getKey().toString(); +// for (MRL mrl : entry.getValue()) { +// if (withArtifacts) { +// for (Artifact artifact : mrl.listArtifacts()) { +// log.info("{} djl://{}", appName, artifact); +// } +// } else { +// log.info("{} {}", appName, mrl); +// } +// } +// } + } + + +} diff --git a/smartjavaai-ocr/pom.xml b/ocr/pom.xml similarity index 96% rename from smartjavaai-ocr/pom.xml rename to ocr/pom.xml index 551e0c6..f2b3263 100644 --- a/smartjavaai-ocr/pom.xml +++ b/ocr/pom.xml @@ -6,16 +6,16 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - smartjavaai-ocr + ocr cn.smartjavaai - smartjavaai-common + common ${project.version} @@ -42,8 +42,8 @@
- 1.0.23 - smartjavaai-ocr + 1.0.24 + ocr SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/DirectionModelConfig.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/DirectionModelConfig.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/DirectionModelConfig.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/DirectionModelConfig.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/OcrDetModelConfig.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrDetModelConfig.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/OcrDetModelConfig.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/OcrDetModelConfig.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecModelConfig.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecModelConfig.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecModelConfig.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecModelConfig.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/OcrRecOptions.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/PlateDetModelConfig.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/PlateDetModelConfig.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/PlateDetModelConfig.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/PlateDetModelConfig.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/PlateRecModelConfig.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/PlateRecModelConfig.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/PlateRecModelConfig.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/PlateRecModelConfig.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/TableStructureConfig.java b/ocr/src/main/java/cn/smartjavaai/ocr/config/TableStructureConfig.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/config/TableStructureConfig.java rename to ocr/src/main/java/cn/smartjavaai/ocr/config/TableStructureConfig.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/DirectionInfo.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/DirectionInfo.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/DirectionInfo.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/DirectionInfo.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/IdCardInfo.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/IdCardInfo.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/IdCardInfo.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/IdCardInfo.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/ImageInfo.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/ImageInfo.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/ImageInfo.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/ImageInfo.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrBox.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrBox.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrBox.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrBox.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrItem.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrItem.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrItem.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrItem.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateInfo.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateInfo.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateInfo.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateInfo.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateResult.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateResult.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateResult.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/PlateResult.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBox.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBox.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBox.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBox.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBoxCompX.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBoxCompX.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBoxCompX.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/RotatedBoxCompX.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/TableStructureResult.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/TableStructureResult.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/entity/TableStructureResult.java rename to ocr/src/main/java/cn/smartjavaai/ocr/entity/TableStructureResult.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/AngleEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/AngleEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/AngleEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/AngleEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonDetModelEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonDetModelEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonDetModelEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonDetModelEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonRecModelEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonRecModelEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonRecModelEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/CommonRecModelEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/DirectionModelEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/DirectionModelEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/DirectionModelEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/DirectionModelEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateDetModelEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateDetModelEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateDetModelEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateDetModelEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateRecModelEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateRecModelEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateRecModelEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateRecModelEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateType.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateType.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateType.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/PlateType.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/TableStructureModelEnum.java b/ocr/src/main/java/cn/smartjavaai/ocr/enums/TableStructureModelEnum.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/enums/TableStructureModelEnum.java rename to ocr/src/main/java/cn/smartjavaai/ocr/enums/TableStructureModelEnum.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/exception/OcrException.java b/ocr/src/main/java/cn/smartjavaai/ocr/exception/OcrException.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/exception/OcrException.java rename to ocr/src/main/java/cn/smartjavaai/ocr/exception/OcrException.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModel.java diff --git a/smartjavaai-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 similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/criteria/OcrCommonDetCriterialFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/criteria/OcrCommonDetCriterialFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/criteria/OcrCommonDetCriterialFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/criteria/OcrCommonDetCriterialFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/translator/PPOCRDetTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/translator/PPOCRDetTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/translator/PPOCRDetTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/translator/PPOCRDetTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/OcrDirectionModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/OcrDirectionModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/OcrDirectionModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/OcrDirectionModel.java diff --git a/smartjavaai-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 similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/criteria/DirectionCriteriaFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/criteria/DirectionCriteriaFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/criteria/DirectionCriteriaFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/criteria/DirectionCriteriaFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/translator/PpWordRotateTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/translator/PpWordRotateTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/translator/PpWordRotateTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/translator/PpWordRotateTranslator.java diff --git a/smartjavaai-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 similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java diff --git a/smartjavaai-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 similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/criteria/OcrCommonRecCriterialFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/criteria/OcrCommonRecCriterialFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/criteria/OcrCommonRecCriterialFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/criteria/OcrCommonRecCriterialFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/translator/PPOCRRecTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/translator/PPOCRRecTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/translator/PPOCRRecTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/translator/PPOCRRecTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java similarity index 93% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java index 5fb00c9..ba95975 100644 --- a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java @@ -275,7 +275,7 @@ public class CRNNPlateRecModel implements PlateRecModel{ } BufferedImage bufferedImage = OpenCVUtils.mat2Image((Mat)img.getWrappedImage()); OcrUtils.drawPlateInfo(bufferedImage, plateResult.getData()); - ImageIO.write(bufferedImage, "jpg", new File(outputPath)); + ImageIO.write(bufferedImage, "png", new File(outputPath)); return R.ok(); } catch (IOException e) { throw new OcrException(e); @@ -291,28 +291,18 @@ public class CRNNPlateRecModel implements PlateRecModel{ if(!ImageUtils.isImageValid(sourceImage)){ return R.fail(R.Status.INVALID_IMAGE); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); try { - R> plateResult = recognize(img); + R> plateResult = recognize(sourceImage); if(!plateResult.isSuccess()){ return R.fail(plateResult.getCode(), plateResult.getMessage()); } if(CollectionUtils.isEmpty(plateResult.getData())){ return R.fail(R.Status.NO_OBJECT_DETECTED); } - OcrUtils.drawPlateInfo((Mat)img.getWrappedImage(), plateResult.getData()); - ByteArrayOutputStream outputStream = new ByteArrayOutputStream(); - // 调用 save 方法将 Image 写入字节流 - img.save(outputStream, "png"); - // 将字节流转换为 BufferedImage - byte[] imageBytes = outputStream.toByteArray(); - return R.ok(ImageIO.read(new ByteArrayInputStream(imageBytes))); - } catch (IOException e) { + OcrUtils.drawPlateInfo(sourceImage, plateResult.getData()); + return R.ok(sourceImage); + } catch (Exception e) { throw new OcrException("导出图片失败", e); - } finally { - if (img != null){ - ((Mat)img.getWrappedImage()).release(); - } } } diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateDetCriterialFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateDetCriterialFactory.java similarity index 79% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateDetCriterialFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateDetCriterialFactory.java index 2b45326..8dfe688 100644 --- a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateDetCriterialFactory.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateDetCriterialFactory.java @@ -10,6 +10,7 @@ import cn.smartjavaai.ocr.config.PlateDetModelConfig; import cn.smartjavaai.ocr.enums.PlateDetModelEnum; import cn.smartjavaai.ocr.model.plate.translator.Yolo5PlateDetectTranslator; import cn.smartjavaai.ocr.model.plate.translator.Yolov7PlateDetectTranslator; +import cn.smartjavaai.ocr.model.plate.translator.Yolov8PlateDetectTranslator; import org.apache.commons.lang3.StringUtils; import java.nio.file.Paths; @@ -55,6 +56,17 @@ public class PlateDetCriterialFactory { .optProgress(new ProgressBar()) .build(); } +// else if (config.getModelEnum() == PlateDetModelEnum.YOLOV8){ +// criteria = +// Criteria.builder() +// .optEngine("OnnxRuntime") +// .setTypes(Image.class, DetectedObjects.class) +// .optModelPath(Paths.get(config.getModelPath())) +// .optTranslator(new Yolov8PlateDetectTranslator(params)) +// .optDevice(device) +// .optProgress(new ProgressBar()) +// .build(); +// } return criteria; } } diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateRecCriterialFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateRecCriterialFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateRecCriterialFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/criteria/PlateRecCriterialFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/CRNNPlateRecTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/CRNNPlateRecTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/CRNNPlateRecTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/CRNNPlateRecTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolo5PlateDetectTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolo5PlateDetectTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolo5PlateDetectTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolo5PlateDetectTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov7PlateDetectTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov7PlateDetectTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov7PlateDetectTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov7PlateDetectTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov8PlateDetectTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov8PlateDetectTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov8PlateDetectTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/plate/translator/Yolov8PlateDetectTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/criteria/StructureCriteriaFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/criteria/StructureCriteriaFactory.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/criteria/StructureCriteriaFactory.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/table/criteria/StructureCriteriaFactory.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/translator/TableStructTranslator.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/translator/TableStructTranslator.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/model/table/translator/TableStructTranslator.java rename to ocr/src/main/java/cn/smartjavaai/ocr/model/table/translator/TableStructTranslator.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java b/ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java rename to ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/ConvertHtml2Excel.java b/ocr/src/main/java/cn/smartjavaai/ocr/utils/ConvertHtml2Excel.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/ConvertHtml2Excel.java rename to ocr/src/main/java/cn/smartjavaai/ocr/utils/ConvertHtml2Excel.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/CrossRangeCellMeta.java b/ocr/src/main/java/cn/smartjavaai/ocr/utils/CrossRangeCellMeta.java similarity index 100% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/CrossRangeCellMeta.java rename to ocr/src/main/java/cn/smartjavaai/ocr/utils/CrossRangeCellMeta.java diff --git a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java b/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java similarity index 98% rename from smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java rename to ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java index d392f34..2c87c7a 100644 --- a/smartjavaai-ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java @@ -4,6 +4,7 @@ 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.Landmark; import ai.djl.modality.cv.util.NDImageUtils; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; @@ -15,6 +16,7 @@ import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.common.utils.PointUtils; import cn.smartjavaai.ocr.entity.*; +import cn.smartjavaai.ocr.entity.RotatedBox; import cn.smartjavaai.ocr.enums.AngleEnum; import cn.smartjavaai.ocr.enums.PlateType; import cn.smartjavaai.ocr.opencv.OcrNDArrayUtils; @@ -60,7 +62,7 @@ public class OcrUtils { /** * 转换为OcrBox - * @param dt_boxes + * @param ndLists * @return */ public static List> convertToOcrBox(List ndLists) { @@ -373,9 +375,11 @@ public class OcrUtils { DetectedObjects.DetectedObject result = (DetectedObjects.DetectedObject)iterator.next(); BoundingBox box = result.getBoundingBox(); List keyPoints = new ArrayList(); - box.getBounds().getPath().forEach(point -> { - keyPoints.add(new Point(point.getX(), point.getY())); - }); + if(box instanceof Landmark){ + box.getBounds().getPath().forEach(point -> { + keyPoints.add(new Point(point.getX(), point.getY())); + }); + } int x = (int)(box.getBounds().getX() * image.getWidth()); int y = (int)(box.getBounds().getY() * image.getHeight()); int width = (int)(box.getBounds().getWidth() * image.getWidth()); @@ -385,7 +389,6 @@ public class OcrUtils { if (y < 0) y = 0; if (x + width > image.getWidth()) width = image.getWidth() - x; if (y + height > image.getHeight()) height = image.getHeight() - y; - PlateInfo plateInfo = new PlateInfo(); plateInfo.setPlateType(PlateType.fromClassName(detectedObjects.getClassNames().get(index))); plateInfo.setScore(detectedObjects.getProbabilities().get(index).floatValue()); diff --git a/pom.xml b/pom.xml index 0037781..fbccb3a 100644 --- a/pom.xml +++ b/pom.xml @@ -7,25 +7,25 @@ SmartJavaAI cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 pom SmartJavaAI - smartjavaai-face - smartjavaai-translate - smartjavaai-common - smartjavaai-objectdetection - smartjavaai-all - smartjavaai-ocr - smartjavaai-bom - smartjavaai-speech + face + translate + common + vision + all + ocr + bom + speech 8 8 UTF-8 - 0.32.0 + 0.34.0 @@ -175,11 +175,11 @@ - - ai.djl.tensorrt - tensorrt - runtime - + + + + + cn.hutool @@ -339,4 +339,5 @@ + diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java b/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java deleted file mode 100644 index 5ffd4ab..0000000 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java +++ /dev/null @@ -1,143 +0,0 @@ -package cn.smartjavaai.common.utils; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDManager; -import org.opencv.core.CvType; -import org.opencv.core.Mat; -import org.opencv.core.Point; -import org.opencv.core.Scalar; -import org.opencv.imgproc.Imgproc; - -import java.awt.*; -import java.awt.image.BufferedImage; -import java.awt.image.DataBufferByte; - -/** - * OpenCV 工具类 - */ -public class OpenCVUtils { - /** - * canny算法,边缘检测 - * - * @param src - * @return - */ - public static Mat canny(Mat src) { - Mat mat = src.clone(); - Imgproc.Canny(src, mat, 100, 200); - return mat; - } - - /** - * 画线 - * - * @param mat - * @param point1 - * @param point2 - */ - public static void line(Mat mat, Point point1, Point point2) { - Imgproc.line(mat, point1, point2, new Scalar(255, 255, 255), 1); - } - - /** - * NDArray to opencv_core.Mat - * - * @param manager - * @param srcPoints - * @param dstPoints - * @return - */ - public static Mat toOpenCVMat(NDManager manager, NDArray srcPoints, NDArray dstPoints) { - NDArray svdMat = SVDUtils.transformationFromPoints(manager, srcPoints, dstPoints); - double[] doubleArray = svdMat.toDoubleArray(); - Mat newSvdMat = new Mat(2, 3, CvType.CV_64F); - for (int i = 0; i < 2; i++) { - for (int j = 0; j < 3; j++) { - newSvdMat.put(i, j, doubleArray[i * 3 + j]); - } - } - return newSvdMat; - } - - /** - * double[][] points array to Mat - * @param points - * @return - */ - public static Mat toOpenCVMat(double[][] points) { - Mat mat = new Mat(5, 2, CvType.CV_64F); - for (int i = 0; i < 5; i++) { - for (int j = 0; j < 2; j++) { - mat.put(i, j, points[i * 5 + j]); - } - } - return mat; - } - - /** - * 变换矩阵的逆矩阵 - * - * @param src - * @return - */ - public static Mat invertAffineTransform(Mat src) { - Mat dst = src.clone(); - Imgproc.invertAffineTransform(src, dst); - return dst; - } - - /** - * Mat to BufferedImage - * - * @param mat - * @return - */ - public static BufferedImage mat2Image(Mat mat) { - int width = mat.width(); - int height = mat.height(); - byte[] data = new byte[width * height * (int) mat.elemSize()]; - Imgproc.cvtColor(mat, mat, 4); - mat.get(0, 0, data); - BufferedImage ret = new BufferedImage(width, height, 5); - ret.getRaster().setDataElements(0, 0, width, height, data); - return ret; - } - - /** - * BufferedImage to Mat - * - * @param img - * @return - */ - public static Mat image2Mat(BufferedImage img) { - int width = img.getWidth(); - int height = img.getHeight(); - - // 强制转换为 TYPE_3BYTE_BGR,自动去除透明通道 - BufferedImage convertedImg = new BufferedImage(width, height, BufferedImage.TYPE_3BYTE_BGR); - Graphics2D g2d = convertedImg.createGraphics(); - g2d.drawImage(img, 0, 0, null); - g2d.dispose(); - - byte[] data = ((DataBufferByte) convertedImg.getRaster().getDataBuffer()).getData(); - Mat mat = new Mat(height, width, CvType.CV_8UC3); - mat.put(0, 0, data); - return mat; - } - - /** - * 透视变换 - * - * @param src - * @param srcPoints - * @param dstPoints - * @return - */ - public static Mat perspectiveTransform(Mat src, Mat srcPoints, Mat dstPoints) { - Mat dst = src.clone(); - Mat warp_mat = Imgproc.getPerspectiveTransform(srcPoints, dstPoints); - Imgproc.warpPerspective(src, dst, warp_mat, dst.size()); - warp_mat.release(); - return dst; - } -} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceDetModelEnum.java b/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceDetModelEnum.java deleted file mode 100644 index fb4fce3..0000000 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/enums/FaceDetModelEnum.java +++ /dev/null @@ -1,37 +0,0 @@ -package cn.smartjavaai.face.enums; - -/** - * 人脸检测模型枚举 - * @author dwj - */ -public enum FaceDetModelEnum { - - RETINA_FACE("RetinaFaceModel"), - ULTRA_LIGHT_FAST_GENERIC_FACE("UltraLightFastGenericFaceModel"), - SEETA_FACE6_MODEL("SeetaFace6Model"); - - private final String modelClassName; - - FaceDetModelEnum(String modelClassName) { - this.modelClassName = modelClassName; - } - - public String getModelClassName() { - return modelClassName; - } - - /** - * 根据名称获取枚举 (忽略大小写和下划线变体) - */ - public static FaceDetModelEnum fromName(String name) { - String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); - for (FaceDetModelEnum model : values()) { - if (model.name().replaceAll("_", "").equals(formatted)) { - return model; - } - } - throw new IllegalArgumentException("未知模型名称: " + name); - } - - -} diff --git a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/criterial/FaceDetCriteriaFactory.java b/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/criterial/FaceDetCriteriaFactory.java deleted file mode 100644 index 38edb36..0000000 --- a/smartjavaai-face/src/main/java/cn/smartjavaai/face/model/facedect/criterial/FaceDetCriteriaFactory.java +++ /dev/null @@ -1,69 +0,0 @@ -package cn.smartjavaai.face.model.facedect.criterial; - -import ai.djl.Device; -import ai.djl.modality.Classifications; -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.face.config.FaceDetConfig; -import cn.smartjavaai.face.config.FaceExpressionConfig; -import cn.smartjavaai.face.constant.FaceDetectConstant; -import cn.smartjavaai.face.constant.RetinaFaceConstant; -import cn.smartjavaai.face.constant.UltraLightFastGenericFaceConstant; -import cn.smartjavaai.face.enums.ExpressionModelEnum; -import cn.smartjavaai.face.enums.FaceDetModelEnum; -import cn.smartjavaai.face.model.expression.translator.DenseNetEmotionTranslator; -import cn.smartjavaai.face.model.expression.translator.FrEmotionTranslator; -import cn.smartjavaai.face.translator.FaceDetectionTranslator; -import org.apache.commons.lang3.StringUtils; - -import java.nio.file.Paths; -import java.util.Objects; - -/** - * 人脸检测 Criteria构建工厂 - * @author dwj - */ -public class FaceDetCriteriaFactory { - - public static Criteria createCriteria(FaceDetConfig config) { - Device device = null; - if(!Objects.isNull(config.getDevice())){ - device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); - } - Criteria criteria = null; - if(config.getModelEnum() == FaceDetModelEnum.RETINA_FACE){ - FaceDetectionTranslator translator = - new FaceDetectionTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), RetinaFaceConstant.variance, FaceDetectConstant.MAX_FACE_LIMIT, RetinaFaceConstant.scales, RetinaFaceConstant.steps); - criteria = - Criteria.builder() - .setTypes(Image.class, DetectedObjects.class) - .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : RetinaFaceConstant.MODEL_URL) - // Load model from local file, e.g: - .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) - .optModelName("retinaface") // specify model file prefix - .optTranslator(translator) - .optDevice(device) - .optProgress(new ProgressBar()) - .optEngine("PyTorch") // Use PyTorch engine - .build(); - }else if (config.getModelEnum() == FaceDetModelEnum.ULTRA_LIGHT_FAST_GENERIC_FACE){ - FaceDetectionTranslator translator = - new FaceDetectionTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), UltraLightFastGenericFaceConstant.variance, FaceDetectConstant.MAX_FACE_LIMIT, UltraLightFastGenericFaceConstant.scales, UltraLightFastGenericFaceConstant.steps); - criteria = - Criteria.builder() - .setTypes(Image.class, DetectedObjects.class) - .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : UltraLightFastGenericFaceConstant.MODEL_URL) - .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) - .optTranslator(translator) - .optProgress(new ProgressBar()) - .optDevice(device) - .optEngine("PyTorch") // Use PyTorch engine - .build(); - } - return criteria; - } - -} diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/utils/DetectorUtils.java b/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/utils/DetectorUtils.java deleted file mode 100644 index c41b220..0000000 --- a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/utils/DetectorUtils.java +++ /dev/null @@ -1,67 +0,0 @@ -package cn.smartjavaai.objectdetection.utils; - -import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.output.BoundingBox; -import ai.djl.modality.cv.output.DetectedObjects; -import cn.smartjavaai.common.entity.DetectionInfo; -import cn.smartjavaai.common.entity.DetectionRectangle; -import cn.smartjavaai.common.entity.DetectionResponse; -import cn.smartjavaai.common.entity.ObjectDetInfo; -import cn.smartjavaai.common.utils.ImageUtils; - -import javax.imageio.ImageIO; -import java.awt.*; -import java.awt.image.BufferedImage; -import java.io.File; -import java.io.IOException; -import java.util.ArrayList; -import java.util.Iterator; -import java.util.List; -import java.util.Objects; - -/** - * 目标检测相关工具类 - * @author dwj - * @date 2025/4/9 - */ -public class DetectorUtils { - - - /** - * 转换为FaceDetectedResult - * @param detection - * @param img - * @return - */ - public static DetectionResponse convertToDetectionResponse(DetectedObjects detection, Image img){ - if(Objects.isNull(detection) || Objects.isNull(detection.getProbabilities()) - || detection.getProbabilities().isEmpty() || Objects.isNull(detection.items()) || detection.items().isEmpty()){ - return null; - } - DetectionResponse detectionResponse = new DetectionResponse(); - List detectionInfoList = new ArrayList(); - List detectedObjectList = detection.items(); - Iterator iterator = detectedObjectList.iterator(); - int index = 0; - while(iterator.hasNext()) { - DetectedObjects.DetectedObject result = (DetectedObjects.DetectedObject)iterator.next(); - String className = result.getClassName(); - BoundingBox box = result.getBoundingBox(); - int x = (int)(box.getBounds().getX() * (double)img.getWidth()); - int y = (int)(box.getBounds().getY() * (double)img.getHeight()); - int width = (int)(box.getBounds().getWidth() * (double)img.getWidth()); - int height = (int)(box.getBounds().getHeight() * (double)img.getHeight()); - DetectionRectangle rectangle = new DetectionRectangle(x, y, width, height); - DetectionInfo detectionInfo = new DetectionInfo(rectangle, detection.getProbabilities().get(index).floatValue()); - ObjectDetInfo objectDetInfo = new ObjectDetInfo(className); - detectionInfo.setObjectDetInfo(objectDetInfo); - detectionInfoList.add(detectionInfo); - index++; - } - detectionResponse.setDetectionInfoList(detectionInfoList); - return detectionResponse; - } - - - -} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/CausalLMOutput.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/CausalLMOutput.java deleted file mode 100644 index acbc11d..0000000 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/CausalLMOutput.java +++ /dev/null @@ -1,33 +0,0 @@ -package cn.smartjavaai.translation.entity; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; - -/** - * 解码输出对象 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class CausalLMOutput { - private NDArray logits; - private NDList pastKeyValuesList; - - public CausalLMOutput(NDArray logits, NDList pastKeyValues) { - this.logits = logits; - this.pastKeyValuesList = pastKeyValues; - } - - public NDArray getLogits() { - return logits; - } - - public void setLogits(NDArray logits) { - this.logits = logits; - } - - public NDList getPastKeyValuesList() { - return pastKeyValuesList; - } -} \ No newline at end of file diff --git a/smartjavaai-speech/pom.xml b/speech/pom.xml similarity index 96% rename from smartjavaai-speech/pom.xml rename to speech/pom.xml index 07a5e29..d554c42 100644 --- a/smartjavaai-speech/pom.xml +++ b/speech/pom.xml @@ -6,10 +6,10 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - smartjavaai-speech + speech 11 @@ -20,7 +20,7 @@ cn.smartjavaai - smartjavaai-common + common ${project.version} @@ -51,8 +51,8 @@ - 1.0.23 - smartjavaai-speech + 1.0.24 + speech SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/audio/SmartAudioFactory.java b/speech/src/main/java/cn/smartjavaai/speech/asr/audio/SmartAudioFactory.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/audio/SmartAudioFactory.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/audio/SmartAudioFactory.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/config/AsrModelConfig.java b/speech/src/main/java/cn/smartjavaai/speech/asr/config/AsrModelConfig.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/config/AsrModelConfig.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/config/AsrModelConfig.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrResult.java b/speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrResult.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrResult.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrResult.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrSegment.java b/speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrSegment.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrSegment.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/entity/AsrSegment.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/RecParams.java b/speech/src/main/java/cn/smartjavaai/speech/asr/entity/RecParams.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/RecParams.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/entity/RecParams.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/VoskParams.java b/speech/src/main/java/cn/smartjavaai/speech/asr/entity/VoskParams.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/VoskParams.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/entity/VoskParams.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/WhisperParams.java b/speech/src/main/java/cn/smartjavaai/speech/asr/entity/WhisperParams.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/entity/WhisperParams.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/entity/WhisperParams.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/enums/AsrModelEnum.java b/speech/src/main/java/cn/smartjavaai/speech/asr/enums/AsrModelEnum.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/enums/AsrModelEnum.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/enums/AsrModelEnum.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/exception/AsrException.java b/speech/src/main/java/cn/smartjavaai/speech/asr/exception/AsrException.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/exception/AsrException.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/exception/AsrException.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java b/speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java b/speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java b/speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java b/speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/pool/WhisperStatePool.java b/speech/src/main/java/cn/smartjavaai/speech/asr/pool/WhisperStatePool.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/asr/pool/WhisperStatePool.java rename to speech/src/main/java/cn/smartjavaai/speech/asr/pool/WhisperStatePool.java diff --git a/smartjavaai-speech/src/main/java/cn/smartjavaai/speech/utils/AudioUtils.java b/speech/src/main/java/cn/smartjavaai/speech/utils/AudioUtils.java similarity index 100% rename from smartjavaai-speech/src/main/java/cn/smartjavaai/speech/utils/AudioUtils.java rename to speech/src/main/java/cn/smartjavaai/speech/utils/AudioUtils.java diff --git a/speech/src/test/java/Test.java b/speech/src/test/java/Test.java new file mode 100644 index 0000000..ed10650 --- /dev/null +++ b/speech/src/test/java/Test.java @@ -0,0 +1,218 @@ +import ai.djl.util.JsonUtils; +import cn.smartjavaai.common.entity.Language; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.speech.asr.config.AsrModelConfig; +import cn.smartjavaai.speech.asr.entity.AsrResult; +import cn.smartjavaai.speech.asr.entity.AsrSegment; +import cn.smartjavaai.speech.asr.entity.RecParams; +import cn.smartjavaai.speech.asr.entity.WhisperParams; +import cn.smartjavaai.speech.asr.enums.AsrModelEnum; +import cn.smartjavaai.speech.asr.model.VoskRecognizer; +import cn.smartjavaai.speech.asr.model.WhisperRecognizer; +import io.github.givimad.whisperjni.WhisperFullParams; +import io.github.givimad.whisperjni.WhisperSamplingStrategy; +import lombok.extern.slf4j.Slf4j; +import ws.schild.jave.Encoder; +import ws.schild.jave.EncoderException; +import ws.schild.jave.InputFormatException; +import ws.schild.jave.MultimediaObject; +import ws.schild.jave.encode.AudioAttributes; +import ws.schild.jave.encode.EncodingAttributes; +import ws.schild.jave.info.AudioInfo; +import ws.schild.jave.info.MultimediaInfo; + +import java.io.*; + +/** + * @author dwj + * @date 2025/8/1 + */ +@Slf4j +public class Test { + + public static void main(String[] args) { + +// System.out.println("TMPDIR = " + System.getProperty("java.io.tmpdir")); +// +//// System.setProperty("io.github.givimad.whisperjni.libdir","/Users/wenjie/smartjavaai_cache/whisper"); +// WhisperRecognizer whisperRecognizer = new WhisperRecognizer(); +// AsrModelConfig config = new AsrModelConfig(); +// config.setModelEnum(AsrModelEnum.WHISPER); +//// config.setModelPath("/Users/wenjie/Downloads/ggml-medium.bin"); +// config.setModelPath("/Users/wenjie/Documents/develop/model/speech/ggml-medium.bin"); +// whisperRecognizer.loadModel(config); +// WhisperParams params = new WhisperParams(); +// WhisperFullParams params1 = new WhisperFullParams(WhisperSamplingStrategy.BEAN_SEARCH); +//// params1.detectLanguage = true; +// params1.language = Language.ZH.getCode(); +// //params1.translate = true; +// params1.initialPrompt = "语音模型"; +// params.setParams(params1); +//// params1.printTimestamps = true; +// params1.printRealtime = true; +// //params.setLanguage(Language.ZH); +// R result = whisperRecognizer.recognize("/Users/wenjie/Downloads/友谊大街.m4a",params); +// if (result.isSuccess()){ +// System.out.println("结果:" + JsonUtils.toJson(result.getData())); +// }else{ +// System.out.println(result.getMessage()); +// } +// +// while (true){ +// try { +// Thread.sleep(10); +// } catch (InterruptedException e) { +// throw new RuntimeException(e); +// } +// } + + testVosk(); + +// testConvert(); + } + + + + public static void testVosk(){ + System.load("/Users/wenjie/Downloads/vosk-arrch64-dylib-main/libvosk.dylib"); + VoskRecognizer voskRecognizer = new VoskRecognizer(); + AsrModelConfig config = new AsrModelConfig(); + config.setModelEnum(AsrModelEnum.VOSK); + config.setModelPath("/Users/wenjie/Documents/develop/model/speech/vosk-model-cn-0.22"); + voskRecognizer.loadModel(config); + + R result = voskRecognizer.recognize("/Users/wenjie/Documents/idea_workplace/SmartJavaAI/examples/speech-examples/src/main/resources/lff_zh.mp3"); + if (result.isSuccess()){ + System.out.println("结果:" + JsonUtils.toJson(result.getData())); + }else{ + System.out.println(result.getMessage()); + } + +// while (true){ +// try { +// Thread.sleep(10); +// } catch (InterruptedException e) { +// throw new RuntimeException(e); +// } +// } + } + + public static void testConvert(){ + getAudioFormatConversionIns("/Users/wenjie/Downloads/中国建设银行(包头当代支行).m4a","/Users/wenjie/Downloads/test1.wav","wav"); + } + + + /** + * 音频格式转换 + * + * @param sourceFilePath + * @param targetFilePath + * @param format wav/mp3/amr + * @return + */ + public static byte[] getAudioFormatConversionBytes(String sourceFilePath,String targetFilePath,String format) { + InputStream fis = null; + ByteArrayOutputStream bos = null; + byte[] bytes = null; + try { + File sourceFile = new File(sourceFilePath); + if (sourceFile.isFile()) { + File targetFile = new File(targetFilePath); + + // 音频格式转换 + audioFormatConversion(sourceFile, targetFile, format); + + fis = new FileInputStream(targetFile); + bos = new ByteArrayOutputStream(); + + byte[] buffer = new byte[1024]; + int bytesRead; + while ((bytesRead = fis.read(buffer)) != -1) { + bos.write(buffer, 0, bytesRead); + } + bytes = bos.toByteArray(); + } + } catch (Exception e) { + log.error("音频格式转换异常:" + e.getMessage(), e); + return null; + } finally { + try { + if (fis != null) { + fis.close(); + } + if (bos != null) { + bos.close(); + } + } catch (IOException e) { + log.error("音频格式转换资源关闭异常:" + e.getMessage(), e); + } + } + return bytes; + } + + /** + * 音频格式转换 + * + * @param sourceFilePath + * @param targetFilePath + * @param format wav/mp3/amr + * @return + */ + public static InputStream getAudioFormatConversionIns(String sourceFilePath, String targetFilePath, String format) { + try { + File sourceFile = new File(sourceFilePath); + if (sourceFile.isFile()) { + File targetFile = new File(targetFilePath); + + // 音频格式转换 + audioFormatConversion(sourceFile, targetFile, format); + + return new FileInputStream(targetFile); + } + } catch (Exception e) { + log.error("音频格式转换异常:" + e.getMessage(), e); + } + return null; + } + + /** + * 音频格式转换 + * @param source 源音频文件 + * @param target 输出的音频文件 + * @param format wav/mp3/amr + */ + public static void audioFormatConversion(File source,File target,String format) { + try { + //Audio Attributes + AudioAttributes audio = new AudioAttributes(); + switch (format) { + case "wav": + audio.setCodec("pcm_s16le"); + break; + case "mp3": + audio.setCodec("libmp3lame"); + break; + case "amr": + audio.setCodec("libvo_amrwbenc"); + break; + default: + log.error("音频格式不合法!"); + return; + } + audio.setBitRate(16000); + audio.setChannels(1); + audio.setSamplingRate(16000); + //Encoding attributes + EncodingAttributes attrs = new EncodingAttributes(); + attrs.setOutputFormat(format); + attrs.setAudioAttributes(audio); + //Encode + Encoder encoder = new Encoder(); + encoder.encode(new MultimediaObject(source), target, attrs); + } catch (Exception e) { + log.error("音频格式转换异常:" + e.getMessage(), e); + } + } + + +} diff --git a/smartjavaai-translate/pom.xml b/translate/pom.xml similarity index 93% rename from smartjavaai-translate/pom.xml rename to translate/pom.xml index 5c70904..7e16e26 100644 --- a/smartjavaai-translate/pom.xml +++ b/translate/pom.xml @@ -9,19 +9,24 @@ 1.0.15 - smartjavaai-translate + translate cn.smartjavaai - smartjavaai-common + common ${project.version} + + + ai.djl.sentencepiece + sentencepiece + - 1.0.23 - smartjavaai-translate + 1.0.24 + translate SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java b/translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java rename to translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java diff --git a/translate/src/main/java/cn/smartjavaai/translation/config/OpusSearchConfig.java b/translate/src/main/java/cn/smartjavaai/translation/config/OpusSearchConfig.java new file mode 100644 index 0000000..dc2a881 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/config/OpusSearchConfig.java @@ -0,0 +1,63 @@ +package cn.smartjavaai.translation.config; +/** + * 配置信息 + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class OpusSearchConfig { + private int maxSeqLength; + private long padTokenId; + private long eosTokenId; + private int beam; + private boolean suffixPadding; + + public OpusSearchConfig() { + this.eosTokenId = 0; + this.padTokenId = 65000; + this.maxSeqLength = 512; + this.beam = 6; + } + + + public void setEosTokenId(long eosTokenId) { + this.eosTokenId = eosTokenId; + } + + public int getMaxSeqLength() { + return maxSeqLength; + } + + public void setMaxSeqLength(int maxSeqLength) { + this.maxSeqLength = maxSeqLength; + } + + public long getPadTokenId() { + return padTokenId; + } + + public void setPadTokenId(long padTokenId) { + this.padTokenId = padTokenId; + } + + public long getEosTokenId() { + return eosTokenId; + } + + public int getBeam() { + return beam; + } + + public void setBeam(int beam) { + this.beam = beam; + } + + public boolean isSuffixPadding() { + return suffixPadding; + } + + public void setSuffixPadding(boolean suffixPadding) { + this.suffixPadding = suffixPadding; + } +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java b/translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java rename to translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/BatchTensorList.java b/translate/src/main/java/cn/smartjavaai/translation/entity/BatchTensorList.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/BatchTensorList.java rename to translate/src/main/java/cn/smartjavaai/translation/entity/BatchTensorList.java diff --git a/translate/src/main/java/cn/smartjavaai/translation/entity/BeamBatchTensorList.java b/translate/src/main/java/cn/smartjavaai/translation/entity/BeamBatchTensorList.java new file mode 100644 index 0000000..2c40021 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/entity/BeamBatchTensorList.java @@ -0,0 +1,61 @@ +package cn.smartjavaai.translation.entity; + +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; + +/** + * beam 搜索张量对象列表 + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class BeamBatchTensorList { + private NDArray nextInputIds; + private NDArray encoderHiddenStates; + private NDArray attentionMask; + private NDList pastKeyValues; + + + public BeamBatchTensorList() { + } + + public BeamBatchTensorList(NDArray nextInputIds, NDArray attentionMask, NDArray encoderHiddenStates, NDList pastKeyValues) { + this.nextInputIds = nextInputIds; + this.attentionMask = attentionMask; + this.pastKeyValues = pastKeyValues; + this.encoderHiddenStates = encoderHiddenStates; + } + + public NDArray getNextInputIds() { + return nextInputIds; + } + + public void setNextInputIds(NDArray nextInputIds) { + this.nextInputIds = nextInputIds; + } + + public NDArray getEncoderHiddenStates() { + return encoderHiddenStates; + } + + public void setEncoderHiddenStates(NDArray encoderHiddenStates) { + this.encoderHiddenStates = encoderHiddenStates; + } + + public NDArray getAttentionMask() { + return attentionMask; + } + + public void setAttentionMask(NDArray attentionMask) { + this.attentionMask = attentionMask; + } + + public NDList getPastKeyValues() { + return pastKeyValues; + } + + public void setPastKeyValues(NDList pastKeyValues) { + this.pastKeyValues = pastKeyValues; + } +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/GreedyBatchTensorList.java b/translate/src/main/java/cn/smartjavaai/translation/entity/GreedyBatchTensorList.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/GreedyBatchTensorList.java rename to translate/src/main/java/cn/smartjavaai/translation/entity/GreedyBatchTensorList.java diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java b/translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java similarity index 76% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java rename to translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java index 6d426c9..aac5b41 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java +++ b/translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java @@ -45,6 +45,16 @@ public class TranslateParam { return R.ok(null); } + public TranslateParam(String input, LanguageCode sourceLanguage, LanguageCode targetLanguage) { + this.input = input; + this.sourceLanguage = sourceLanguage; + this.targetLanguage = targetLanguage; + } + public TranslateParam(String input) { + this.input = input; + } + public TranslateParam() { + } } diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java b/translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java rename to translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java b/translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java similarity index 91% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java rename to translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java index e558a48..c23b0b3 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java +++ b/translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java @@ -6,7 +6,11 @@ package cn.smartjavaai.translation.enums; */ public enum TranslationModeEnum { - NLLB_MODEL; + NLLB_MODEL, + + OPUS_MT_ZH_EN, + + OPUS_MT_EN_ZH; /** * 根据名称获取枚举 (忽略大小写和下划线变体) diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/exception/TranslationException.java b/translate/src/main/java/cn/smartjavaai/translation/exception/TranslationException.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/exception/TranslationException.java rename to translate/src/main/java/cn/smartjavaai/translation/exception/TranslationException.java diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java b/translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java similarity index 72% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java rename to translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java index 5ca6797..e13ad9e 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java +++ b/translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java @@ -4,8 +4,10 @@ import cn.smartjavaai.common.config.Config; import cn.smartjavaai.translation.config.TranslationModelConfig; +import cn.smartjavaai.translation.enums.TranslationModeEnum; import cn.smartjavaai.translation.exception.TranslationException; import cn.smartjavaai.translation.model.NllbModel; +import cn.smartjavaai.translation.model.OpusMtModel; import cn.smartjavaai.translation.model.TranslationModel; import lombok.extern.slf4j.Slf4j; @@ -23,14 +25,14 @@ public class TranslationModelFactory { // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 private static volatile TranslationModelFactory instance; - private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); + private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); /** * 检测模型注册表 */ - private static final Map> modelRegistry = + private static final Map> modelRegistry = new ConcurrentHashMap<>(); @@ -49,11 +51,11 @@ public class TranslationModelFactory { /** * 注册翻译模型 - * @param name + * @param translationModeEnum * @param clazz */ - private static void registerCommonDetModel(String name, Class clazz) { - modelRegistry.put(name.toLowerCase(), clazz); + private static void registerCommonDetModel(TranslationModeEnum translationModeEnum, Class clazz) { + modelRegistry.put(translationModeEnum, clazz); } /** @@ -65,7 +67,7 @@ public class TranslationModelFactory { if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){ throw new TranslationException("未配置OCR模型"); } - return modelMap.computeIfAbsent(config.getModelEnum().name(), k -> { + return modelMap.computeIfAbsent(config.getModelEnum(), k -> { return createModel(config); }); } @@ -77,7 +79,7 @@ public class TranslationModelFactory { * @return */ private TranslationModel createModel(TranslationModelConfig config) { - Class clazz = modelRegistry.get(config.getModelEnum().name().toLowerCase()); + Class clazz = modelRegistry.get(config.getModelEnum()); if(clazz == null){ throw new TranslationException("Unsupported model"); } @@ -94,7 +96,9 @@ public class TranslationModelFactory { // 初始化默认算法 static { - registerCommonDetModel("NLLB_MODEL", NllbModel.class); + registerCommonDetModel(TranslationModeEnum.NLLB_MODEL, NllbModel.class); + registerCommonDetModel(TranslationModeEnum.OPUS_MT_EN_ZH, OpusMtModel.class); + registerCommonDetModel(TranslationModeEnum.OPUS_MT_ZH_EN, OpusMtModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/BeamHypotheses.java b/translate/src/main/java/cn/smartjavaai/translation/model/BeamHypotheses.java new file mode 100644 index 0000000..6ae5ce4 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/model/BeamHypotheses.java @@ -0,0 +1,121 @@ +package cn.smartjavaai.translation.model; + +import ai.djl.util.Pair; + +import java.util.ArrayList; + +/** + * Beam hypothesis + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class BeamHypotheses { + float length_penalty; + boolean early_stopping; + int num_beams; + ArrayList> beams; + float worst_score = 1e9f; + + public BeamHypotheses(float length_penalty, boolean early_stopping, int num_beams) { + this.length_penalty = length_penalty; + this.early_stopping = early_stopping; + this.num_beams = num_beams; + beams = new ArrayList<>(); + } + + /** + * Get length + * + * @return + */ + public int getLen() { + return beams.size(); + } + + /** + * Add a new hypothesis to the list. + * + * @param sum_logprobs + * @param hyp + */ + public void add(float sum_logprobs, long[] hyp) { + float score = sum_logprobs / (float) (Math.pow(hyp.length, this.length_penalty)); + + if (getLen() < this.num_beams || score > this.worst_score) { + this.beams.add(new Pair<>(score, hyp)); + if (getLen() > this.num_beams) { + int index = min(); + this.beams.remove(index); + index = min(); + this.worst_score = this.beams.get(index).getKey(); + }else { + this.worst_score = Math.min(score, this.worst_score); + } + } + } + + /** + * Get Pair + * @param index + * @return + */ + public Pair getPair(int index) { + return beams.get(index); + } + + /** + * Get index for minmum score value + * + * @return + */ + public int min() { + float min = beams.get(0).getKey(); + int index = 0; + for (int i = 1; i < beams.size(); ++i) { + if (beams.get(i).getKey() < min) { + min = beams.get(i).getKey(); + index = i; + } + } + return index; + } + + /** + * Get index for maximum score value + * @return + */ + public int max() { + float max = beams.get(0).getKey(); + int index = 0; + for (int i = 1; i < beams.size(); ++i) { + if (beams.get(i).getKey() > max) { + max = beams.get(i).getKey(); + index = i; + } + } + return index; + } + + /** + * If there are enough hypotheses and that none of the hypotheses being generated can become better than the worst + * one in the heap, then we are done with this sentence. + * + * @param best_sum_logprobs + * @param cur_len + * @return + */ + public boolean isDone(float best_sum_logprobs, long cur_len) { + if (getLen() < this.num_beams) + return false; + + if (this.early_stopping) + return true; + else { + float highest_attainable_score = best_sum_logprobs / (float) Math.pow(cur_len, this.length_penalty); + boolean ret = (this.worst_score >= highest_attainable_score); + return ret; + } + } +} diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/BeamSearchScorer.java b/translate/src/main/java/cn/smartjavaai/translation/model/BeamSearchScorer.java new file mode 100644 index 0000000..48ed946 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/model/BeamSearchScorer.java @@ -0,0 +1,118 @@ +package cn.smartjavaai.translation.model; + +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.index.NDIndex; +import ai.djl.ndarray.types.Shape; +import ai.djl.util.Pair; + +/** + * Implementing standard beam search decoding. + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class BeamSearchScorer { + private int num_beams; + private float length_penalty; + private boolean do_early_stopping; + private int num_beam_hyps_to_keep; + private int num_beam_groups; + private BeamHypotheses beam_hyp; + private boolean _done; + + public BeamSearchScorer(int num_beams, float length_penalty, boolean do_early_stopping, int num_beam_hyps_to_keep, int num_beam_groups) { + this.num_beams = num_beams; + this.length_penalty = length_penalty; + this.do_early_stopping = do_early_stopping; + this.num_beam_hyps_to_keep = num_beam_hyps_to_keep; + this.num_beam_groups = num_beam_groups; + beam_hyp = new BeamHypotheses(length_penalty, do_early_stopping, num_beams); + _done = false; + } + + public boolean isDone() { + return _done; + } + + public NDList process(NDManager manager, NDArray input_ids, NDArray next_scores, NDArray next_tokens, NDArray next_indices, long pad_token_id, long eos_token_id) { + + float[] next_scores_arr = next_scores.toFloatArray(); + long[] next_indices_arr = next_indices.toLongArray(); + long[] next_tokens_arr = next_tokens.toLongArray(); + + NDArray next_beam_scores = manager.zeros(new Shape(1, this.num_beams), next_scores.getDataType()); + NDArray next_beam_tokens = manager.zeros(new Shape(1, this.num_beams), next_tokens.getDataType()); + NDArray next_beam_indices = manager.zeros(new Shape(1, this.num_beams), next_indices.getDataType()); + + + // next tokens for this sentence + int beam_idx = 0; + float maxScore = Float.NEGATIVE_INFINITY; + for (int i = 0; i < next_scores_arr.length; ++i) { + int beam_token_rank = i; + long next_token = next_tokens_arr[i]; + float next_score = next_scores_arr[i]; + if (maxScore < next_score) { + maxScore = next_score; + } + long next_index = next_indices_arr[i]; + + long batch_beam_idx = next_index; + + // add to generated hypotheses if end of sentence + if (next_token == eos_token_id) { + // if beam_token does not belong to top num_beams tokens, it should not be added + if (beam_token_rank >= this.num_beams) + continue; + long[] arr = input_ids.get(batch_beam_idx).toLongArray(); + // Add a new hypothesis to the list. + beam_hyp.add(next_score, arr); + } else { + // add next predicted token since it is not eos_token + next_beam_scores.set(new NDIndex(0, beam_idx), next_score); + next_beam_tokens.set(new NDIndex(0, beam_idx), next_token); + next_beam_indices.set(new NDIndex(0, beam_idx), batch_beam_idx); + beam_idx += 1; + } + + // once the beam for next step is full, don't add more tokens to it. + if (beam_idx == this.num_beams) + break; + } + + long cur_len = input_ids.getShape().getLastDimension(); + this._done = this._done || beam_hyp.isDone(maxScore, cur_len); + + NDList list = new NDList(); + list.add(next_beam_scores); + list.add(next_beam_tokens); + list.add(next_beam_indices); + + return list; + } + + public long[] finalize(int max_length, long eos_token_id) { + + // best_hyp_tuple + Pair pair = beam_hyp.getPair(beam_hyp.max()); + float best_score = pair.getKey(); + long[] best_hyp = pair.getValue(); + int sent_length = best_hyp.length; + + // prepare for adding eos + int sent_max_len = Math.min(sent_length + 1, max_length); + long[] decodedArr = new long[sent_max_len]; + + for (int i = 0; i < sent_length; ++i) { + decodedArr[i] = best_hyp[i]; + } + if (sent_length < max_length) { + decodedArr[sent_length] = eos_token_id; + } + + return decodedArr; + } +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java b/translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java similarity index 99% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java rename to translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java index 71653f7..1b9f22b 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java @@ -6,6 +6,7 @@ import ai.djl.engine.Engine; import ai.djl.huggingface.tokenizers.Encoding; import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; import ai.djl.inference.Predictor; +import ai.djl.modality.nlp.generate.CausalLMOutput; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; import ai.djl.ndarray.NDManager; @@ -20,7 +21,6 @@ import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.pool.CommonPredictorFactory; import cn.smartjavaai.translation.config.TranslationModelConfig; import cn.smartjavaai.translation.config.NllbSearchConfig; -import cn.smartjavaai.translation.entity.CausalLMOutput; import cn.smartjavaai.translation.entity.GreedyBatchTensorList; import cn.smartjavaai.translation.entity.TranslateParam; import cn.smartjavaai.translation.exception.TranslationException; @@ -40,7 +40,7 @@ import java.nio.file.Paths; import java.util.Objects; /** - * 机器翻译通用检测模型 + * Nllb机器翻译模型 * * @author lwx * @date 2025/6/05 diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java b/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java new file mode 100644 index 0000000..60f6608 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java @@ -0,0 +1,442 @@ +package cn.smartjavaai.translation.model; + +import ai.djl.Device; +import ai.djl.MalformedModelException; +import ai.djl.engine.Engine; +import ai.djl.huggingface.tokenizers.Encoding; +import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; +import ai.djl.inference.Predictor; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.nlp.generate.CausalLMOutput; +import ai.djl.modality.nlp.generate.SearchConfig; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.index.NDIndex; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +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.sentencepiece.SpTokenizer; +import ai.djl.translate.NoopTranslator; +import ai.djl.translate.TranslateException; +import ai.djl.util.Utils; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.pool.CommonPredictorFactory; +import cn.smartjavaai.translation.config.NllbSearchConfig; +import cn.smartjavaai.translation.config.OpusSearchConfig; +import cn.smartjavaai.translation.config.TranslationModelConfig; +import cn.smartjavaai.translation.entity.BeamBatchTensorList; +import cn.smartjavaai.translation.entity.GreedyBatchTensorList; +import cn.smartjavaai.translation.entity.TranslateParam; +import cn.smartjavaai.translation.exception.TranslationException; +import cn.smartjavaai.translation.model.translator.NllbDecoder2Translator; +import cn.smartjavaai.translation.model.translator.NllbDecoderTranslator; +import cn.smartjavaai.translation.model.translator.NllbEncoderTranslator; +import cn.smartjavaai.translation.model.translator.opus.Decoder2Translator; +import cn.smartjavaai.translation.model.translator.opus.DecoderTranslator; +import cn.smartjavaai.translation.model.translator.opus.EncoderTranslator; +import cn.smartjavaai.translation.utils.NDArrayUtils; +import cn.smartjavaai.translation.utils.TokenUtils; +import com.google.gson.Gson; +import com.google.gson.reflect.TypeToken; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.lang3.StringUtils; +import org.apache.commons.pool2.impl.GenericObjectPool; + +import java.io.IOException; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.*; +import java.util.concurrent.ConcurrentHashMap; + +/** + * OpusMt机器翻译模型 + * + * @author dwj + */ +@Slf4j +public class OpusMtModel implements TranslationModel{ + + private GenericObjectPool> encodePredictorPool; + + private GenericObjectPool> decodePredictorPool; + + private GenericObjectPool> decode2PredictorPool; + + private ZooModel model; + private SpTokenizer sourceTokenizer; + + private OpusSearchConfig searchConfig; + private TranslationModelConfig config; + + private ConcurrentHashMap map; + + private ConcurrentHashMap reverseMap; + + private float length_penalty = 1.0f; + private boolean do_early_stopping = false; + private int num_beam_hyps_to_keep = 1; + private int num_beam_groups = 1; + + + + @Override + public void loadModel(TranslationModelConfig config) { + if (StringUtils.isBlank(config.getModelPath())) { + throw new TranslationException("modelPath is null"); + } + Device device = null; + if (!Objects.isNull(config.getDevice())) { + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + this.config = config; + Path modelPath = Paths.get(config.getModelPath()); + Criteria criteria = + Criteria.builder() + .setTypes(NDList.class, NDList.class) + .optModelPath(modelPath) + .optEngine("PyTorch") + .optDevice(device) + .optTranslator(new NoopTranslator()) + .build(); + try { + model = ModelZoo.loadModel(criteria); + encodePredictorPool = new GenericObjectPool<>(new CommonPredictorFactory(model,new EncoderTranslator())); + decodePredictorPool = new GenericObjectPool<>(new CommonPredictorFactory(model,new DecoderTranslator())); + decode2PredictorPool = new GenericObjectPool<>(new CommonPredictorFactory(model,new Decoder2Translator())); + + Path tokenizerPath = modelPath.getParent().resolve("source.spm"); + sourceTokenizer = new SpTokenizer(tokenizerPath); + List words = Utils.readLines(modelPath.getParent().resolve("vocab.txt")); + String jsonStr = ""; + for (String line : words) { + jsonStr = jsonStr + line; + } + map = new Gson().fromJson(jsonStr, new TypeToken>() { + }.getType()); + reverseMap = new ConcurrentHashMap<>(); + Iterator it = map.entrySet().iterator(); + while (it.hasNext()) { + Map.Entry next = (Map.Entry) it.next(); + reverseMap.put(next.getValue(), next.getKey()); + } + //初始化searchConfig + this.searchConfig = new OpusSearchConfig(); + int predictorPoolSize = config.getPredictorPoolSize(); + if(config.getPredictorPoolSize() <= 0){ + predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数 + } + encodePredictorPool.setMaxTotal(predictorPoolSize); + decodePredictorPool.setMaxTotal(predictorPoolSize); + decode2PredictorPool.setMaxTotal(predictorPoolSize); + log.debug("当前设备: " + model.getNDManager().getDevice()); + log.debug("当前引擎: " + Engine.getInstance().getEngineName()); + log.debug("模型推理器线程池最大数量: " + predictorPoolSize); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new TranslationException("模型加载失败", e); + } + } + + @Override + public R translate(TranslateParam translateParam) { + if(translateParam == null){ + return R.fail(R.Status.PARAM_ERROR); + } + //验证 + if (StringUtils.isBlank(translateParam.getInput())) { + return R.fail(R.Status.PARAM_ERROR.getCode(), "输入文本不能为空"); + } + return R.ok(translateLanguage(translateParam)); + } + + @Override + public R translate(String input) { + //验证 + if (StringUtils.isBlank(input)) { + return R.fail(R.Status.PARAM_ERROR.getCode(), "输入文本不能为空"); + } + return R.ok(translateLanguage(new TranslateParam(input))); + } + + private String translateLanguage(TranslateParam translateParam) { + try (NDManager manager = NDManager.newBaseManager()) { + long numBeam = searchConfig.getBeam(); + BeamSearchScorer beamSearchScorer = new BeamSearchScorer((int) numBeam, length_penalty, do_early_stopping, num_beam_hyps_to_keep, num_beam_groups); + // 1. Encode + List tokens = sourceTokenizer.tokenize(translateParam.getInput()); + String[] strs = tokens.toArray(new String[]{}); + log.info("Tokens: " + Arrays.toString(strs)); + int[] sourceIds = new int[tokens.size() + 1]; + sourceIds[tokens.size()] = 0; + for (int i = 0; i < tokens.size(); i++) { + sourceIds[i] = map.get(tokens.get(i)).intValue(); + } + NDArray encoder_hidden_states = encoder(sourceIds); + encoder_hidden_states = NDArrayUtils.expand(encoder_hidden_states, searchConfig.getBeam()); + + NDArray decoder_input_ids = manager.create(new long[]{65000}).reshape(1, 1); + decoder_input_ids = NDArrayUtils.expand(decoder_input_ids, numBeam); + + + long[] attentionMask = new long[sourceIds.length]; + Arrays.fill(attentionMask, 1); + NDArray attentionMaskArray = manager.create(attentionMask).expandDims(0); + NDArray new_attention_mask = NDArrayUtils.expand(attentionMaskArray, searchConfig.getBeam()); + NDList decoderInput = new NDList(decoder_input_ids, encoder_hidden_states, new_attention_mask); + + + // 2. Initial Decoder + CausalLMOutput modelOutput = decoder(decoderInput); + modelOutput.getLogits().attach(manager); + modelOutput.getPastKeyValuesList().attach(manager); + + NDArray beam_scores = manager.zeros(new Shape(1, numBeam), DataType.FLOAT32); + beam_scores.set(new NDIndex(":, 1:"), -1e9); + beam_scores = beam_scores.reshape(numBeam, 1); + + NDArray input_ids = decoder_input_ids; + BeamBatchTensorList searchState = new BeamBatchTensorList(null, new_attention_mask, encoder_hidden_states, modelOutput.getPastKeyValuesList()); + NDArray next_tokens; + NDArray next_indices; + while (true) { + if (searchState.getNextInputIds() != null) { + decoder_input_ids = searchState.getNextInputIds().get(new NDIndex(":, -1:")); + decoderInput = new NDList(decoder_input_ids, searchState.getEncoderHiddenStates(), searchState.getAttentionMask()); + decoderInput.addAll(searchState.getPastKeyValues()); + // 3. Decoder loop + modelOutput = decoder2(decoderInput); + } + + NDArray next_token_logits = modelOutput.getLogits().get(":, -1, :"); + + // hack: adjust tokens for Marian. For Marian we have to make sure that the `pad_token_id` + // cannot be generated both before and after the `nn.functional.log_softmax` operation. + NDArray new_next_token_logits = manager.create(next_token_logits.getShape(), next_token_logits.getDataType()); + next_token_logits.copyTo(new_next_token_logits); + new_next_token_logits.set(new NDIndex(":," + searchConfig.getPadTokenId()), Float.NEGATIVE_INFINITY); + + NDArray next_token_scores = new_next_token_logits.logSoftmax(1); + + // next_token_scores = logits_processor(input_ids, next_token_scores) + // 1. NoBadWordsLogitsProcessor + next_token_scores.set(new NDIndex(":," + searchConfig.getPadTokenId()), Float.NEGATIVE_INFINITY); + + // 2. MinLengthLogitsProcessor 没生效 + // 3. ForcedEOSTokenLogitsProcessor + long cur_len = input_ids.getShape().getLastDimension(); + if (cur_len == (searchConfig.getMaxSeqLength() - 1)) { + long num_tokens = next_token_scores.getShape().getLastDimension(); + for (long i = 0; i < num_tokens; i++) { + if(i != searchConfig.getEosTokenId()){ + next_token_scores.set(new NDIndex(":," + i), Float.NEGATIVE_INFINITY); + } + } + next_token_scores.set(new NDIndex(":," + searchConfig.getEosTokenId()), 0); + } + + long vocab_size = next_token_scores.getShape().getLastDimension(); + beam_scores = beam_scores.repeat(1, vocab_size); + next_token_scores = next_token_scores.add(beam_scores); + + // reshape for beam search + next_token_scores = next_token_scores.reshape(1, numBeam * vocab_size); + + // [batch, beam] + NDList topK = next_token_scores.topK(Math.toIntExact(numBeam) * 2, 1, true, true); + + next_token_scores = topK.get(0); + next_tokens = topK.get(1); + + // next_indices = next_tokens // vocab_size + next_indices = next_tokens.div(vocab_size).toType(DataType.INT64, true); + + // next_tokens = next_tokens % vocab_size + next_tokens = next_tokens.mod(vocab_size); + + // stateless + NDList beam_outputs = beamSearchScorer.process(manager, input_ids, next_token_scores, next_tokens, next_indices, searchConfig.getPadTokenId(), searchConfig.getEosTokenId()); + + beam_scores = beam_outputs.get(0).reshape(numBeam, 1); + NDArray beam_next_tokens = beam_outputs.get(1); + NDArray beam_idx = beam_outputs.get(2); + + // input_ids = torch.cat([input_ids[beam_idx, :], beam_next_tokens.unsqueeze(-1)], dim=-1) + long[] beam_next_tokens_arr = beam_next_tokens.toLongArray(); + long[] beam_idx_arr = beam_idx.toLongArray(); + NDList inputList = new NDList(); + for (int i = 0; i < numBeam; i++) { + long index = beam_idx_arr[i]; + NDArray ndArray = input_ids.get(index).reshape(1, input_ids.getShape().getLastDimension()); + ndArray = ndArray.concat(manager.create(beam_next_tokens_arr[i]).reshape(1, 1), 1); + inputList.add(ndArray); + } + input_ids = NDArrays.concat(inputList, 0); + searchState.setNextInputIds(input_ids); + searchState.setPastKeyValues(modelOutput.getPastKeyValuesList()); + + boolean maxLengthCriteria = (input_ids.getShape().getLastDimension() >= searchConfig.getMaxSeqLength()); + if (beamSearchScorer.isDone() || maxLengthCriteria) { + break; + } + + } + + long[] sequences = beamSearchScorer.finalize(searchConfig.getMaxSeqLength(), searchConfig.getEosTokenId()); + String result = TokenUtils.decode(reverseMap, sequences); + return result; + } catch (Exception e) { + throw new TranslationException("翻译错误", e); + } + } + + public NDArray encoder(int[] ids) { + Predictor predictor = null; + try { + predictor = (Predictor)encodePredictorPool.borrowObject(); + return predictor.predict(ids); + } catch (Exception e) { + throw new TranslationException("机器翻译编码错误", e); + }finally { + if (predictor != null) { + try { + encodePredictorPool.returnObject(predictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + public CausalLMOutput decoder(NDList input) throws TranslateException { + Predictor predictor = null; + try { + predictor = (Predictor)decodePredictorPool.borrowObject(); + return predictor.predict(input); + } catch (Exception e) { + throw new TranslationException("机器翻译编码错误", e); + }finally { + if (predictor != null) { + try { + decodePredictorPool.returnObject(predictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + public CausalLMOutput decoder2(NDList input) throws TranslateException { + Predictor predictor = null; + try { + predictor = (Predictor)decode2PredictorPool.borrowObject(); + return predictor.predict(input); + } catch (Exception e) { + throw new TranslationException("机器翻译编码错误", e); + }finally { + if (predictor != null) { + try { + decode2PredictorPool.returnObject(predictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + public NDArray greedyStepGen(NllbSearchConfig config, NDArray pastOutputIds, NDArray next_token_scores, NDManager manager) { + next_token_scores = next_token_scores.get(":, -1, :"); + + NDArray new_next_token_scores = manager.create(next_token_scores.getShape(), next_token_scores.getDataType()); + next_token_scores.copyTo(new_next_token_scores); + + // LogitsProcessor 1. ForcedBOSTokenLogitsProcessor + // 设置目标语言 + long cur_len = pastOutputIds.getShape().getLastDimension(); + if (cur_len == 1) { + long num_tokens = new_next_token_scores.getShape().getLastDimension(); + for (long i = 0; i < num_tokens; i++) { + if (i != config.getForcedBosTokenId()) { + new_next_token_scores.set(new NDIndex(":," + i), Float.NEGATIVE_INFINITY); + } + } + new_next_token_scores.set(new NDIndex(":," + config.getForcedBosTokenId()), 0); + } + + NDArray probs = new_next_token_scores.softmax(-1); + NDArray next_tokens = probs.argMax(-1); + + return next_tokens.expandDims(0); + } + + public GenericObjectPool> getEncodePredictorPool() { + return encodePredictorPool; + } + + public GenericObjectPool> getDecodePredictorPool() { + return decodePredictorPool; + } + + public GenericObjectPool> getDecode2PredictorPool() { + return decode2PredictorPool; + } + + @Override + public void close() throws Exception { + try { + if (model != null) { + model.close(); + } + } catch (Exception e) { + log.warn("关闭 model 失败", e); + } + try { + if (sourceTokenizer != null) { + sourceTokenizer.close(); + } + } catch (Exception e) { + log.warn("关闭 tokenizer 失败", e); + } + try { + if (encodePredictorPool != null) { + encodePredictorPool.close(); + } + } catch (Exception e) { + log.warn("关闭 encodePredictorPool 失败", e); + } + try { + if (decodePredictorPool != null) { + decodePredictorPool.close(); + } + } catch (Exception e) { + log.warn("关闭 decodePredictorPool 失败", e); + } + try { + if (decode2PredictorPool != null) { + decode2PredictorPool.close(); + } + } catch (Exception e) { + log.warn("关闭 decode2PredictorPool 失败", e); + } + } +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java b/translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java similarity index 76% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java rename to translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java index f0b35d2..a9d55ec 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java @@ -28,6 +28,15 @@ public interface TranslationModel extends AutoCloseable{ } + /** + * 机器翻译 + * @param input 输入文本 + * @return + */ + default R translate(String input) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java b/translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java similarity index 95% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java rename to translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java index 6e18541..f98e4aa 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java @@ -1,10 +1,10 @@ package cn.smartjavaai.translation.model.translator; +import ai.djl.modality.nlp.generate.CausalLMOutput; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; import ai.djl.translate.NoBatchifyTranslator; import ai.djl.translate.TranslatorContext; -import cn.smartjavaai.translation.entity.CausalLMOutput; /** diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoderTranslator.java b/translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoderTranslator.java similarity index 95% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoderTranslator.java rename to translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoderTranslator.java index 7867c87..717b49c 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoderTranslator.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoderTranslator.java @@ -1,10 +1,10 @@ package cn.smartjavaai.translation.model.translator; +import ai.djl.modality.nlp.generate.CausalLMOutput; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; import ai.djl.translate.NoBatchifyTranslator; import ai.djl.translate.TranslatorContext; -import cn.smartjavaai.translation.entity.CausalLMOutput; /** * 解碼器,參數沒有 pastKeyValues diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbEncoderTranslator.java b/translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbEncoderTranslator.java similarity index 100% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbEncoderTranslator.java rename to translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbEncoderTranslator.java diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/Decoder2Translator.java b/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/Decoder2Translator.java new file mode 100644 index 0000000..434361d --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/Decoder2Translator.java @@ -0,0 +1,52 @@ +package cn.smartjavaai.translation.model.translator.opus; + +import ai.djl.modality.nlp.generate.CausalLMOutput; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.translate.NoBatchifyTranslator; +import ai.djl.translate.TranslatorContext; +/** + * 解碼器,參數支持 pastKeyValues + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class Decoder2Translator implements NoBatchifyTranslator { + private String tupleName; + + public Decoder2Translator() { + tupleName = "past_key_values(" + 6 + ',' + 4 + ')'; + } + + @Override + public NDList processInput(TranslatorContext ctx, NDList input) { + + NDArray placeholder = ctx.getNDManager().create(0); + placeholder.setName("module_method:decoder2"); + + input.add(placeholder); + + return input; + } + + @Override + public CausalLMOutput processOutput(TranslatorContext ctx, NDList output) { + NDArray logitsOutput = output.get(0); + NDList pastKeyValuesOutput = output.subNDList(1, 6 * 4 + 1); +// if (ctx.getAttachment("initialCall") != null) { +// NDIndex index2 = new NDIndex(":, :, 1:, ..."); +// pastKeyValuesOutput = +// new NDList( +// pastKeyValuesOutput.stream() +// .map(object -> object.get(index2)) +// .collect(Collectors.toList())); +// } + + for (NDArray array : pastKeyValuesOutput) { + array.setName(tupleName); + } + + return new CausalLMOutput(logitsOutput, pastKeyValuesOutput); + } +} diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/DecoderTranslator.java b/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/DecoderTranslator.java new file mode 100644 index 0000000..e40f305 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/DecoderTranslator.java @@ -0,0 +1,54 @@ +package cn.smartjavaai.translation.model.translator.opus; + +import ai.djl.modality.nlp.generate.CausalLMOutput; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.translate.NoBatchifyTranslator; +import ai.djl.translate.TranslatorContext; +/** + * 解碼器,參數沒有 pastKeyValues + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class DecoderTranslator implements NoBatchifyTranslator { + private String tupleName; + + public DecoderTranslator() { + tupleName = "past_key_values(" + 6 + ',' + 4 + ')'; + } + + @Override + public NDList processInput(TranslatorContext ctx, NDList input) { + + NDArray placeholder = ctx.getNDManager().create(0); + placeholder.setName("module_method:decoder"); + + input.add(placeholder); + + return input; + } + + @Override + public CausalLMOutput processOutput(TranslatorContext ctx, NDList output) { + NDArray logitsOutput = output.get(0); + NDList pastKeyValuesOutput = output.subNDList(1, 6 * 4 + 1); + +// if (ctx.getAttachment("initialCall") != null) { +// NDIndex index2 = new NDIndex(":, :, 1:, ..."); +// pastKeyValuesOutput = +// new NDList( +// pastKeyValuesOutput.stream() +// .map(object -> object.get(index2)) +// .collect(Collectors.toList())); +// } + + for (NDArray array : pastKeyValuesOutput) { + array.setName(tupleName); + } + logitsOutput.detach(); + pastKeyValuesOutput.detach(); + return new CausalLMOutput(logitsOutput, pastKeyValuesOutput); + } +} diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/EncoderTranslator.java b/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/EncoderTranslator.java new file mode 100644 index 0000000..037e731 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/model/translator/opus/EncoderTranslator.java @@ -0,0 +1,49 @@ +package cn.smartjavaai.translation.model.translator.opus; + +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; +import ai.djl.translate.NoBatchifyTranslator; +import ai.djl.translate.TranslatorContext; + +import java.util.Arrays; + +/** + * 编码器前后处理 + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public class EncoderTranslator implements NoBatchifyTranslator { + + + public EncoderTranslator() { + } + + @Override + public NDList processInput(TranslatorContext ctx, int[] input) throws Exception { + NDManager manager = ctx.getNDManager(); + NDArray inputIdArray = manager.create(input).expandDims(0).toType(DataType.INT64, false); + inputIdArray.setName("input_ids"); + + long[] attentionMask = new long[input.length]; + Arrays.fill(attentionMask, 1); + NDArray attentionMaskArray = manager.create(attentionMask).expandDims(0); + attentionMaskArray.setName("attention_mask"); + + NDArray placeholder = ctx.getNDManager().create(0); + placeholder.setName("module_method:encoder"); + + return new NDList(inputIdArray, attentionMaskArray, placeholder); + } + + @Override + public NDArray processOutput(TranslatorContext ctx, NDList list) { + NDArray encoder_hidden_states = list.get(0); + encoder_hidden_states.detach(); + return encoder_hidden_states; + } + +} diff --git a/translate/src/main/java/cn/smartjavaai/translation/utils/NDArrayUtils.java b/translate/src/main/java/cn/smartjavaai/translation/utils/NDArrayUtils.java new file mode 100644 index 0000000..f159698 --- /dev/null +++ b/translate/src/main/java/cn/smartjavaai/translation/utils/NDArrayUtils.java @@ -0,0 +1,28 @@ +package cn.smartjavaai.translation.utils; + +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDArrays; +import ai.djl.ndarray.NDList; + +/** + * NDArray 工具类 + * + * @author Calvin + * @mail 179209347@qq.com + * @website www.aias.top + */ +public final class NDArrayUtils { + + private NDArrayUtils() { + } + + public static NDArray expand(NDArray array, long beam) { + NDList list = new NDList(); + for (long i = 0; i < beam; i++) { + list.add(array); + } + NDArray result = NDArrays.concat(list, 0); + + return result; + } +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java b/translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java similarity index 59% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java rename to translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java index d613432..f19dc4b 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java +++ b/translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java @@ -6,6 +6,8 @@ import cn.smartjavaai.translation.config.NllbSearchConfig; import java.util.ArrayList; +import java.util.Arrays; +import java.util.Map; /** * @@ -44,4 +46,30 @@ public final class TokenUtils { String text = tokenizer.decode(ids); return text; } + + /** + * Token 解码 + * 根据语言的类型更新下面的方法 + * + * @param reverseMap + * @param outputIds + * @return + */ + public static String decode(Map reverseMap, long[] outputIds) { + int[] intArray = Arrays.stream(outputIds).mapToInt(l -> (int) l).toArray(); + + StringBuffer sb = new StringBuffer(); + for (int value : intArray) { + // 65000 + // 0 + if (value == 65000 || value == 0 || value == 8) + continue; + String text = reverseMap.get(Long.valueOf(value)); + sb.append(text); + } + + String result = sb.toString(); + result = result.replaceAll("▁"," "); + return result; + } } diff --git a/translate/src/main/test/java/Test.java b/translate/src/main/test/java/Test.java new file mode 100644 index 0000000..d602703 --- /dev/null +++ b/translate/src/main/test/java/Test.java @@ -0,0 +1,35 @@ +import ai.djl.util.JsonUtils; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.translation.config.TranslationModelConfig; +import cn.smartjavaai.translation.entity.TranslateParam; +import cn.smartjavaai.translation.enums.LanguageCode; +import cn.smartjavaai.translation.enums.TranslationModeEnum; +import cn.smartjavaai.translation.factory.TranslationModelFactory; +import cn.smartjavaai.translation.model.TranslationModel; + +/** + * @author dwj + * @date 2025/6/17 + */ +public class Test { + + public static void main(String[] args) { + TranslationModelConfig config = new TranslationModelConfig(); + config.setModelEnum(TranslationModeEnum.NLLB_MODEL); + config.setModelPath("/Users/wenjie/Documents/develop/model/trans/traced_translation_cpu.pt"); + // 输入文字 + String input2 = "我爱你"; + String input = "你好,欢迎使用SmartJavaAI!"; + TranslationModel detModel = TranslationModelFactory.getInstance().getModel(config); + TranslateParam translateParam = new TranslateParam(); + translateParam.setInput(input2); +// translateParam.setSourceLanguage(LanguageCode.ZHO_HANS); +// translateParam.setTargetLanguage(LanguageCode.ENG_LATN); + translateParam.setSourceLanguage(LanguageCode.ZHO_HANS); + translateParam.setTargetLanguage(LanguageCode.KOR_HANG); + R result = detModel.translate(translateParam); + System.out.println(JsonUtils.toJson(result)); + + } + +} diff --git a/smartjavaai-objectdetection/pom.xml b/vision/pom.xml similarity index 81% rename from smartjavaai-objectdetection/pom.xml rename to vision/pom.xml index dbe79d0..8c906ab 100644 --- a/smartjavaai-objectdetection/pom.xml +++ b/vision/pom.xml @@ -6,12 +6,12 @@ cn.smartjavaai smartjavaai-parent - 1.0.23 + 1.0.24 - smartjavaai-objectdetection - 1.0.23 - smartjavaai-objectdetection + vision + 1.0.24 + vision SmartJavaAI https://github.com/geekwenjie/SmartJavaAI @@ -32,9 +32,36 @@ cn.smartjavaai - smartjavaai-common + common ${project.version} + + + org.bytedeco + javacpp + 1.5.10 + macosx-arm64 + + + org.bytedeco + ffmpeg + 6.1.1-1.5.10 + macosx-arm64 + + + + org.bytedeco + openblas + 0.3.26-1.5.10 + macosx-arm64 + + + + org.bytedeco + opencv + 4.9.0-1.5.10 + macosx-arm64 + diff --git a/vision/src/main/java/cn/smartjavaai/action/config/ActionRecModelConfig.java b/vision/src/main/java/cn/smartjavaai/action/config/ActionRecModelConfig.java new file mode 100644 index 0000000..fad2123 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/config/ActionRecModelConfig.java @@ -0,0 +1,58 @@ +package cn.smartjavaai.action.config; + +import cn.smartjavaai.action.enums.ActionRecModelEnum; +import cn.smartjavaai.common.config.ModelConfig; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.objectdetection.constant.DetectorConstant; +import cn.smartjavaai.objectdetection.enums.DetectorModelEnum; +import lombok.Data; + +import java.util.List; + +/** + * 人类动作识别模型参数配置 + * + * @author dwj + */ +@Data +public class ActionRecModelConfig extends ModelConfig { + + /** + * 模型 + */ + private ActionRecModelEnum modelEnum; + + + + /** + * 模型路径 + */ + private String modelPath; + + + /** + * 允许的分类列表 + */ + private List allowedClasses; + + /** + * 置信度阈值 + */ + private float threshold = 0.3f; + + + + public ActionRecModelConfig() { + } + + public ActionRecModelConfig(ActionRecModelEnum modelEnum, DeviceEnum device) { + this.modelEnum = modelEnum; + setDevice(device); + } + + public ActionRecModelConfig(ActionRecModelEnum modelEnum) { + this.modelEnum = modelEnum; + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/action/criteria/ActionRecCriteriaFactory.java b/vision/src/main/java/cn/smartjavaai/action/criteria/ActionRecCriteriaFactory.java new file mode 100644 index 0000000..c8eec8b --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/criteria/ActionRecCriteriaFactory.java @@ -0,0 +1,67 @@ +package cn.smartjavaai.action.criteria; + +import ai.djl.Device; +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.Image; +import ai.djl.repository.zoo.Criteria; +import ai.djl.training.util.ProgressBar; +import cn.smartjavaai.action.config.ActionRecModelConfig; +import cn.smartjavaai.action.enums.ActionRecModelEnum; +import cn.smartjavaai.action.model.CommonActionTranslator; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.action.exception.ActionException; +import org.apache.commons.lang3.StringUtils; + +import java.nio.file.Paths; +import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; + +/** + * 动作识别模型工厂 + * @author dwj + */ +public class ActionRecCriteriaFactory { + + + public static Criteria createCriteria(ActionRecModelConfig config) { + Device device = null; + if(!Objects.isNull(config.getDevice())){ + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + Criteria criteria = null; + ConcurrentHashMap params = new ConcurrentHashMap(); + params.putAll(config.getCustomParams()); + if(config.getModelEnum() == ActionRecModelEnum.VIT_BASE_PATCH16_224){ + criteria = + Criteria.builder() + .setTypes(Image.class, Classifications.class) + .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : + config.getModelEnum().getModelUri()) + .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) + .optEngine("PyTorch") + .optDevice(device) + .optProgress(new ProgressBar()) + .build(); + }else { + if (StringUtils.isBlank(config.getModelPath())){ + throw new ActionException("请指定模型路径"); + } + int width = 224; + int height = 224; + if(config.getModelEnum() == ActionRecModelEnum.INCEPTIONV3_KINETICS400){ + width = 299; + height = 299; + } + criteria = + Criteria.builder() + .setTypes(Image.class, Classifications.class) + .optTranslator(new CommonActionTranslator(width, height)) + .optEngine("OnnxRuntime") + .optModelPath(Paths.get(config.getModelPath())) + .optDevice(device) + .optProgress(new ProgressBar()) + .build(); + } + return criteria; + } +} diff --git a/vision/src/main/java/cn/smartjavaai/action/enums/ActionRecModelEnum.java b/vision/src/main/java/cn/smartjavaai/action/enums/ActionRecModelEnum.java new file mode 100644 index 0000000..dbb0605 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/enums/ActionRecModelEnum.java @@ -0,0 +1,48 @@ +package cn.smartjavaai.action.enums; + +/** + * 人类动作识别模型枚举 + * @author dwj + */ +public enum ActionRecModelEnum { + + VIT_BASE_PATCH16_224("djl://ai.djl.pytorch/Human-Action-Recognition-VIT-Base-patch16-224"), + + INCEPTIONV3_KINETICS400(""), + + INCEPTIONV1_KINETICS400(""), + + RESNET18_V1B_KINETICS400(""), + + RESNET34_V1B_KINETICS400(""), + + RESNET50_V1B_KINETICS400(""), + + RESNET101_V1B_KINETICS400(""), + + RESNET152_V1B_KINETICS400(""); + + /** + * 根据名称获取枚举 (忽略大小写和下划线变体) + */ + public static ActionRecModelEnum fromName(String name) { + String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); + for (ActionRecModelEnum model : values()) { + if (model.name().replaceAll("_", "").equals(formatted)) { + return model; + } + } + throw new IllegalArgumentException("未知模型名称: " + name); + } + + private final String modelUri; + + ActionRecModelEnum(String modelUri) { + this.modelUri = modelUri; + } + + public String getModelUri() { + return modelUri; + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/action/exception/ActionException.java b/vision/src/main/java/cn/smartjavaai/action/exception/ActionException.java new file mode 100644 index 0000000..0d01d68 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/exception/ActionException.java @@ -0,0 +1,29 @@ +package cn.smartjavaai.action.exception; + +/** + * 动作检测异常 + * @author dwj + */ +public class ActionException extends RuntimeException{ + + public ActionException() { + super(); + } + + public ActionException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) { + super(message, cause, enableSuppression, writableStackTrace); + } + + public ActionException(String message, Throwable cause) { + super(message, cause); + } + + public ActionException(String message) { + super(message); + } + + public ActionException(Throwable cause) { + super(cause); + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java b/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java new file mode 100644 index 0000000..b708ee0 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java @@ -0,0 +1,46 @@ +package cn.smartjavaai.action.model; + +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.Image; +import cn.smartjavaai.action.config.ActionRecModelConfig; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.objectdetection.config.DetectorModelConfig; + +import java.awt.image.BufferedImage; +import java.io.InputStream; + +/** + * 人类动作识别模型 + * @author dwj + * @date 2025/8/12 + */ +public interface ActionRecModel extends AutoCloseable{ + + + /** + * 加载模型 + * @param config + */ + void loadModel(ActionRecModelConfig config); + + /** + * 动作检测 + * @param base64Image + * @return + */ + default R detectBase64(String base64Image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 动作检测 + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java b/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java new file mode 100644 index 0000000..cd784c5 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java @@ -0,0 +1,142 @@ +package cn.smartjavaai.action.model; + +import ai.djl.MalformedModelException; +import ai.djl.engine.Engine; +import ai.djl.inference.Predictor; +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.ImageFactory; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ModelNotFoundException; +import ai.djl.repository.zoo.ZooModel; +import cn.hutool.core.img.ImgUtil; +import cn.smartjavaai.action.config.ActionRecModelConfig; +import cn.smartjavaai.action.criteria.ActionRecCriteriaFactory; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.Base64ImageUtils; +import cn.smartjavaai.objectdetection.criteria.CriteriaBuilderFactory; +import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.vision.utils.ClassificationFilter; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.collections.CollectionUtils; +import org.apache.commons.lang3.StringUtils; +import org.apache.commons.pool2.impl.GenericObjectPool; + +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.util.Collections; +import java.util.Objects; + +/** + * 通用人类动作识别模型 + * @author dwj + */ +@Slf4j +public class CommonActionRecModel implements ActionRecModel{ + + + private ActionRecModelConfig config; + + private ZooModel model; + + private GenericObjectPool> predictorPool; + + @Override + public void loadModel(ActionRecModelConfig config) { + if(Objects.isNull(config.getModelEnum())){ + throw new DetectionException("未配置模型枚举"); + } + Criteria criteria = ActionRecCriteriaFactory.createCriteria(config); + this.config = config; + try { + model = criteria.loadModel(); + // 创建池子:每个线程独享 Predictor + this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model)); + int predictorPoolSize = config.getPredictorPoolSize(); + if(config.getPredictorPoolSize() <= 0){ + predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数 + } + predictorPool.setMaxTotal(predictorPoolSize); + log.debug("当前设备: " + model.getNDManager().getDevice()); + log.debug("当前引擎: " + Engine.getInstance().getEngineName()); + log.debug("模型推理器线程池最大数量: " + predictorPoolSize); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new DetectionException("模型加载失败", e); + } + } + + @Override + public R detectBase64(String base64Image) { + if(StringUtils.isBlank(base64Image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + try { + byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); + Image image = ImageFactory.getInstance().fromInputStream(new ByteArrayInputStream(imageData)); + return detect(image); + } catch (IOException e) { + throw new DetectionException("读取图片异常", e); + } + } + + @Override + public R detect(Image image) { + Classifications classifications = detectCore(image); + // 过滤 + if(config.getThreshold() > 0 && CollectionUtils.isNotEmpty(config.getAllowedClasses()) + && Objects.nonNull(classifications) && !classifications.items().isEmpty()){ + classifications = new ClassificationFilter(config.getAllowedClasses(), config.getThreshold()).filter(classifications); + } + return R.ok(classifications); + } + + /** + * 模型核心推理方法 + * @param image + * @return + */ + public Classifications detectCore(Image image) { + Predictor predictor = null; + try { + predictor = predictorPool.borrowObject(); + return predictor.predict(image); + } catch (Exception e) { + throw new DetectionException("动作识别错误", e); + }finally { + if (predictor != null) { + try { + predictorPool.returnObject(predictor); //归还 + log.debug("释放资源"); + } 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); + } + } +} diff --git a/vision/src/main/java/cn/smartjavaai/action/model/CommonActionTranslator.java b/vision/src/main/java/cn/smartjavaai/action/model/CommonActionTranslator.java new file mode 100644 index 0000000..796bec1 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/action/model/CommonActionTranslator.java @@ -0,0 +1,140 @@ +package cn.smartjavaai.action.model; + +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.Image; +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.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.translate.Batchifier; +import ai.djl.translate.Translator; +import ai.djl.translate.TranslatorContext; +import cn.smartjavaai.common.utils.LetterBoxUtils; + +import java.util.Arrays; + +/** + * @author dwj + */ +public class CommonActionTranslator implements Translator { + + private int width = 224; + + private int height = 224; + + private String[] labels; + + public static String rawLabels = "'abseiling','air_drumming','answering_questions','applauding','applying_cream','archery'," + + "'arm_wrestling','arranging_flowers','assembling_computer','auctioning','baby_waking_up','baking_cookies'," + + "'balloon_blowing','bandaging','barbequing','bartending','beatboxing','bee_keeping','belly_dancing'," + + "'bench_pressing','bending_back','bending_metal','biking_through_snow','blasting_sand','blowing_glass'," + + "'blowing_leaves','blowing_nose','blowing_out_candles','bobsledding','bookbinding','bouncing_on_trampoline'," + + "'bowling','braiding_hair','breading_or_breadcrumbing','breakdancing','brush_painting','brushing_hair'," + + "'brushing_teeth','building_cabinet','building_shed','bungee_jumping','busking','canoeing_or_kayaking'," + + "'capoeira','carrying_baby','cartwheeling','carving_pumpkin','catching_fish','catching_or_throwing_baseball'," + + "'catching_or_throwing_frisbee','catching_or_throwing_softball','celebrating','changing_oil','changing_wheel'," + + "'checking_tires','cheerleading','chopping_wood','clapping','clay_pottery_making','clean_and_jerk'," + + "'cleaning_floor','cleaning_gutters','cleaning_pool','cleaning_shoes','cleaning_toilet','cleaning_windows'," + + "'climbing_a_rope','climbing_ladder','climbing_tree','contact_juggling','cooking_chicken','cooking_egg'," + + "'cooking_on_campfire','cooking_sausages','counting_money','country_line_dancing','cracking_neck','crawling_baby'," + + "'crossing_river','crying','curling_hair','cutting_nails','cutting_pineapple','cutting_watermelon'," + + "'dancing_ballet','dancing_charleston','dancing_gangnam_style','dancing_macarena','deadlifting'," + + "'decorating_the_christmas_tree','digging','dining','disc_golfing','diving_cliff','dodgeball','doing_aerobics'," + + "'doing_laundry','doing_nails','drawing','dribbling_basketball','drinking','drinking_beer','drinking_shots'," + + "'driving_car','driving_tractor','drop_kicking','drumming_fingers','dunking_basketball','dying_hair'," + + "'eating_burger','eating_cake','eating_carrots','eating_chips','eating_doughnuts','eating_hotdog'," + + "'eating_ice_cream','eating_spaghetti','eating_watermelon','egg_hunting','exercising_arm'," + + "'exercising_with_an_exercise_ball','extinguishing_fire','faceplanting','feeding_birds','feeding_fish'," + + "'feeding_goats','filling_eyebrows','finger_snapping','fixing_hair','flipping_pancake','flying_kite'," + + "'folding_clothes','folding_napkins','folding_paper','front_raises','frying_vegetables','garbage_collecting'," + + "'gargling','getting_a_haircut','getting_a_tattoo','giving_or_receiving_award','golf_chipping','golf_driving'," + + "'golf_putting','grinding_meat','grooming_dog','grooming_horse','gymnastics_tumbling','hammer_throw'," + + "'headbanging','headbutting','high_jump','high_kick','hitting_baseball','hockey_stop','holding_snake'," + + "'hopscotch','hoverboarding','hugging','hula_hooping','hurdling','hurling_-sport-','ice_climbing','ice_fishing'," + + "'ice_skating','ironing','javelin_throw','jetskiing','jogging','juggling_balls','juggling_fire'," + + "'juggling_soccer_ball','jumping_into_pool','jumpstyle_dancing','kicking_field_goal','kicking_soccer_ball'," + + "'kissing','kitesurfing','knitting','krumping','laughing','laying_bricks','long_jump','lunge','making_a_cake'," + + "'making_a_sandwich','making_bed','making_jewelry','making_pizza','making_snowman','making_sushi','making_tea'," + + "'marching','massaging_back','massaging_feet','massaging_legs','massaging_person's_head','milking_cow'," + + "'mopping_floor','motorcycling','moving_furniture','mowing_lawn','news_anchoring','opening_bottle'," + + "'opening_present','paragliding','parasailing','parkour','passing_American_football_-in_game-'," + + "'passing_American_football_-not_in_game-','peeling_apples','peeling_potatoes','petting_animal_-not_cat-'," + + "'petting_cat','picking_fruit','planting_trees','plastering','playing_accordion','playing_badminton'," + + "'playing_bagpipes','playing_basketball','playing_bass_guitar','playing_cards','playing_cello','playing_chess'," + + "'playing_clarinet','playing_controller','playing_cricket','playing_cymbals','playing_didgeridoo','playing_drums'," + + "'playing_flute','playing_guitar','playing_harmonica','playing_harp','playing_ice_hockey','playing_keyboard'," + + "'playing_kickball','playing_monopoly','playing_organ','playing_paintball','playing_piano','playing_poker'," + + "'playing_recorder','playing_saxophone','playing_squash_or_racquetball','playing_tennis','playing_trombone'," + + "'playing_trumpet','playing_ukulele','playing_violin','playing_volleyball','playing_xylophone','pole_vault'," + + "'presenting_weather_forecast','pull_ups','pumping_fist','pumping_gas','punching_bag','punching_person_-boxing-'," + + "'push_up','pushing_car','pushing_cart','pushing_wheelchair','reading_book','reading_newspaper','recording_music'," + + "'riding_a_bike','riding_camel','riding_elephant','riding_mechanical_bull','riding_mountain_bike','riding_mule'," + + "'riding_or_walking_with_horse','riding_scooter','riding_unicycle','ripping_paper','robot_dancing','rock_climbing'," + + "'rock_scissors_paper','roller_skating','running_on_treadmill','sailing','salsa_dancing','sanding_floor'," + + "'scrambling_eggs','scuba_diving','setting_table','shaking_hands','shaking_head','sharpening_knives'," + + "'sharpening_pencil','shaving_head','shaving_legs','shearing_sheep','shining_shoes','shooting_basketball'," + + "'shooting_goal_-soccer-','shot_put','shoveling_snow','shredding_paper','shuffling_cards','side_kick'," + + "'sign_language_interpreting','singing','situp','skateboarding','ski_jumping','skiing_-not_slalom_or_crosscountry-'," + + "'skiing_crosscountry','skiing_slalom','skipping_rope','skydiving','slacklining','slapping','sled_dog_racing'," + + "'smoking','smoking_hookah','snatch_weight_lifting','sneezing','sniffing','snorkeling','snowboarding','snowkiting'," + + "'snowmobiling','somersaulting','spinning_poi','spray_painting','spraying','springboard_diving','squat'," + + "'sticking_tongue_out','stomping_grapes','stretching_arm','stretching_leg','strumming_guitar','surfing_crowd'," + + "'surfing_water','sweeping_floor','swimming_backstroke','swimming_breast_stroke','swimming_butterfly_stroke'," + + "'swing_dancing','swinging_legs','swinging_on_something','sword_fighting','tai_chi','taking_a_shower','tango_dancing'," + + "'tap_dancing','tapping_guitar','tapping_pen','tasting_beer','tasting_food','testifying','texting','throwing_axe'," + + "'throwing_ball','throwing_discus','tickling','tobogganing','tossing_coin','tossing_salad','training_dog'," + + "'trapezing','trimming_or_shaving_beard','trimming_trees','triple_jump','tying_bow_tie','tying_knot_-not_on_a_tie-'," + + "'tying_tie','unboxing','unloading_truck','using_computer','using_remote_controller_-not_gaming-','using_segway'," + + "'vault','waiting_in_line','walking_the_dog','washing_dishes','washing_feet','washing_hair','washing_hands'," + + "'water_skiing','water_sliding','watering_plants','waxing_back','waxing_chest','waxing_eyebrows','waxing_legs'," + + "'weaving_basket','welding','whistling','windsurfing','wrapping_present','wrestling','writing','yawning','yoga','zumba'"; + + public CommonActionTranslator(int width, int height) { + this.width = width; + this.height = height; + } + + @Override + public void prepare(TranslatorContext ctx) throws Exception { + labels = rawLabels.replace("'", "").split(","); + labels = Arrays.stream(labels) + .map(String::trim) + .toArray(String[]::new); + } + + @Override + public NDList processInput(TranslatorContext ctx, Image input) { + NDManager manager = ctx.getNDManager(); + NDArray array = input.toNDArray(manager, Image.Flag.COLOR); // HWC, uint8 + array = NDImageUtils.resize(array, width, height); + // 转 float32 + array = array.toType(DataType.FLOAT32, false); + // 自己归一化,注意均值和标准差要乘以255,因为原始是0~255 + float[] mean = {0.485f * 255, 0.456f * 255, 0.406f * 255}; + float[] std = {0.229f * 255, 0.224f * 255, 0.225f * 255}; + // 增加 batch 维度,变成 (1, H, W, C) + array = array.expandDims(0); + System.out.println(Arrays.toString(array.getShape().getShape())); + return new NDList(array); + } + + @Override + public Classifications processOutput(TranslatorContext ctx, NDList list) { + NDArray output = list.singletonOrThrow(); + output = output.softmax(1); // 计算概率 + Classifications classifications = new Classifications(Arrays.asList(labels), output); + return classifications; + + } + + @Override + public Batchifier getBatchifier() { + return null; + } + + + + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/config/InstanceSegModelConfig.java b/vision/src/main/java/cn/smartjavaai/instanceseg/config/InstanceSegModelConfig.java new file mode 100644 index 0000000..dbbfbe0 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/config/InstanceSegModelConfig.java @@ -0,0 +1,55 @@ +package cn.smartjavaai.instanceseg.config; + +import cn.smartjavaai.common.config.ModelConfig; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.instanceseg.enums.InstanceSegModelEnum; +import cn.smartjavaai.objectdetection.constant.DetectorConstant; +import lombok.Data; + +import java.util.List; + +/** + * 实例分割模型参数配置 + * + * @author dwj + */ +@Data +public class InstanceSegModelConfig extends ModelConfig { + + /** + * 模型 + */ + private InstanceSegModelEnum modelEnum; + + + /** + * 模型路径 + */ + private String modelPath; + + + /** + * 允许的分类列表 + */ + private List allowedClasses; + + /** + * 置信度阈值 + */ + private float threshold = 0.3f; + + + public InstanceSegModelConfig() { + } + + public InstanceSegModelConfig(InstanceSegModelEnum modelEnum, DeviceEnum device) { + this.modelEnum = modelEnum; + setDevice(device); + } + + public InstanceSegModelConfig(InstanceSegModelEnum modelEnum) { + this.modelEnum = modelEnum; + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/criteria/InstanceSegCriteriaFactory.java b/vision/src/main/java/cn/smartjavaai/instanceseg/criteria/InstanceSegCriteriaFactory.java new file mode 100644 index 0000000..940ab0a --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/criteria/InstanceSegCriteriaFactory.java @@ -0,0 +1,47 @@ +package cn.smartjavaai.instanceseg.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.instanceseg.config.InstanceSegModelConfig; +import cn.smartjavaai.instanceseg.translator.YoloSegmentationTranslatorFactory2; +import org.apache.commons.lang3.StringUtils; + +import java.nio.file.Paths; +import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; + +/** + * 实例分割Criteria工厂 + * @author dwj + */ +public class InstanceSegCriteriaFactory { + + + public static Criteria createCriteria(InstanceSegModelConfig config) { + Device device = null; + if(!Objects.isNull(config.getDevice())){ + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + Criteria criteria = null; + ConcurrentHashMap params = new ConcurrentHashMap(); + params.putAll(config.getCustomParams()); +// YoloV5Translator.Builder builder = new YoloV5Translator.Builder() +// .optSynsetArtifactName("synset.txt").setPipeline() + criteria = + Criteria.builder() + .setTypes(Image.class, DetectedObjects.class) + .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : + config.getModelEnum().getModelUri()) + .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) + .optDevice(device) + .optEngine("PyTorch") + .optTranslatorFactory(new YoloSegmentationTranslatorFactory2()) + .optProgress(new ProgressBar()) + .build(); + return criteria; + } +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/entity/DetectParams.java b/vision/src/main/java/cn/smartjavaai/instanceseg/entity/DetectParams.java new file mode 100644 index 0000000..8dbbdba --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/entity/DetectParams.java @@ -0,0 +1,17 @@ +package cn.smartjavaai.instanceseg.entity; + +import lombok.Data; + +/** + * 检测参数 + * @author dwj + */ +@Data +public class DetectParams { + + /** + * 置信度阈值 + */ + private float threshold = 0.3f; + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/enums/InstanceSegModelEnum.java b/vision/src/main/java/cn/smartjavaai/instanceseg/enums/InstanceSegModelEnum.java new file mode 100644 index 0000000..d87ccb6 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/enums/InstanceSegModelEnum.java @@ -0,0 +1,43 @@ +package cn.smartjavaai.instanceseg.enums; + +/** + * 实例分割模型枚举 + * @author dwj + */ +public enum InstanceSegModelEnum { + + SEG_YOLO11N_PYTORCH("djl://ai.djl.pytorch/yolo11n-seg"), + + SEG_YOLOV8N_PYTORCH("djl://ai.djl.pytorch/yolo11n-seg"), + + SEG_YOLO11N_ONNX("djl://ai.djl.onnxruntime/yolo11n-seg"), + + SEG_YOLOV8N_ONNX("djl://ai.djl.onnxruntime/yolov8n-seg"), + + SEG_MASK_RCNN("djl://ai.djl.mxnet/mask_rcnn"); + + + /** + * 根据名称获取枚举 (忽略大小写和下划线变体) + */ + public static InstanceSegModelEnum fromName(String name) { + String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); + for (InstanceSegModelEnum model : values()) { + if (model.name().replaceAll("_", "").equals(formatted)) { + return model; + } + } + throw new IllegalArgumentException("未知模型名称: " + name); + } + + private final String modelUri; + + InstanceSegModelEnum(String modelUri) { + this.modelUri = modelUri; + } + + public String getModelUri() { + return modelUri; + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/exception/InstanceSegException.java b/vision/src/main/java/cn/smartjavaai/instanceseg/exception/InstanceSegException.java new file mode 100644 index 0000000..6e859f7 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/exception/InstanceSegException.java @@ -0,0 +1,29 @@ +package cn.smartjavaai.instanceseg.exception; + +/** + * 实例分割异常 + * @author dwj + */ +public class InstanceSegException extends RuntimeException{ + + public InstanceSegException() { + super(); + } + + public InstanceSegException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) { + super(message, cause, enableSuppression, writableStackTrace); + } + + public InstanceSegException(String message, Throwable cause) { + super(message, cause); + } + + public InstanceSegException(String message) { + super(message); + } + + public InstanceSegException(Throwable cause) { + super(cause); + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java b/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java new file mode 100644 index 0000000..6e3324f --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java @@ -0,0 +1,153 @@ +package cn.smartjavaai.instanceseg.model; + +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.output.DetectedObjects; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ModelNotFoundException; +import ai.djl.repository.zoo.ZooModel; +import cn.smartjavaai.common.cv.SmartImageFactory; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig; +import cn.smartjavaai.instanceseg.criteria.InstanceSegCriteriaFactory; +import cn.smartjavaai.instanceseg.exception.InstanceSegException; +import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.vision.utils.DetectedObjectsFilter; +import cn.smartjavaai.vision.utils.DetectorUtils; +import cn.smartjavaai.vision.utils.CategoryMaskFilter; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.collections.CollectionUtils; +import org.apache.commons.pool2.impl.GenericObjectPool; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.Objects; + +/** + * 实例分割模型 + * @author dwj + */ +@Slf4j +public class CommonInstanceSegModel implements InstanceSegModel { + + + private InstanceSegModelConfig config; + + private ZooModel model; + + private GenericObjectPool> predictorPool; + + @Override + public void loadModel(InstanceSegModelConfig config) { + if(Objects.isNull(config.getModelEnum())){ + throw new DetectionException("未配置模型枚举"); + } + Criteria criteria = InstanceSegCriteriaFactory.createCriteria(config); + this.config = config; + try { + model = criteria.loadModel(); + // 创建池子:每个线程独享 Predictor + this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model)); + int predictorPoolSize = config.getPredictorPoolSize(); + if(config.getPredictorPoolSize() <= 0){ + predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数 + } + predictorPool.setMaxTotal(predictorPoolSize); + log.debug("当前设备: " + model.getNDManager().getDevice()); + log.debug("当前引擎: " + Engine.getInstance().getEngineName()); + log.debug("模型推理器线程池最大数量: " + predictorPoolSize); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new DetectionException("模型加载失败", e); + } + } + + @Override + public R detect(Image image) { + DetectedObjects detectedObjects = detectCore(image); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image); + return R.ok(detectionResponse); + } + + /** + * 模型核心推理方法 + * @param image + * @return + */ + @Override + public DetectedObjects detectCore(Image image) { + Predictor predictor = null; + try { + predictor = predictorPool.borrowObject(); + DetectedObjects detectedObjects = predictor.predict(image); + //过滤 + if(Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){ + DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getThreshold()); + detectedObjects = detectedObjectsFilter.filter(detectedObjects); + } + return detectedObjects; + } catch (Exception e) { + throw new DetectionException("实例分割错误", e); + }finally { + if (predictor != null) { + try { + predictorPool.returnObject(predictor); //归还 + log.debug("释放资源"); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + @Override + public R detectAndDraw(Image image) { + DetectedObjects detectedObjects = detectCore(image); + image.drawBoundingBoxes(detectedObjects); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image); + detectionResponse.setDrawnImage(image); + return R.ok(detectionResponse); + } + + @Override + public R detectAndDraw(String imagePath, String outputPath) { + try { + Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detectedObjects = detectCore(img); + img.drawBoundingBoxes(detectedObjects); + img.save(Files.newOutputStream(Paths.get(outputPath)), "png"); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, img); + return R.ok(detectionResponse); + } catch (IOException e) { + throw new InstanceSegException(e); + } + } + + @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); + } + } +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java b/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java new file mode 100644 index 0000000..1ad627f --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java @@ -0,0 +1,47 @@ +package cn.smartjavaai.instanceseg.model; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.CategoryMask; +import ai.djl.modality.cv.output.DetectedObjects; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig; + +import java.awt.image.BufferedImage; + +/** + * 实例分割模型 + * @author dwj + */ +public interface InstanceSegModel extends AutoCloseable{ + + + /** + * 加载模型 + * @param config + */ + void loadModel(InstanceSegModelConfig config); + + /** + * 实例分割 + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + default DetectedObjects detectCore(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + default R detectAndDraw(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + default R detectAndDraw(String imagePath, String outputPath){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/translator/YoloSegmentationTranslator2.java b/vision/src/main/java/cn/smartjavaai/instanceseg/translator/YoloSegmentationTranslator2.java new file mode 100644 index 0000000..35ef997 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/translator/YoloSegmentationTranslator2.java @@ -0,0 +1,182 @@ +/* + * Copyright 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file except in compliance + * with the License. A copy of the License is located at + * + * http://aws.amazon.com/apache2.0/ + * + * or in the "license" file accompanying this file. This file is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES + * OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions + * and limitations under the License. + */ +package cn.smartjavaai.instanceseg.translator; + +import ai.djl.modality.cv.output.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.output.Mask; +import ai.djl.modality.cv.output.Rectangle; +import ai.djl.modality.cv.transform.Resize; +import ai.djl.modality.cv.transform.ToTensor; +import ai.djl.modality.cv.translator.YoloV5Translator; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.types.DataType; +import ai.djl.translate.ArgumentsUtil; +import ai.djl.translate.Pipeline; +import ai.djl.translate.TranslatorContext; + +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +/** A translator for Yolov8 instance segmentation models. */ +public class YoloSegmentationTranslator2 extends YoloV5Translator { + + private static final int[] AXIS_0 = {0}; + private static final int[] AXIS_1 = {1}; + + private float threshold; + private float nmsThreshold; + + /** + * Creates the instance segmentation translator from the given builder. + * + * @param builder the builder for the translator + */ + public YoloSegmentationTranslator2(Builder builder) { + super(builder); + this.threshold = 0.25f; + this.nmsThreshold = 0.4F; + } + + /** {@inheritDoc} */ + @Override + public DetectedObjects processOutput(TranslatorContext ctx, NDList list) { + NDArray pred = list.get(0); + NDArray protos = list.get(1); + int maskIndex = classes.size() + 4; + NDArray candidates = pred.get("4:" + maskIndex).max(AXIS_0).gt(threshold); + pred = pred.transpose(); + NDArray sub = pred.get("..., :4"); + sub = xywh2xyxy(sub); + pred = sub.concat(pred.get("..., 4:"), -1); + pred = pred.get(candidates); + + NDList split = pred.split(new long[] {4, maskIndex}, 1); + NDArray box = split.get(0); + + int numBox = Math.toIntExact(box.getShape().get(0)); + + float[] buf = box.toFloatArray(); + float[] confidences = split.get(1).max(AXIS_1).toFloatArray(); + long[] ids = split.get(1).argMax(1).toLongArray(); + + List boxes = new ArrayList<>(numBox); + List scores = new ArrayList<>(numBox); + for (int i = 0; i < numBox; ++i) { + float xPos = buf[i * 4]; + float yPos = buf[i * 4 + 1]; + float w = buf[i * 4 + 2] - xPos; + float h = buf[i * 4 + 3] - yPos; + Rectangle rect = new Rectangle(xPos, yPos, w, h); + boxes.add(rect); + scores.add((double) confidences[i]); + } + List nms = Rectangle.nms(boxes, scores, nmsThreshold); + long[] idx = nms.stream().mapToLong(Integer::longValue).toArray(); + NDArray selected = box.getManager().create(idx); + NDArray masks = split.get(2).get(selected); + + int maskW = Math.toIntExact(protos.getShape().get(2)); + int maskH = Math.toIntExact(protos.getShape().get(1)); + + protos = protos.reshape(32, (long) maskH * maskW); + masks = + masks.matMul(protos) + .reshape(nms.size(), maskH, maskW) + .gt(0f) + .toType(DataType.FLOAT32, true); + + float[] maskArray = masks.toFloatArray(); + box = box.get(selected); + buf = box.toFloatArray(); + + List retClasses = new ArrayList<>(); + List retProbs = new ArrayList<>(); + List retBB = new ArrayList<>(); + for (int i = 0; i < idx.length; ++i) { + float x = buf[i * 4] / width; + float y = buf[i * 4 + 1] / height; + float w = buf[i * 4 + 2] / width - x; + float h = buf[i * 4 + 3] / width - y; + int id = nms.get(i); + retClasses.add(classes.get((int) ids[id])); + retProbs.add((double) confidences[id]); + + float[][] maskFloat = new float[maskH][maskW]; + int pos = i * maskH * maskW; + for (int j = 0; j < maskH; j++) { + System.arraycopy(maskArray, pos + j * maskW, maskFloat[j], 0, maskW); + } + Mask bb = new Mask(x, y, w, h, maskFloat, true); + retBB.add(bb); + } + return new DetectedObjects(retClasses, retProbs, retBB); + } + + private NDArray xywh2xyxy(NDArray array) { + NDArray xy = array.get("..., :2", new Object[0]); + NDArray wh = array.get("..., 2:", new Object[0]).div(2); + return xy.sub(wh).concat(xy.add(wh), -1); + } + + /** + * Creates a builder to build a {@code YoloSegmentationTranslator}. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Creates a builder to build a {@code YoloSegmentationTranslator} with specified arguments. + * + * @param arguments arguments to specify builder options + * @return a new builder + */ + public static Builder builder(Map arguments) { + Builder builder = new Builder(); + builder.optSynsetArtifactName("synset.txt"); + builder.setImageSize(640, 640); + return builder; + } + + /** The builder for instance segmentation translator. */ + public static class Builder extends YoloV5Translator.Builder { + + Builder() {} + + /** {@inheritDoc} */ + @Override + protected Builder self() { + return this; + } + + /** {@inheritDoc} */ + @Override + public YoloSegmentationTranslator2 build() { + pipeline = new Pipeline(); + pipeline.add(new Resize(640, 640)); + pipeline.add(new ToTensor()); +// validate(); + return new YoloSegmentationTranslator2(this); + } + } + + + + + +} diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/translator/YoloSegmentationTranslatorFactory2.java b/vision/src/main/java/cn/smartjavaai/instanceseg/translator/YoloSegmentationTranslatorFactory2.java new file mode 100644 index 0000000..c8d632c --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/translator/YoloSegmentationTranslatorFactory2.java @@ -0,0 +1,37 @@ +/* + * Copyright 2024 Amazon.com, Inc. or its affiliates. All Rights Reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file except in compliance + * with the License. A copy of the License is located at + * + * http://aws.amazon.com/apache2.0/ + * + * or in the "license" file accompanying this file. This file is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES + * OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions + * and limitations under the License. + */ +package cn.smartjavaai.instanceseg.translator; + +import ai.djl.Model; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.translator.ObjectDetectionTranslatorFactory; +import ai.djl.translate.Translator; + +import java.io.Serializable; +import java.util.Map; + +/** A translatorFactory that creates a {@link ai.djl.modality.cv.translator.YoloSegmentationTranslator} instance. */ +public class YoloSegmentationTranslatorFactory2 extends ObjectDetectionTranslatorFactory + implements Serializable { + + private static final long serialVersionUID = 1L; + + /** {@inheritDoc} */ + @Override + protected Translator buildBaseTranslator( + Model model, Map arguments) { + Translator translator = YoloSegmentationTranslator2.builder(arguments).build(); + return translator; + } +} diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/config/DetectorModelConfig.java b/vision/src/main/java/cn/smartjavaai/objectdetection/config/DetectorModelConfig.java similarity index 100% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/config/DetectorModelConfig.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/config/DetectorModelConfig.java diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/config/PersonDetModelConfig.java b/vision/src/main/java/cn/smartjavaai/objectdetection/config/PersonDetModelConfig.java new file mode 100644 index 0000000..b3f3f5f --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/config/PersonDetModelConfig.java @@ -0,0 +1,62 @@ +package cn.smartjavaai.objectdetection.config; + +import cn.smartjavaai.common.config.ModelConfig; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.objectdetection.constant.DetectorConstant; +import cn.smartjavaai.objectdetection.enums.DetectorModelEnum; +import cn.smartjavaai.objectdetection.enums.PersonDetectorModelEnum; +import lombok.Data; + +import java.util.List; + +/** + * 行人检测模型参数配置 + * + * @author dwj + * @date 2025/4/4 + */ +@Data +public class PersonDetModelConfig extends ModelConfig { + + /** + * 模型 + */ + private PersonDetectorModelEnum modelEnum; + + /** + * 置信度阈值 + */ + private float threshold; + + + /** + * 模型路径 + */ + private String modelPath; + + + /** + * 允许的分类列表 + */ + private List allowedClasses; + + /** + * 按置信度分数排序后,最多保留的检测框数量 + */ + private int topK; + + + public PersonDetModelConfig() { + } + + public PersonDetModelConfig(PersonDetectorModelEnum modelEnum, DeviceEnum device) { + this.modelEnum = modelEnum; + setDevice(device); + } + + public PersonDetModelConfig(PersonDetectorModelEnum modelEnum) { + this.modelEnum = modelEnum; + } + + +} diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/constant/DetectorConstant.java b/vision/src/main/java/cn/smartjavaai/objectdetection/constant/DetectorConstant.java similarity index 100% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/constant/DetectorConstant.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/constant/DetectorConstant.java diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderFactory.java b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderFactory.java similarity index 88% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderFactory.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderFactory.java index 44cc0d2..da8aca1 100644 --- a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderFactory.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderFactory.java @@ -20,7 +20,8 @@ public class CriteriaBuilderFactory { if(config.getModelEnum() == DetectorModelEnum.YOLOV8_OFFICIAL || config.getModelEnum() == DetectorModelEnum.YOLOV12_OFFICIAL || config.getModelEnum() == DetectorModelEnum.YOLOV8_CUSTOM || - config.getModelEnum() == DetectorModelEnum.YOLOV12_CUSTOM){ + config.getModelEnum() == DetectorModelEnum.YOLOV12_CUSTOM || + config.getModelEnum() == DetectorModelEnum.TENSORFLOW2_OFFICIAL){ if(StringUtils.isBlank(config.getModelPath())){ throw new DetectionException("modelPath is null"); } @@ -34,6 +35,8 @@ public class CriteriaBuilderFactory { return new YoloCriteriaBuilder().buildCriteria(config); case YOLOV12_CUSTOM: return new YoloCriteriaBuilder().buildCriteria(config); + case TENSORFLOW2_OFFICIAL: + return new Tensorflow2CriteriaBuilder().buildCriteria(config); // 其他类型 default: return new DJLModelCriteriaBuilder().buildCriteria(config); diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderStrategy.java b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderStrategy.java similarity index 100% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderStrategy.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/criteria/CriteriaBuilderStrategy.java diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/DJLModelCriteriaBuilder.java b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/DJLModelCriteriaBuilder.java similarity index 100% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/DJLModelCriteriaBuilder.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/criteria/DJLModelCriteriaBuilder.java diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/PersonDetCriteriaFactory.java b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/PersonDetCriteriaFactory.java new file mode 100644 index 0000000..0be1097 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/PersonDetCriteriaFactory.java @@ -0,0 +1,76 @@ +package cn.smartjavaai.objectdetection.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 ai.djl.translate.Translator; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.obb.config.ObbDetModelConfig; +import cn.smartjavaai.obb.entity.ObbResult; +import cn.smartjavaai.obb.enums.ObbDetModelEnum; +import cn.smartjavaai.obb.exception.ObbDetException; +import cn.smartjavaai.obb.translator.YoloV11OddTranslator; +import cn.smartjavaai.objectdetection.config.PersonDetModelConfig; +import cn.smartjavaai.objectdetection.enums.PersonDetectorModelEnum; +import cn.smartjavaai.objectdetection.translator.YoloV8PersonDetTranslator; +import org.apache.commons.lang3.StringUtils; + +import java.nio.file.Paths; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; + +/** + * 行人检测Criteria创建工厂 + * @author dwj + */ +public class PersonDetCriteriaFactory { + + + /** + * 创建行人检测Criteria + * @param config + * @return + */ + public static Criteria createCriteria(PersonDetModelConfig config) { + Device device = null; + if(!Objects.isNull(config.getDevice())){ + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + Translator translator = getTranslator(config); + //检查模型路径 + if (StringUtils.isBlank(config.getModelPath())){ + throw new ObbDetException("请指定模型路径"); + } + Criteria criteria = + Criteria.builder() + .setTypes(Image.class, DetectedObjects.class) + .optModelPath(Paths.get(config.getModelPath())) + .optTranslator(translator) + .optDevice(device) + .optProgress(new ProgressBar()) + .optEngine(config.getModelEnum().getEngine()) + .build(); + return criteria; + } + + + /** + * 获取行人检测Translator + * @param config + * @return + */ + public static Translator getTranslator(PersonDetModelConfig config) { + Translator translator = null; + if (config.getModelEnum() == PersonDetectorModelEnum.YOLOV8_PERSON){ + translator = YoloV8PersonDetTranslator.builder() + .setImageSize(config.getModelEnum().getInputSize(), config.getModelEnum().getInputSize()) + .optThreshold(config.getThreshold() > 0 ? config.getThreshold() : 0.5f) + .optNmsThreshold(0.45f).build(); + } + return translator; + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/Tensorflow2CriteriaBuilder.java b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/Tensorflow2CriteriaBuilder.java new file mode 100644 index 0000000..691c15c --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/Tensorflow2CriteriaBuilder.java @@ -0,0 +1,119 @@ +package cn.smartjavaai.objectdetection.criteria; + +import ai.djl.Application; +import ai.djl.Device; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.translator.YoloV8TranslatorFactory; +import ai.djl.repository.zoo.Criteria; +import ai.djl.training.util.ProgressBar; +import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.objectdetection.config.DetectorModelConfig; +import cn.smartjavaai.objectdetection.constant.DetectorConstant; +import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.objectdetection.translator.TensorflowTranslator; +import cn.smartjavaai.vision.utils.TensorflowSynsetUtils; +import org.apache.commons.lang3.StringUtils; + +import java.io.IOException; +import java.net.URL; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; + +/** + * YOLO模型Criteria 构建器 + * @author dwj + * @date 2025/5/14 + */ +public class Tensorflow2CriteriaBuilder implements CriteriaBuilderStrategy { + @Override + public Criteria buildCriteria(DetectorModelConfig config) { + Device device = null; + if(!Objects.isNull(config.getDevice())){ + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + + Map customParams = getDefaultConfig(); + // 合并用户自定义参数(如有重复,覆盖默认默认值) + if (config.getCustomParams() != null) { + customParams.putAll(config.getCustomParams()); + } + //解析synset + String synsetUrl = (String)customParams.get("synsetUrl"); + String synsetPath = (String)customParams.get("synsetPath"); + String synsetFileName = (String)customParams.get("synsetFileName"); + Map classes = null; + if(StringUtils.isNotBlank(synsetUrl)){ + try { + classes = TensorflowSynsetUtils.loadSynset(new URL(synsetUrl)); + } catch (IOException e) { + throw new DetectionException("加载synset异常", e); + } + }else if(StringUtils.isNotBlank(synsetPath)){ + try { + if(!Files.exists(Paths.get(synsetPath))){ + throw new DetectionException("synsetPath:" + synsetPath + "不存在"); + } + classes = TensorflowSynsetUtils.loadSynset(Paths.get(synsetPath)); + } catch (IOException e) { + throw new DetectionException("加载synset异常", e); + } + }else if(StringUtils.isNotBlank(synsetFileName)){ + if(StringUtils.isBlank(config.getModelPath())){ + throw new DetectionException("指定synsetFileName,需同时指定modelPath"); + } + try { + Path modelPath = Paths.get(config.getModelPath()); + Path synset = modelPath.resolve(synsetFileName); + if(!Files.exists(synset)){ + throw new DetectionException(synset.toAbsolutePath().toString() + " 不存在"); + } + //模型同目录下存在synsetFileName + classes = TensorflowSynsetUtils.loadSynset(synset); + } catch (IOException e) { + throw new DetectionException("加载synset异常", e); + } + }else{ + if(StringUtils.isBlank(config.getModelPath())){ + throw new DetectionException("modelPath is null"); + } + try { + Path modelPath = Paths.get(config.getModelPath()); + synsetFileName = "mscoco_label_map.pbtxt"; + Path synset = modelPath.resolve(synsetFileName); + if(!Files.exists(synset)){ + throw new DetectionException(synset.toAbsolutePath().toString() + " 不存在"); + } + //模型同目录下存在synsetFileName + classes = TensorflowSynsetUtils.loadSynset(synset); + } catch (IOException e) { + throw new DetectionException("加载synset异常", e); + } + } + customParams.put("classes", classes); + Criteria.Builder criteriaBuilder = Criteria.builder() + .optApplication(Application.CV.OBJECT_DETECTION) + .setTypes(Image.class, DetectedObjects.class) + .optModelPath(Paths.get(config.getModelPath())) + .optModelName("saved_model") + .optEngine("TensorFlow") + .optDevice(device) + .optTranslator(new TensorflowTranslator(customParams)) + .optProgress(new ProgressBar()); + if(config.getMaxBox() > 0){ + criteriaBuilder.optArgument("maxBox", config.getMaxBox()); + } + Criteria criteria = criteriaBuilder.build(); + return criteria; + } + + public Map getDefaultConfig(){ + return new HashMap<>(); + } +} diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/YoloCriteriaBuilder.java b/vision/src/main/java/cn/smartjavaai/objectdetection/criteria/YoloCriteriaBuilder.java similarity index 100% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/criteria/YoloCriteriaBuilder.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/criteria/YoloCriteriaBuilder.java diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/enums/DetectorModelEnum.java b/vision/src/main/java/cn/smartjavaai/objectdetection/enums/DetectorModelEnum.java similarity index 96% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/enums/DetectorModelEnum.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/enums/DetectorModelEnum.java index 4fb6d8a..c03cc6c 100644 --- a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/enums/DetectorModelEnum.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/enums/DetectorModelEnum.java @@ -41,7 +41,10 @@ public enum DetectorModelEnum { YOLOV8_CUSTOM(""), - YOLOV12_CUSTOM(""); + YOLOV12_CUSTOM(""), + + // TensorFlow 2.x 官方模型 + TENSORFLOW2_OFFICIAL(""); /** * 根据名称获取枚举 (忽略大小写和下划线变体) diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/enums/PersonDetectorModelEnum.java b/vision/src/main/java/cn/smartjavaai/objectdetection/enums/PersonDetectorModelEnum.java new file mode 100644 index 0000000..44db37f --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/enums/PersonDetectorModelEnum.java @@ -0,0 +1,51 @@ +package cn.smartjavaai.objectdetection.enums; + +/** + * 目标检测模型枚举 + * @author dwj + * @date 2025/4/4 + */ +public enum PersonDetectorModelEnum { + + YOLOV8_PERSON("OnnxRuntime", 1280); + + + /** + * 模型输入尺寸 + */ + private final int inputSize; + + /** + * 模型引擎 + */ + private final String engine; + + PersonDetectorModelEnum(String engine, int inputSize) { + this.inputSize = inputSize; + this.engine = engine; + } + + + public int getInputSize() { + return inputSize; + } + + public String getEngine() { + return engine; + } + + /** + * 根据名称获取枚举 (忽略大小写和下划线变体) + */ + public static PersonDetectorModelEnum fromName(String name) { + String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); + for (PersonDetectorModelEnum model : values()) { + if (model.name().replaceAll("_", "").equals(formatted)) { + return model; + } + } + throw new IllegalArgumentException("未知模型名称: " + name); + } + + +} diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/exception/DetectionException.java b/vision/src/main/java/cn/smartjavaai/objectdetection/exception/DetectionException.java similarity index 100% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/exception/DetectionException.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/exception/DetectionException.java diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java b/vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java similarity index 79% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java index b4e3b40..fb16e38 100644 --- a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java @@ -7,29 +7,23 @@ 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.translator.YoloV8TranslatorFactory; import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ZooModel; -import ai.djl.training.util.ProgressBar; -import cn.smartjavaai.common.entity.DetectionInfo; -import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; -import cn.smartjavaai.common.entity.Point; -import cn.smartjavaai.common.entity.face.FaceInfo; import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.utils.FileUtils; +import cn.smartjavaai.common.utils.FrameConverterUtil; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.objectdetection.config.DetectorModelConfig; -import cn.smartjavaai.objectdetection.constant.DetectorConstant; import cn.smartjavaai.objectdetection.criteria.CriteriaBuilderFactory; import cn.smartjavaai.objectdetection.exception.DetectionException; -import cn.smartjavaai.objectdetection.utils.DetectorUtils; +import cn.smartjavaai.vision.utils.CategoryMaskFilter; +import cn.smartjavaai.vision.utils.DetectedObjectsFilter; +import cn.smartjavaai.vision.utils.DetectorUtils; import lombok.extern.slf4j.Slf4j; import org.apache.commons.collections.CollectionUtils; -import org.apache.commons.lang3.StringUtils; -import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.impl.GenericObjectPool; import org.opencv.core.Mat; @@ -38,10 +32,8 @@ import java.awt.image.BufferedImage; import java.io.*; import java.nio.file.Paths; import java.util.ArrayList; -import java.util.Iterator; import java.util.List; import java.util.Objects; -import java.util.stream.Collectors; /** * 目标检测模型 @@ -56,6 +48,15 @@ public class DetectorModel implements AutoCloseable{ private DetectorModelConfig config; + private boolean fromFactory = false; + + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + public void loadModel(DetectorModelConfig config){ if(Objects.isNull(config.getModelEnum())){ throw new DetectionException("未配置模型枚举"); @@ -212,7 +213,11 @@ public class DetectorModel implements AutoCloseable{ try { predictor = predictorPool.borrowObject(); DetectedObjects detectedObjects = predictor.predict(image); - detectedObjects = filterDetections(detectedObjects); + if(CollectionUtils.isNotEmpty(config.getAllowedClasses()) + && Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){ + DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getThreshold(),config.getTopK()); + detectedObjects = detectedObjectsFilter.filter(detectedObjects); + } return detectedObjects; } catch (Exception e) { throw new DetectionException("目标检测错误", e); @@ -233,45 +238,6 @@ public class DetectorModel implements AutoCloseable{ } } - /** - * 筛选检测结果 - * @param detectedObjects - * @return - */ - private DetectedObjects filterDetections(DetectedObjects detectedObjects) { - if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0) { - return detectedObjects; - } - List items = detectedObjects.items(); - // 按照允许的类别进行过滤 - List filtered = new ArrayList<>(); - //过滤类别 - if(!CollectionUtils.isEmpty(config.getAllowedClasses())){ - for (DetectedObjects.DetectedObject obj : items) { - if(config.getAllowedClasses().contains(obj.getClassName())){ - filtered.add(obj); - } - } - }else{ - filtered = items; - } - // 按照概率进行排序 - filtered.sort((o1, o2) -> Double.compare(o2.getProbability(), o1.getProbability())); - if(config.getTopK() > 0 && filtered.size() > config.getTopK()){ - filtered = filtered.subList(0, config.getTopK()); - } - // 构建新的 DetectedObjects 返回 - List names = new ArrayList<>(); - List probs = new ArrayList<>(); - List boxes = new ArrayList<>(); - - for (DetectedObjects.DetectedObject obj : filtered) { - names.add(obj.getClassName()); - probs.add(obj.getProbability()); - boxes.add(obj.getBoundingBox()); - } - return new DetectedObjects(names, probs, boxes); - } public GenericObjectPool> getPool() { return predictorPool; @@ -283,6 +249,9 @@ public class DetectorModel implements AutoCloseable{ */ @Override public void close() { + if (fromFactory) { + ObjectDetectionModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); diff --git a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/model/ObjectDetectionModelFactory.java b/vision/src/main/java/cn/smartjavaai/objectdetection/model/ObjectDetectionModelFactory.java similarity index 82% rename from smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/model/ObjectDetectionModelFactory.java rename to vision/src/main/java/cn/smartjavaai/objectdetection/model/ObjectDetectionModelFactory.java index d65b3ac..8122587 100644 --- a/smartjavaai-objectdetection/src/main/java/cn/smartjavaai/objectdetection/model/ObjectDetectionModelFactory.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/model/ObjectDetectionModelFactory.java @@ -19,7 +19,7 @@ public class ObjectDetectionModelFactory { // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 private static volatile ObjectDetectionModelFactory instance; - private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); + private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); static{ log.debug("缓存目录:{}", Config.getCachePath()); @@ -49,9 +49,10 @@ public class ObjectDetectionModelFactory { if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){ throw new DetectionException("未配置模型"); } - return modelMap.computeIfAbsent(config.getModelEnum().name(), k -> { + return modelMap.computeIfAbsent(config.getModelEnum(), k -> { DetectorModel model = new DetectorModel(); model.loadModel(config); + model.setFromFactory(true); return model; }); } @@ -72,6 +73,15 @@ public class ObjectDetectionModelFactory { */ public void closeAll() { modelMap.values().forEach(DetectorModel::close); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(DetectorModelEnum modelEnum) { + modelMap.remove(modelEnum); } } diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/CommonPersonDetModel.java b/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/CommonPersonDetModel.java new file mode 100644 index 0000000..1fa7347 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/CommonPersonDetModel.java @@ -0,0 +1,160 @@ +package cn.smartjavaai.objectdetection.model.person; + +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.output.DetectedObjects; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ModelNotFoundException; +import ai.djl.repository.zoo.ZooModel; +import cn.smartjavaai.common.cv.SmartImageFactory; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.obb.config.ObbDetModelConfig; +import cn.smartjavaai.obb.criteria.ObbDetCriteriaFactory; +import cn.smartjavaai.obb.entity.ObbResult; +import cn.smartjavaai.obb.exception.ObbDetException; +import cn.smartjavaai.objectdetection.config.PersonDetModelConfig; +import cn.smartjavaai.objectdetection.criteria.PersonDetCriteriaFactory; +import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.vision.utils.DetectedObjectsFilter; +import cn.smartjavaai.vision.utils.DetectorUtils; +import cn.smartjavaai.vision.utils.ObbResultFilter; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.collections.CollectionUtils; +import org.apache.commons.pool2.impl.GenericObjectPool; +import org.opencv.core.Mat; + +import java.io.ByteArrayOutputStream; +import java.io.FileOutputStream; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Paths; +import java.util.Objects; + +/** + * 行人检测模型 + * @author dwj + */ +@Slf4j +public class CommonPersonDetModel implements PersonDetModel { + + + private PersonDetModelConfig config; + + private ZooModel model; + + private GenericObjectPool> predictorPool; + + @Override + public void loadModel(PersonDetModelConfig config) { + if(Objects.isNull(config.getModelEnum())){ + throw new DetectionException("未配置模型枚举"); + } + Criteria criteria = PersonDetCriteriaFactory.createCriteria(config); + this.config = config; + try { + model = criteria.loadModel(); + // 创建池子:每个线程独享 Predictor + this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model)); + int predictorPoolSize = config.getPredictorPoolSize(); + if(config.getPredictorPoolSize() <= 0){ + predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数 + } + predictorPool.setMaxTotal(predictorPoolSize); + log.debug("当前设备: " + model.getNDManager().getDevice()); + log.debug("当前引擎: " + Engine.getInstance().getEngineName()); + log.debug("模型推理器线程池最大数量: " + predictorPoolSize); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new DetectionException("模型加载失败", e); + } + } + + @Override + public R detect(Image image) { + DetectedObjects detectedObjects = detectCore(image); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image); + return R.ok(detectionResponse); + } + + /** + * 模型核心推理方法 + * @param image + * @return + */ + @Override + public DetectedObjects detectCore(Image image) { + Predictor predictor = null; + try { + predictor = predictorPool.borrowObject(); + DetectedObjects detectedObjects = predictor.predict(image); + //过滤 + if(Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){ + DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getTopK()); + detectedObjects = detectedObjectsFilter.filter(detectedObjects); + } + return detectedObjects; + } catch (Exception e) { + throw new DetectionException("行人检测错误", e); + }finally { + if (predictor != null) { + try { + predictorPool.returnObject(predictor); //归还 + log.debug("释放资源"); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + @Override + public R detectAndDraw(Image image) { + DetectedObjects detectedObjects = detectCore(image); + image.drawBoundingBoxes(detectedObjects); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image); + detectionResponse.setDrawnImage(image); + return R.ok(detectionResponse); + } + + @Override + public R detectAndDraw(String imagePath, String outputPath) { + try { + Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detectedObjects = detectCore(img); + img.drawBoundingBoxes(detectedObjects); + ByteArrayOutputStream outputStream = new ByteArrayOutputStream(); + // 调用 save 方法将 Image 写入字节流 + img.save(new FileOutputStream(Paths.get(outputPath).toAbsolutePath().toString()), "png"); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, img); + return R.ok(detectionResponse); + } catch (IOException e) { + throw new ObbDetException(e); + } + } + + @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); + } + } +} diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/PersonDetModel.java b/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/PersonDetModel.java new file mode 100644 index 0000000..4a4367e --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/PersonDetModel.java @@ -0,0 +1,62 @@ +package cn.smartjavaai.objectdetection.model.person; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.DetectedObjects; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.obb.config.ObbDetModelConfig; +import cn.smartjavaai.obb.entity.ObbResult; +import cn.smartjavaai.objectdetection.config.PersonDetModelConfig; + +/** + * 行人检测模型 + * @author dwj + */ +public interface PersonDetModel extends AutoCloseable{ + + + /** + * 加载模型 + * @param config + */ + void loadModel(PersonDetModelConfig config); + + /** + * 旋转框 + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 行人检测 核心方法 + * @param image + * @return + */ + default DetectedObjects detectCore(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 行人检测并绘制 + * @param image + * @return + */ + default R detectAndDraw(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 行人检测并绘制 + * @param imagePath + * @param outputPath + * @return + */ + default R detectAndDraw(String imagePath, String outputPath){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetectionListener.java b/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetectionListener.java new file mode 100644 index 0000000..6a80f0f --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetectionListener.java @@ -0,0 +1,16 @@ +package cn.smartjavaai.objectdetection.stream; + +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.entity.DetectionInfo; + +import java.util.List; + +/** + * 视频流目标检测监听器 + * @author dwj + */ +public interface StreamDetectionListener { + + + void onObjectDetected(List detectionInfoList, Image image); +} diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java b/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java new file mode 100644 index 0000000..4532422 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java @@ -0,0 +1,280 @@ +package cn.smartjavaai.objectdetection.stream; + +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.Mask; +import ai.djl.modality.cv.output.Rectangle; +import cn.hutool.core.lang.UUID; +import cn.smartjavaai.common.entity.*; +import cn.smartjavaai.common.enums.VideoSourceType; +import cn.smartjavaai.common.utils.FrameConverterUtil; +import cn.smartjavaai.common.utils.ImageUtils; +import cn.smartjavaai.common.utils.OpenCVUtils; +import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.objectdetection.model.DetectorModel; +import cn.smartjavaai.vision.utils.DetectorUtils; +import lombok.extern.slf4j.Slf4j; +import nu.pattern.OpenCV; +import org.apache.commons.lang3.StringUtils; +import org.bytedeco.javacv.*; +import org.bytedeco.opencv.global.opencv_imgcodecs; +import org.opencv.core.Core; +import org.opencv.core.Mat; +import org.opencv.imgcodecs.Imgcodecs; + +import java.util.*; +import java.util.concurrent.*; + +/** + * 视频流目标检测器 + * @author dwj + */ +@Slf4j +public class StreamDetector implements AutoCloseable{ + + static { + OpenCV.loadLocally(); + } + + + private DetectorModel detectorModel; + private String streamUrl; + private ExecutorService grabberExecutor; // 专门抓帧的线程 + private ExecutorService processorExecutor; // 专门处理帧的线程 + private int frameDetectionInterval = 1; + private int repeatGap = 5; // 秒 + private volatile boolean isRunning; + private FrameGrabber grabber; + private StreamDetectionListener listener; + private OpenCVFrameConverter.ToOrgOpenCvCoreMat converterToMat; + private VideoSourceType sourceType = VideoSourceType.STREAM; // 默认流 + private int cameraIndex = 0; // 默认第一个摄像头 + + private Map lastDetectTime = new ConcurrentHashMap<>(); + private BlockingQueue frameQueue = new LinkedBlockingQueue<>(100); + + public static Builder builder() { return new Builder(); } + + private StreamDetector(Builder builder) { + this.detectorModel = builder.detectorModel; + this.streamUrl = builder.streamUrl; + this.frameDetectionInterval = builder.frameDetectionInterval; + this.listener = builder.listener; + this.sourceType = builder.sourceType; + this.cameraIndex = builder.cameraIndex; + this.converterToMat = new OpenCVFrameConverter.ToOrgOpenCvCoreMat(); + + + } + + private void initializeGrabber() throws FrameGrabber.Exception { + if (sourceType == VideoSourceType.CAMERA) { + // 使用用户传入的摄像头索引 + grabber = new OpenCVFrameGrabber(cameraIndex); + } else { + grabber = new FFmpegFrameGrabber(streamUrl); + if (sourceType == VideoSourceType.STREAM) { + grabber.setOption("rtsp_transport", "tcp"); + grabber.setOption("buffer_size", "1024000"); + grabber.setOption("stimeout", "20000000"); + grabber.setOption("max_delay", "500000"); + } + } + grabber.start(); + } + + public void startDetection() { + if (isRunning) return; + isRunning = true; + + // 初始化抓帧线程池 + if (grabberExecutor == null) grabberExecutor = Executors.newSingleThreadExecutor(); + // 初始化帧处理线程池 + if (processorExecutor == null) processorExecutor = Executors.newSingleThreadExecutor(); + + // 初始化抓帧线程 + grabberExecutor.submit(() -> { + try { + initializeGrabber(); + processFrames(); + } catch (Exception e) { + log.error("视频流处理异常", e); + } finally { + release(); + } + }); + log.info("视频流处理已启动"); + // 初始化队列处理线程:解决回调比较耗时,导致线程池爆满 + startFrameProcessor(); + } + + /** + * 负责抓取视频帧到队列 + */ + private void processFrames() { + int frameCount = 0; + while (isRunning) { + try { + Frame frame = grabber.grab(); + if (frame == null || frame.image == null) continue; + + frameCount++; + if (frameCount % frameDetectionInterval != 0) continue; + + Frame currentFrame = frame.clone(); + frameQueue.offer(currentFrame); // 队列满则丢弃,可改为 put 阻塞 + } catch (Exception e) { + log.error("抓取视频帧异常", e); + if (e instanceof FFmpegFrameGrabber.Exception) reconnect(); + } + } + } + + /** + * 负责处理视频帧 + */ + private void startFrameProcessor() { + processorExecutor.submit(() -> { + log.info("帧处理线程已启动"); + while (isRunning || !frameQueue.isEmpty()) { + try { + Frame frame = frameQueue.poll(100, TimeUnit.MILLISECONDS); + if (frame != null) processFrame(frame); + } catch (InterruptedException e) { + e.printStackTrace(); + } catch (Exception e) { + log.error("帧处理异常", e); + } + } + }); + } + + private void processFrame(Frame frame) { + Mat mat = null; + try { + mat = converterToMat.convert(frame); + if (mat == null) return; + + Image image = ImageFactory.getInstance().fromImage(mat); + DetectedObjects detectedObjects = detectorModel.detect(image); +// log.debug("检测结果:{}", detectedObjects.toString()); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image); + List filtered = filterRepeatedObjects(detectionResponse); + if (!filtered.isEmpty() && listener != null) { + listener.onObjectDetected(filtered, image); // 同帧多物体一次回调 + } + } catch (Throwable e) { + e.printStackTrace(); + log.error("单帧处理异常", e); + } finally { + if (mat != null) mat.release(); + } + } + + private List filterRepeatedObjects(DetectionResponse response) { + List result = new ArrayList<>(); + long now = System.currentTimeMillis(); + for (DetectionInfo info : response.getDetectionInfoList()) { + String name = info.getObjectDetInfo().getClassName(); + Long last = lastDetectTime.get(name); + if (last == null || (now - last) > repeatGap * 1000) { + lastDetectTime.put(name, now); + result.add(info); + } + } + return result; + } + + private void reconnect() { + log.info("尝试重新连接视频流"); + try { + release(); + Thread.sleep(5000); + initializeGrabber(); + } catch (Exception e) { + log.error("重新连接RTSP流失败", e); + } + } + + public void stopDetection() { isRunning = false; } + + private void release() { + if (grabber != null) { + try { grabber.stop(); grabber.release(); } + catch (FrameGrabber.Exception e) { log.error("释放Grabber失败", e); } + } + } + + @Override + public void close() { + stopDetection(); + if (grabberExecutor != null) grabberExecutor.shutdownNow(); + if (processorExecutor != null) processorExecutor.shutdownNow(); + release(); + } + + public static class Builder { + private DetectorModel detectorModel; + private String streamUrl; + private ExecutorService executorService; + private int frameDetectionInterval = 1; + private StreamDetectionListener listener; + private VideoSourceType sourceType = VideoSourceType.STREAM; // 默认流 + private int cameraIndex = 0; // 默认第一个摄像头 + + public Builder detectorModel(DetectorModel m) { this.detectorModel = m; return this; } + public Builder streamUrl(String url) { this.streamUrl = url; return this; } + public Builder executorService(ExecutorService es) { this.executorService = es; return this; } + public Builder listener(StreamDetectionListener listener) { this.listener = listener; return this; } + public Builder sourceType(VideoSourceType sourceType) { + this.sourceType = sourceType; + return this; + } + public Builder cameraIndex(int cameraIndex) { + this.cameraIndex = cameraIndex; + return this; + } + public Builder frameDetectionInterval(int interval) { + if (interval < 1) throw new IllegalArgumentException("frameDetectionInterval >= 1"); + this.frameDetectionInterval = interval; + return this; + } + + public StreamDetector build() { + if (detectorModel == null) { + throw new DetectionException("detectorModel 不能为空"); + } + + if (sourceType == null) { + throw new DetectionException("sourceType 不能为空"); + } + + // 根据 sourceType 校验不同的参数 + switch (sourceType) { + case STREAM: + case FILE: + if (StringUtils.isBlank(streamUrl)) { + throw new DetectionException("streamUrl 不能为空"); + } + break; + case CAMERA: + if (cameraIndex < 0) { + throw new DetectionException("cameraIndex 必须 >= 0"); + } + break; + default: + throw new DetectionException("不支持的视频源类型: " + sourceType); + } + + if (executorService == null) { + executorService = Executors.newFixedThreadPool(2); // 至少2个线程 + } + + return new StreamDetector(this); + } + + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/translator/TensorflowTranslator.java b/vision/src/main/java/cn/smartjavaai/objectdetection/translator/TensorflowTranslator.java new file mode 100644 index 0000000..2bbced9 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/translator/TensorflowTranslator.java @@ -0,0 +1,111 @@ +package cn.smartjavaai.objectdetection.translator; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +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.types.DataType; +import ai.djl.translate.NoBatchifyTranslator; +import ai.djl.translate.TranslatorContext; + +import java.io.IOException; +import java.util.ArrayList; +import java.util.List; +import java.util.Map; +import java.util.Objects; + +/** + * @author dwj + * @date 2025/8/18 + */ +public class TensorflowTranslator implements NoBatchifyTranslator { + + private Map classes; + private int maxBoxes; + private float threshold; + + + public TensorflowTranslator(Map arguments) { + maxBoxes = + arguments.containsKey("maxBoxes") + ? Integer.parseInt(arguments.get("maxBoxes").toString()) + : 10; + threshold = + arguments.containsKey("threshold") + ? Float.parseFloat(arguments.get("threshold").toString()) + : 0.7f; + + classes = (Map)arguments.get("classes"); + } + + /** {@inheritDoc} */ + @Override + public NDList processInput(TranslatorContext ctx, Image input) { + // input to tf object-detection models is a list of tensors, hence NDList + NDArray array = input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR); + // optionally resize the image for faster processing + array = NDImageUtils.resize(array, 224); + // tf object-detection models expect 8 bit unsigned integer tensor + array = array.toType(DataType.UINT8, true); + array = array.expandDims(0); // tf object-detection models expect a 4 dimensional input + return new NDList(array); + } + + /** {@inheritDoc} */ + @Override + public void prepare(TranslatorContext ctx) throws IOException { + } + + /** {@inheritDoc} */ + @Override + public DetectedObjects processOutput(TranslatorContext ctx, NDList list) { + // output of tf object-detection models is a list of tensors, hence NDList in djl + // output NDArray order in the list are not guaranteed + + int[] classIds = null; + float[] probabilities = null; + NDArray boundingBoxes = null; + for (NDArray array : list) { + if ("detection_boxes".equals(array.getName())) { + boundingBoxes = array.get(0); + } else if ("detection_scores".equals(array.getName())) { + probabilities = array.get(0).toFloatArray(); + } else if ("detection_classes".equals(array.getName())) { + // class id is between 1 - number of classes + classIds = array.get(0).toType(DataType.INT32, true).toIntArray(); + } + } + Objects.requireNonNull(classIds); + Objects.requireNonNull(probabilities); + Objects.requireNonNull(boundingBoxes); + + List retNames = new ArrayList<>(); + List retProbs = new ArrayList<>(); + List retBB = new ArrayList<>(); + + // result are already sorted + for (int i = 0; i < Math.min(classIds.length, maxBoxes); ++i) { + int classId = classIds[i]; + double probability = probabilities[i]; + // classId starts from 1, -1 means background + if (classId > 0 && probability > threshold) { + String className = classes.getOrDefault(classId, "#" + classId); + float[] box = boundingBoxes.get(i).toFloatArray(); + float yMin = box[0]; + float xMin = box[1]; + float yMax = box[2]; + float xMax = box[3]; + Rectangle rect = new Rectangle(xMin, yMin, xMax - xMin, yMax - yMin); + retNames.add(className); + retProbs.add(probability); + retBB.add(rect); + } + } + + return new DetectedObjects(retNames, retProbs, retBB); + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/translator/YoloV8PersonDetTranslator.java b/vision/src/main/java/cn/smartjavaai/objectdetection/translator/YoloV8PersonDetTranslator.java new file mode 100644 index 0000000..3fe5393 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/translator/YoloV8PersonDetTranslator.java @@ -0,0 +1,457 @@ +package cn.smartjavaai.objectdetection.translator; + +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.modality.cv.transform.*; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; +import ai.djl.ndarray.types.Shape; +import ai.djl.translate.*; +import cn.smartjavaai.common.utils.LetterBoxUtils; + +import java.nio.file.Files; +import java.nio.file.Paths; +import java.util.*; + + +/** + * YoloV8 行人检测 Translator + */ +public class YoloV8PersonDetTranslator implements Translator { + + private int maxBoxes; + + private YoloOutputType yoloOutputLayerType; + private float nmsThreshold; + + protected float threshold; +// private BaseImageTranslator.SynsetLoader synsetLoader; + protected List classes; + protected boolean applyRatio; + protected boolean removePadding; + + protected Pipeline pipeline; + private Image.Flag flag; + private Batchifier batchifier; + protected int width; + protected int height; + + + /** + * Constructs an ImageTranslator with the provided builder. + * + * @param builder the data to build with + */ + protected YoloV8PersonDetTranslator(Builder builder) { + this.yoloOutputLayerType = builder.outputType; + this.nmsThreshold = builder.nmsThreshold; + maxBoxes = builder.maxBox; + this.threshold = builder.threshold; +// this.synsetLoader = builder.synsetLoader; + this.applyRatio = builder.applyRatio; + this.removePadding = builder.removePadding; + this.flag = builder.flag; + this.pipeline = builder.pipeline; + this.batchifier = builder.batchifier; + this.width = builder.width; + this.height = builder.height; + classes = Arrays.asList("person"); + } + + /** + * Creates a builder to build a {@code YoloV8Translator} with specified arguments. + * + * @return a new builder + */ + public static Builder builder() { + return new Builder(); + } + + /** + * Creates a builder to build a {@code YoloV8Translator} with specified arguments. + * + * @param arguments arguments to specify builder options + * @return a new builder + */ + public static Builder builder(Map arguments) { + Builder builder = new Builder(); + builder.configPreProcess(arguments); + builder.configPostProcess(arguments); + return builder; + } + + @Override + public NDList processInput(TranslatorContext ctx, Image input) throws Exception { + NDManager manager = ctx.getNDManager(); + NDArray array = input.toNDArray(manager, Image.Flag.COLOR); + //Letter box resize 640x640 with padding (保持比例,补边缘) + LetterBoxUtils.ResizeResult letterBoxResult = LetterBoxUtils.letterbox(manager, array, width, height, 114f, LetterBoxUtils.PaddingPosition.CENTER); + array = letterBoxResult.image; + // 转为 float32 且归一化到 0~1 + array = array.toType(DataType.FLOAT32, false).div(255f); // HWC + // HWC -> CHW + array = array.transpose(2, 0, 1); // CHW + + ctx.setAttachment("width", input.getWidth()); + ctx.setAttachment("height", input.getHeight()); + ctx.setAttachment("processedWidth", width); + ctx.setAttachment("processedHeight", height); + ctx.setAttachment("scale", letterBoxResult.r); + return new NDList(array); + } + + @Override + public DetectedObjects processOutput(TranslatorContext ctx, NDList list) throws Exception { + //原图宽高 + int imageWidth = (Integer) ctx.getAttachment("width"); + int imageHeight = (Integer) ctx.getAttachment("height"); + float scale = (Float) ctx.getAttachment("scale"); + switch (yoloOutputLayerType) { + case DETECT: + return processFromDetectOutput(); + case AUTO: + if (list.get(0).getShape().dimension() > 2) { + return processFromDetectOutput(); + } else { + return processFromBoxOutput(imageWidth, imageHeight, list, scale); + } + case BOX: + default: + return processFromBoxOutput(imageWidth, imageHeight, list, scale); + } + } + + /** {@inheritDoc} */ + protected DetectedObjects processFromBoxOutput(int origImageWidth, int origImageHeight, NDList list, float scale) { + + NDArray rawResult = list.get(0); + NDArray reshapedResult = rawResult.transpose(); + Shape shape = reshapedResult.getShape(); + float[] buf = reshapedResult.toFloatArray(); + int numberRows = Math.toIntExact(shape.get(0)); + int nClasses = Math.toIntExact(shape.get(1)); + int padding = nClasses - classes.size(); + System.out.println(Arrays.toString(reshapedResult.get(0).toFloatArray())); + if (padding != 0 && padding != 4) { + throw new IllegalStateException( + "Expected classes: " + (nClasses - 4) + ", got " + classes.size()); + } + + ArrayList boxes = new ArrayList<>(); + ArrayList scores = new ArrayList<>(); + ArrayList classIds = new ArrayList<>(); + + // reverse order search in heap; searches through #maxBoxes for optimization when set + for (int i = numberRows - 1; i > numberRows - maxBoxes; --i) { + int index = i * nClasses; + float maxClassProb = buf[index + 4]; + + if (maxClassProb > threshold) { + float xPos = buf[index]; // center x + float yPos = buf[index + 1]; // center y + float w = buf[index + 2]; + float h = buf[index + 3]; + Rectangle rect = + new Rectangle(Math.max(0, xPos - w / 2), Math.max(0, yPos - h / 2), w, h); + scores.add(maxClassProb); + classIds.add(0); + boxes.add(rect); + } + } + + return nms(origImageWidth, origImageHeight, boxes, classIds, scores, scale); + } + + private DetectedObjects processFromDetectOutput() { + throw new UnsupportedOperationException( + "detect layer output is not supported yet, check correct YoloV5 export format"); + } + + + + + + protected DetectedObjects nms( + int origImageWidth, + int origImageHeight, + List boxes, + List classIds, + List scores, float scale) { + List retClasses = new ArrayList<>(); + List retProbs = new ArrayList<>(); + List retBB = new ArrayList<>(); + + for (int classId = 0; classId < classes.size(); classId++) { + List r = new ArrayList<>(); + List s = new ArrayList<>(); + List map = new ArrayList<>(); + for (int j = 0; j < classIds.size(); ++j) { + if (classIds.get(j) == classId) { + r.add(boxes.get(j)); + s.add(scores.get(j).doubleValue()); + map.add(j); + } + } + if (r.isEmpty()) { + continue; + } + List nms = Rectangle.nms(r, s, nmsThreshold); + for (int index : nms) { + int pos = map.get(index); + int id = classIds.get(pos); + retClasses.add(classes.get(id)); + retProbs.add(scores.get(pos).doubleValue()); +// Rectangle rect = boxes.get(pos); + Rectangle rect = boxes.get(pos); + if (removePadding) { + rect = + LetterBoxUtils.restoreBox(rect, scale, origImageWidth, origImageHeight, width, height); + } else if (applyRatio) { + rect = + new Rectangle( + rect.getX() / width, + rect.getY() / height, + rect.getWidth() / width, + rect.getHeight() / height); + } + retBB.add(rect); + } + } + return new DetectedObjects(retClasses, retProbs, retBB); + } + + public static class Builder { + + private int maxBox = 8400; + + YoloOutputType outputType; + float nmsThreshold; + + protected float threshold = 0.2F; + protected boolean applyRatio; + protected boolean removePadding; + + protected int width = 224; + protected int height = 224; + protected Image.Flag flag; + protected Pipeline pipeline; + protected Batchifier batchifier; + + public Builder() { + this.outputType = YoloOutputType.AUTO; + this.nmsThreshold = 0.45F; + } + + public Builder optOutputType(YoloOutputType outputType) { + this.outputType = outputType; + return this; + } + + public Builder optNmsThreshold(float nmsThreshold) { + this.nmsThreshold = nmsThreshold; + return this; + } + + /** + * Builds the translator. + * + * @return the new translator + */ + public YoloV8PersonDetTranslator build() { + if (pipeline == null) { + addTransform( + array -> array.transpose(2, 0, 1).toType(DataType.FLOAT32, false).div(255)); + } +// validate(); + return new YoloV8PersonDetTranslator(this); + } + + protected Builder self() { + return this; + } + + public Builder addTransform(Transform transform) { + if (this.pipeline == null) { + this.pipeline = new Pipeline(); + } + + this.pipeline.add(transform); + return this.self(); + } + + public Builder optApplyRatio(boolean value) { + this.applyRatio = value; + return this.self(); + } + + public Builder optFlag(Image.Flag flag) { + this.flag = flag; + return this.self(); + } + + public Builder setPipeline(Pipeline pipeline) { + this.pipeline = pipeline; + return this.self(); + } + + public Builder setImageSize(int width, int height) { + this.width = width; + this.height = height; + return this.self(); + } + + + public Builder optBatchifier(Batchifier batchifier) { + this.batchifier = batchifier; + return this.self(); + } + + public Builder optThreshold(float threshold) { + this.threshold = threshold; + return this.self(); + } + + /** {@inheritDoc} */ + protected void configPostProcess(Map arguments) { + if (ArgumentsUtil.booleanValue(arguments, "optApplyRatio") || ArgumentsUtil.booleanValue(arguments, "applyRatio")) { + this.optApplyRatio(true); + } + this.threshold = ArgumentsUtil.floatValue(arguments, "threshold", 0.2F); + String centerFit = ArgumentsUtil.stringValue(arguments, "centerFit", "false"); + this.removePadding = "true".equals(centerFit); + String type = ArgumentsUtil.stringValue(arguments, "outputType", "AUTO"); + this.outputType = YoloOutputType.valueOf(type.toUpperCase(Locale.ENGLISH)); + this.nmsThreshold = ArgumentsUtil.floatValue(arguments, "nmsThreshold", 0.45F); + maxBox = ArgumentsUtil.intValue(arguments, "maxBox", 8400); + } + + protected void configPreProcess(Map arguments) { + if (this.pipeline == null) { + this.pipeline = new Pipeline(); + } + + this.width = ArgumentsUtil.intValue(arguments, "width", 224); + this.height = ArgumentsUtil.intValue(arguments, "height", 224); + if (arguments.containsKey("flag")) { + this.flag = Image.Flag.valueOf(arguments.get("flag").toString()); + } + + String pad = ArgumentsUtil.stringValue(arguments, "pad", "false"); + if ("true".equals(pad)) { + this.addTransform(new Pad(0.0)); + } else if (!"false".equals(pad)) { + double padding = Double.parseDouble(pad); + this.addTransform(new Pad(padding)); + } + + String resize = ArgumentsUtil.stringValue(arguments, "resize", "false"); + int w; + int shortEdge; + if ("true".equals(resize)) { + this.addTransform(new Resize(this.width, this.height)); + } else if (!"false".equals(resize)) { + String[] tokens = resize.split("\\s*,\\s*"); + w = (int)Double.parseDouble(tokens[0]); + if (tokens.length > 1) { + shortEdge = (int)Double.parseDouble(tokens[1]); + } else { + shortEdge = w; + } + + Image.Interpolation interpolation; + if (tokens.length > 2) { + interpolation = Image.Interpolation.valueOf(tokens[2]); + } else { + interpolation = Image.Interpolation.BILINEAR; + } + + this.addTransform(new Resize(w, shortEdge, interpolation)); + } + + String resizeShort = ArgumentsUtil.stringValue(arguments, "resizeShort", "false"); + if ("true".equals(resizeShort)) { + w = Math.max(this.width, this.height); + this.addTransform(new ResizeShort(w)); + } else if (!"false".equals(resizeShort)) { + String[] tokens = resizeShort.split("\\s*,\\s*"); + shortEdge = (int)Double.parseDouble(tokens[0]); + int longEdge; + if (tokens.length > 1) { + longEdge = (int)Double.parseDouble(tokens[1]); + } else { + longEdge = -1; + } + + Image.Interpolation interpolation; + if (tokens.length > 2) { + interpolation = Image.Interpolation.valueOf(tokens[2]); + } else { + interpolation = Image.Interpolation.BILINEAR; + } + + this.addTransform(new ResizeShort(shortEdge, longEdge, interpolation)); + } + + if (ArgumentsUtil.booleanValue(arguments, "centerCrop", false)) { + this.addTransform(new CenterCrop(this.width, this.height)); + } + + if (ArgumentsUtil.booleanValue(arguments, "centerFit")) { + this.addTransform(new CenterFit(this.width, this.height)); + } + + if (ArgumentsUtil.booleanValue(arguments, "toTensor", true)) { + this.addTransform(new ToTensor()); + } + + String normalize = ArgumentsUtil.stringValue(arguments, "normalize", "false"); + if ("true".equals(normalize)) { + float[] MEAN = new float[]{0.485F, 0.456F, 0.406F}; + float[] STD = new float[]{0.229F, 0.224F, 0.225F}; + this.addTransform(new Normalize(MEAN, STD)); + } else if (!"false".equals(normalize)) { + String[] tokens = normalize.split("\\s*,\\s*"); + if (tokens.length != 6) { + throw new IllegalArgumentException("Invalid normalize value: " + normalize); + } + + float[] mean = new float[]{Float.parseFloat(tokens[0]), Float.parseFloat(tokens[1]), Float.parseFloat(tokens[2])}; + float[] std = new float[]{Float.parseFloat(tokens[3]), Float.parseFloat(tokens[4]), Float.parseFloat(tokens[5])}; + this.addTransform(new Normalize(mean, std)); + } + + String range = (String)arguments.get("range"); + if ("0,1".equals(range)) { + this.addTransform((a) -> { + return a.div(255.0F); + }); + } else if ("-1,1".equals(range)) { + this.addTransform((a) -> { + return a.div(128.0F).sub(1); + }); + } + + if (arguments.containsKey("batchifier")) { + this.batchifier = Batchifier.fromString((String)arguments.get("batchifier")); + } + + } + } + + public static enum YoloOutputType { + BOX, + DETECT, + AUTO; + + private YoloOutputType() { + } + } + + + +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/config/SemSegModelConfig.java b/vision/src/main/java/cn/smartjavaai/semseg/config/SemSegModelConfig.java new file mode 100644 index 0000000..27319f4 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/config/SemSegModelConfig.java @@ -0,0 +1,49 @@ +package cn.smartjavaai.semseg.config; + +import cn.smartjavaai.common.config.ModelConfig; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.semseg.enums.SemSegModelEnum; +import lombok.Data; + +import java.util.List; + +/** + * 语义分割模型参数配置 + * + * @author dwj + */ +@Data +public class SemSegModelConfig extends ModelConfig { + + /** + * 模型 + */ + private SemSegModelEnum modelEnum; + + + /** + * 模型路径 + */ + private String modelPath; + + + /** + * 允许的分类列表 + */ + private List allowedClasses; + + + public SemSegModelConfig() { + } + + public SemSegModelConfig(SemSegModelEnum modelEnum, DeviceEnum device) { + this.modelEnum = modelEnum; + setDevice(device); + } + + public SemSegModelConfig(SemSegModelEnum modelEnum) { + this.modelEnum = modelEnum; + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/criteria/SemSegCriteriaFactory.java b/vision/src/main/java/cn/smartjavaai/semseg/criteria/SemSegCriteriaFactory.java new file mode 100644 index 0000000..00262a1 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/criteria/SemSegCriteriaFactory.java @@ -0,0 +1,48 @@ +package cn.smartjavaai.semseg.criteria; + +import ai.djl.Device; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.CategoryMask; +import ai.djl.modality.cv.translator.SemanticSegmentationTranslatorFactory; +import ai.djl.repository.zoo.Criteria; +import ai.djl.training.util.ProgressBar; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.semseg.config.SemSegModelConfig; +import cn.smartjavaai.semseg.enums.SemSegModelEnum; +import org.apache.commons.lang3.StringUtils; + +import java.nio.file.Paths; +import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; + +/** + * 语义分割Criteria工厂 + * @author dwj + */ +public class SemSegCriteriaFactory { + + + public static Criteria createCriteria(SemSegModelConfig config) { + Device device = null; + if(!Objects.isNull(config.getDevice())){ + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); + } + Criteria criteria = null; + ConcurrentHashMap params = new ConcurrentHashMap(); + params.putAll(config.getCustomParams()); + if(config.getModelEnum() == SemSegModelEnum.DEEPLABV3){ + criteria = + Criteria.builder() + .setTypes(Image.class, CategoryMask.class) + .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : + config.getModelEnum().getModelUri()) + .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) + .optTranslatorFactory(new SemanticSegmentationTranslatorFactory()) + .optEngine("PyTorch") + .optDevice(device) + .optProgress(new ProgressBar()) + .build(); + } + return criteria; + } +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/entity/DetectParams.java b/vision/src/main/java/cn/smartjavaai/semseg/entity/DetectParams.java new file mode 100644 index 0000000..cd59af8 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/entity/DetectParams.java @@ -0,0 +1,17 @@ +package cn.smartjavaai.semseg.entity; + +import lombok.Data; + +/** + * 检测参数 + * @author dwj + */ +@Data +public class DetectParams { + + /** + * 置信度阈值 + */ + private float threshold = 0.3f; + +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/enums/SemSegModelEnum.java b/vision/src/main/java/cn/smartjavaai/semseg/enums/SemSegModelEnum.java new file mode 100644 index 0000000..b5ed5bc --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/enums/SemSegModelEnum.java @@ -0,0 +1,34 @@ +package cn.smartjavaai.semseg.enums; + +/** + * 语义分割模型枚举 + * @author dwj + */ +public enum SemSegModelEnum { + + DEEPLABV3("djl://ai.djl.pytorch/deeplabv3"); + + /** + * 根据名称获取枚举 (忽略大小写和下划线变体) + */ + public static SemSegModelEnum fromName(String name) { + String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); + for (SemSegModelEnum model : values()) { + if (model.name().replaceAll("_", "").equals(formatted)) { + return model; + } + } + throw new IllegalArgumentException("未知模型名称: " + name); + } + + private final String modelUri; + + SemSegModelEnum(String modelUri) { + this.modelUri = modelUri; + } + + public String getModelUri() { + return modelUri; + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/exception/SemSegException.java b/vision/src/main/java/cn/smartjavaai/semseg/exception/SemSegException.java new file mode 100644 index 0000000..02dcfe4 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/exception/SemSegException.java @@ -0,0 +1,29 @@ +package cn.smartjavaai.semseg.exception; + +/** + * 语义分割异常 + * @author dwj + */ +public class SemSegException extends RuntimeException{ + + public SemSegException() { + super(); + } + + public SemSegException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) { + super(message, cause, enableSuppression, writableStackTrace); + } + + public SemSegException(String message, Throwable cause) { + super(message, cause); + } + + public SemSegException(String message) { + super(message); + } + + public SemSegException(Throwable cause) { + super(cause); + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java b/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java new file mode 100644 index 0000000..a392f54 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java @@ -0,0 +1,137 @@ +package cn.smartjavaai.semseg.model; + +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.CategoryMask; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ModelNotFoundException; +import ai.djl.repository.zoo.ZooModel; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.Base64ImageUtils; +import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.semseg.config.SemSegModelConfig; +import cn.smartjavaai.semseg.criteria.SemSegCriteriaFactory; +import cn.smartjavaai.vision.utils.CategoryMaskFilter; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.collections.CollectionUtils; +import org.apache.commons.lang3.StringUtils; +import org.apache.commons.pool2.impl.GenericObjectPool; + +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.util.Objects; + +/** + * 语义分割模型 + * @author dwj + */ +@Slf4j +public class CommonSemSegModel implements SemSegModel { + + + private SemSegModelConfig config; + + private ZooModel model; + + private GenericObjectPool> predictorPool; + + @Override + public void loadModel(SemSegModelConfig config) { + if(Objects.isNull(config.getModelEnum())){ + throw new DetectionException("未配置模型枚举"); + } + Criteria criteria = SemSegCriteriaFactory.createCriteria(config); + this.config = config; + try { + model = criteria.loadModel(); + // 创建池子:每个线程独享 Predictor + this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model)); + int predictorPoolSize = config.getPredictorPoolSize(); + if(config.getPredictorPoolSize() <= 0){ + predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数 + } + predictorPool.setMaxTotal(predictorPoolSize); + log.debug("当前设备: " + model.getNDManager().getDevice()); + log.debug("当前引擎: " + Engine.getInstance().getEngineName()); + log.debug("模型推理器线程池最大数量: " + predictorPoolSize); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new DetectionException("模型加载失败", e); + } + } + + @Override + public R detectBase64(String base64Image) { + if(StringUtils.isBlank(base64Image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + try { + byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); + Image image = ImageFactory.getInstance().fromInputStream(new ByteArrayInputStream(imageData)); + return detect(image); + } catch (IOException e) { + throw new DetectionException("读取图片异常", e); + } + } + + @Override + public R detect(Image image) { + CategoryMask categoryMask = detectCore(image); + // 过滤 + if(CollectionUtils.isNotEmpty(config.getAllowedClasses()) + && Objects.nonNull(categoryMask) && !categoryMask.getClasses().isEmpty()){ + categoryMask = new CategoryMaskFilter(config.getAllowedClasses()).filter(categoryMask); + } + return R.ok(categoryMask); + } + + /** + * 模型核心推理方法 + * @param image + * @return + */ + public CategoryMask detectCore(Image image) { + Predictor predictor = null; + try { + predictor = predictorPool.borrowObject(); + return predictor.predict(image); + } catch (Exception e) { + throw new DetectionException("语义分割错误", e); + }finally { + if (predictor != null) { + try { + predictorPool.returnObject(predictor); //归还 + log.debug("释放资源"); + } 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); + } + } +} diff --git a/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java b/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java new file mode 100644 index 0000000..dfac361 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java @@ -0,0 +1,42 @@ +package cn.smartjavaai.semseg.model; + +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.CategoryMask; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.semseg.config.SemSegModelConfig; + +/** + * 语义分割模型 + * @author dwj + */ +public interface SemSegModel extends AutoCloseable{ + + + /** + * 加载模型 + * @param config + */ + void loadModel(SemSegModelConfig config); + + /** + * 语义分割 + * @param base64Image + * @return + */ + default R detectBase64(String base64Image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 语义分割 + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/CategoryMaskFilter.java b/vision/src/main/java/cn/smartjavaai/vision/utils/CategoryMaskFilter.java new file mode 100644 index 0000000..e48d7cb --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/CategoryMaskFilter.java @@ -0,0 +1,102 @@ +package cn.smartjavaai.vision.utils; + +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.output.CategoryMask; +import lombok.Data; + +import java.util.ArrayList; +import java.util.List; + +/** + * CategoryMask过滤器 + * @author dwj + */ +@Data +public class CategoryMaskFilter { + + private List allowedClasses; + + /** + * 构造函数,初始化允许的分类列表和置信度阈值 + * + * @param allowedClasses 允许的分类列表,null表示不过滤分类 + */ + public CategoryMaskFilter(List allowedClasses) { + this.allowedClasses = allowedClasses != null ? new ArrayList<>(allowedClasses) : null; + } + + + /** + * 对给定的CategoryMask进行过滤,只保留允许的分类。 + * + * @param categoryMask 待过滤的CategoryMask + * @return 过滤后的CategoryMask + */ + public CategoryMask filter(CategoryMask categoryMask) { + if (categoryMask == null) { + throw new NullPointerException("Input categoryMask cannot be null"); + } + + // If allowedClasses is null, return a copy of the original CategoryMask + if (allowedClasses == null) { + List classNamesCopy = new ArrayList<>(categoryMask.getClasses()); + int[][] maskCopy = copyMask(categoryMask.getMask()); + return new CategoryMask(classNamesCopy, maskCopy); + } + + // Get original class names and mask + List originalClassNames = categoryMask.getClasses(); + int[][] originalMask = categoryMask.getMask(); + + // Create a new class names list for allowed classes + List filteredClassNames = new ArrayList<>(); + + // Map original indices to new indices for allowed classes + List allowedOriginalIndices = new ArrayList<>(); + for (int i = 0; i < originalClassNames.size(); i++) { + String className = originalClassNames.get(i); + if (allowedClasses.contains(className)) { + filteredClassNames.add(className); + allowedOriginalIndices.add(i); + } + } + + // Create a new mask with the same dimensions + int height = originalMask.length; + int width = (height > 0) ? originalMask[0].length : 0; + int[][] filteredMask = new int[height][width]; + + // Update mask: remap pixels of allowed classes to new indices, others to 0 + for (int h = 0; h < height; h++) { + for (int w = 0; w < width; w++) { + int originalIndex = originalMask[h][w]; + int newIndex = allowedOriginalIndices.indexOf(originalIndex); + if (newIndex != -1) { + filteredMask[h][w] = newIndex; + } else { + filteredMask[h][w] = 0; // Background + } + } + } + + return new CategoryMask(filteredClassNames, filteredMask); + } + + /** + * Helper method to copy the mask array. + * + * @param originalMask the original mask to copy + * @return a deep copy of the mask + */ + private int[][] copyMask(int[][] originalMask) { + int height = originalMask.length; + int width = (height > 0) ? originalMask[0].length : 0; + int[][] copy = new int[height][width]; + for (int h = 0; h < height; h++) { + System.arraycopy(originalMask[h], 0, copy[h], 0, width); + } + return copy; + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/ClassificationFilter.java b/vision/src/main/java/cn/smartjavaai/vision/utils/ClassificationFilter.java new file mode 100644 index 0000000..2ed0528 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/ClassificationFilter.java @@ -0,0 +1,62 @@ +package cn.smartjavaai.vision.utils; + +import ai.djl.modality.Classifications; +import lombok.Data; + +import java.util.ArrayList; +import java.util.List; +import java.util.stream.Collectors; + +/** + * Classification过滤器 + * @author dwj + */ +@Data +public class ClassificationFilter { + + private List allowedClasses; + private float threshold; + + /** + * 构造函数,初始化允许的分类列表和置信度阈值 + * + * @param allowedClasses 允许的分类列表,null表示不过滤分类 + * @param threshold 置信度阈值 + */ + public ClassificationFilter(List allowedClasses, float threshold) { + this.allowedClasses = allowedClasses != null ? new ArrayList<>(allowedClasses) : null; + this.threshold = threshold; + } + + /** + * 过滤分类结果并返回新的Classifications对象 + * + * @param classifications 原始分类结果 + * @return 过滤后的Classifications对象 + */ + public Classifications filter(Classifications classifications) { + // 获取原始分类和概率 + List originalClasses = classifications.getClassNames(); + List originalProbabilities = classifications.getProbabilities(); + + // 过滤分类和概率 + List filteredClasses = new ArrayList<>(); + List filteredProbabilities = new ArrayList<>(); + + for (int i = 0; i < originalClasses.size(); i++) { + String className = originalClasses.get(i); + double probability = originalProbabilities.get(i); + + // 检查是否满足置信度阈值和允许的分类列表 + if (probability >= threshold && (allowedClasses == null || allowedClasses.contains(className))) { + filteredClasses.add(className); + filteredProbabilities.add(probability); + } + } + + // 创建新的Classifications对象 + return new Classifications(filteredClasses, filteredProbabilities); + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/DetectedObjectsFilter.java b/vision/src/main/java/cn/smartjavaai/vision/utils/DetectedObjectsFilter.java new file mode 100644 index 0000000..3de3f6d --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/DetectedObjectsFilter.java @@ -0,0 +1,84 @@ +package cn.smartjavaai.vision.utils; + +import ai.djl.modality.Classifications; +import ai.djl.modality.cv.output.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +import lombok.Data; +import org.apache.commons.collections.CollectionUtils; + +import java.util.ArrayList; +import java.util.Comparator; +import java.util.List; +import java.util.Objects; + +/** + * Classification过滤器 + * @author dwj + */ +@Data +public class DetectedObjectsFilter { + + private List allowedClasses; + private float threshold; + private int topK; + + /** + * 构造函数,初始化允许的分类列表和置信度阈值 + * + * @param allowedClasses 允许的分类列表,null表示不过滤分类 + * @param threshold 置信度阈值 + */ + public DetectedObjectsFilter(List allowedClasses, float threshold) { + this.allowedClasses = allowedClasses != null ? new ArrayList<>(allowedClasses) : null; + this.threshold = threshold; + } + + public DetectedObjectsFilter(List allowedClasses, float threshold, int topK) { + this.allowedClasses = allowedClasses; + this.threshold = threshold; + this.topK = topK; + } + + /** + * 对给定的 DetectedObjects 进行过滤,返回过滤后的结果 + * + * @param detectedObjects 待过滤的 DetectedObjects 对象 + * @return 过滤后的 DetectedObjects 对象 + */ + public DetectedObjects filter(DetectedObjects detectedObjects) { + if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0) { + return detectedObjects; + } + List items = detectedObjects.items(); + // 按照允许的类别进行过滤 + List filtered = new ArrayList<>(); + //过滤类别 + if(!CollectionUtils.isEmpty(allowedClasses)){ + for (DetectedObjects.DetectedObject obj : items) { + if(allowedClasses.contains(obj.getClassName())){ + filtered.add(obj); + } + } + }else{ + filtered = items; + } + if(topK > 0 && filtered.size() > topK){ + // 按照概率进行排序 + filtered.sort(Comparator.comparingDouble(Classifications.Classification::getProbability).reversed()); + filtered = filtered.subList(0, topK); + } + // 构建新的 DetectedObjects 返回 + List names = new ArrayList<>(); + List probs = new ArrayList<>(); + List boxes = new ArrayList<>(); + + for (DetectedObjects.DetectedObject obj : filtered) { + names.add(obj.getClassName()); + probs.add(obj.getProbability()); + boxes.add(obj.getBoundingBox()); + } + return new DetectedObjects(names, probs, boxes); + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java b/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java new file mode 100644 index 0000000..447b786 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java @@ -0,0 +1,114 @@ +package cn.smartjavaai.vision.utils; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.output.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.output.Mask; +import ai.djl.modality.cv.output.Rectangle; +import cn.smartjavaai.common.entity.*; +import cn.smartjavaai.common.utils.ImageUtils; +import cn.smartjavaai.obb.entity.ObbResult; +import cn.smartjavaai.obb.entity.YoloRotatedBox; +import org.apache.commons.collections.CollectionUtils; +import org.opencv.core.Mat; +import org.opencv.core.Scalar; +import org.opencv.imgproc.Imgproc; + +import javax.imageio.ImageIO; +import java.awt.image.BufferedImage; +import java.io.File; +import java.io.IOException; +import java.util.ArrayList; +import java.util.Iterator; +import java.util.List; +import java.util.Objects; + +/** + * 目标检测相关工具类 + * @author dwj + * @date 2025/4/9 + */ +public class DetectorUtils { + + + /** + * 转换为FaceDetectedResult + * @param detection + * @param img + * @return + */ + public static DetectionResponse convertToDetectionResponse(DetectedObjects detection, Image img){ + if(Objects.isNull(detection) || Objects.isNull(detection.getProbabilities()) + || detection.getProbabilities().isEmpty() || Objects.isNull(detection.items()) || detection.items().isEmpty()){ + return null; + } + DetectionResponse detectionResponse = new DetectionResponse(); + List detectionInfoList = new ArrayList(); + List detectedObjectList = detection.items(); + Iterator iterator = detectedObjectList.iterator(); + int index = 0; + while(iterator.hasNext()) { + DetectedObjects.DetectedObject result = (DetectedObjects.DetectedObject)iterator.next(); + String className = result.getClassName(); + BoundingBox box = result.getBoundingBox(); + int x = (int)(box.getBounds().getX() * (double)img.getWidth()); + int y = (int)(box.getBounds().getY() * (double)img.getHeight()); + int width = (int)(box.getBounds().getWidth() * (double)img.getWidth()); + int height = (int)(box.getBounds().getHeight() * (double)img.getHeight()); + DetectionRectangle rectangle = new DetectionRectangle(x, y, width, height); + DetectionInfo detectionInfo = new DetectionInfo(rectangle, detection.getProbabilities().get(index).floatValue()); + //目标检测 + if(box instanceof Rectangle){ + ObjectDetInfo objectDetInfo = new ObjectDetInfo(className); + detectionInfo.setObjectDetInfo(objectDetInfo); + }else if(box instanceof Mask){ + Mask mask = (Mask)box; + InstanceSegInfo instanceSegInfo = new InstanceSegInfo(className, mask.getProbDist()); + detectionInfo.setInstanceSegInfo(instanceSegInfo); + } + detectionInfoList.add(detectionInfo); + index++; + } + detectionResponse.setDetectionInfoList(detectionInfoList); + return detectionResponse; + } + + public static DetectionResponse obbToToDetectionResponse(ObbResult obbResult){ + if(Objects.isNull(obbResult) || CollectionUtils.isEmpty(obbResult.getRotatedBoxeList())){ + return null; + } + DetectionResponse detectionResponse = new DetectionResponse(); + List detectionInfoList = new ArrayList(); + for (YoloRotatedBox box : obbResult.getRotatedBoxeList()){ + List points = box.toPoints(); + RotatedBox rotatedBox = new RotatedBox(points.get(0), points.get(1), points.get(2), points.get(3)); + ObbDetInfo obbDetInfo = new ObbDetInfo(box.className, rotatedBox); + DetectionInfo detectionInfo = new DetectionInfo(); + detectionInfo.setScore(box.score); + detectionInfo.setObbDetInfo(obbDetInfo); + detectionInfoList.add(detectionInfo); + } + detectionResponse.setDetectionInfoList(detectionInfoList); + return detectionResponse; + } + + + /** + * 绘制文本框及文本 + * @param srcMat + * @param rotatedBoxeList + */ + public static void drawRectWithText(Mat srcMat, List rotatedBoxeList) { + for(YoloRotatedBox box : rotatedBoxeList){ + List points = box.toPoints(); + Imgproc.line(srcMat, points.get(0).toCvPoint(), points.get(1).toCvPoint(), new Scalar(0, 255, 0), 1); + Imgproc.line(srcMat, points.get(1).toCvPoint(), points.get(2).toCvPoint(), new Scalar(0, 255, 0),1); + Imgproc.line(srcMat, points.get(2).toCvPoint(), points.get(3).toCvPoint(), new Scalar(0, 255, 0),1); + Imgproc.line(srcMat, points.get(3).toCvPoint(), points.get(1).toCvPoint(), new Scalar(0, 255, 0), 1); + // 中文乱码 + Imgproc.putText(srcMat, box.className, points.get(0).toCvPoint(), Imgproc.FONT_HERSHEY_SCRIPT_SIMPLEX, 1.0, new Scalar(0, 255, 0), 1); + } + } + + +} diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/ObbResultFilter.java b/vision/src/main/java/cn/smartjavaai/vision/utils/ObbResultFilter.java new file mode 100644 index 0000000..298d139 --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/ObbResultFilter.java @@ -0,0 +1,45 @@ +package cn.smartjavaai.vision.utils; + +import cn.smartjavaai.obb.entity.ObbResult; +import cn.smartjavaai.obb.entity.YoloRotatedBox; + +import java.util.Collections; +import java.util.List; +import java.util.stream.Collectors; + +/** + * @author dwj + * @date 2025/8/24 + */ +public class ObbResultFilter { + + private List allowedClasses; + private int topK; + + public ObbResultFilter(List allowedClasses, int topK) { + this.allowedClasses = allowedClasses; + this.topK = topK; + } + + public ObbResult filter(ObbResult result) { + if (result == null || result.getRotatedBoxeList() == null) { + return new ObbResult(Collections.emptyList()); + } + + // 1. 按类别过滤 + List filtered = result.getRotatedBoxeList().stream() + .filter(box -> allowedClasses.contains(box.className)) + .collect(Collectors.toList()); + + // 2. 按 score 排序 +// filtered.sort((a, b) -> Float.compare(b.score, a.score)); + + // 3. 取前 topK + if (topK > 0 && filtered.size() > topK) { + filtered = filtered.subList(0, topK); + } + + return new ObbResult(filtered); + } + +} diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/TensorflowSynsetUtils.java b/vision/src/main/java/cn/smartjavaai/vision/utils/TensorflowSynsetUtils.java new file mode 100644 index 0000000..b97875d --- /dev/null +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/TensorflowSynsetUtils.java @@ -0,0 +1,78 @@ +package cn.smartjavaai.vision.utils; + +import ai.djl.util.JsonUtils; +import com.google.gson.annotations.SerializedName; +import org.apache.commons.io.FileUtils; + +import java.io.BufferedInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.net.URL; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.Map; +import java.util.Scanner; +import java.util.concurrent.ConcurrentHashMap; + +/** + * Tensorflow synset utils + * @author dwj + */ +public class TensorflowSynsetUtils { + + /** + * 加载synset + * @param synsetUrl + * @return + * @throws IOException + */ + public static Map loadSynset(URL synsetUrl) throws IOException { + return loadSynset(synsetUrl.openStream()); + } + + + /** + * 加载synset + * @param synsetPath + * @return + * @throws IOException + */ + public static Map loadSynset(Path synsetPath) throws IOException { + return loadSynset(Files.newInputStream(synsetPath)); + } + + + /** + * 加载synset + * @param inputStream + * @return + * @throws IOException + */ + public static Map loadSynset(InputStream inputStream) throws IOException { + Map map = new ConcurrentHashMap<>(); + int maxId = 0; + try (InputStream is = new BufferedInputStream(inputStream); + Scanner scanner = new Scanner(is, StandardCharsets.UTF_8.name())) { + scanner.useDelimiter("item "); + while (scanner.hasNext()) { + String content = scanner.next(); + content = content.replaceAll("(\"|\\d)\\n\\s", "$1,"); + Item item = JsonUtils.GSON.fromJson(content, Item.class); + map.put(item.id, item.displayName); + if (item.id > maxId) { + maxId = item.id; + } + } + } + return map; + } + + private static final class Item { + int id; + + @SerializedName("display_name") + String displayName; + } +}