diff --git a/README.md b/README.md index 97af8a4..d42be81 100644 --- a/README.md +++ b/README.md @@ -426,6 +426,7 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的 ### 2、Maven 在项目的 `pom.xml` 的 `dependencies` 中可以一次性引入全部功能(如下所示)。 + ⚠️ **注意:不推荐直接引入全部依赖**,更推荐根据实际需求,按功能模块单独引入,避免引入不必要的包。 详细引入方式请查看 [文档](http://doc.smartjavaai.cn/install.html)、或查看[示例代码](https://gitee.com/dengwenjie/SmartJavaAI/tree/master/examples) @@ -434,7 +435,7 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的 cn.smartjavaai smartjavaai-all - 1.0.24 + 1.0.25 ``` @@ -577,6 +578,7 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的 | YOLOV8-SEG | OnnxRuntime | Ultralytics在COCO 数据集 上训练的模型 | [Github](https://docs.ultralytics.com/zh/tasks/segment/) | | YOLOV11-SEG | OnnxRuntime | Ultralytics在COCO 数据集 上训练的模型 | [Github](https://docs.ultralytics.com/zh/tasks/segment/) | | Mask R-CNN | MXNet | Mask R-CNN 是一种在目标检测基础上,同时为每个物体生成像素级分割区域的深度学习模型 | 无 | + --- #### OBB旋转框目标检测模型 @@ -733,6 +735,18 @@ SmartJavaAI是专为JAVA 开发者打造的一个功能丰富、开箱即用的 ## 近期更新日志 +## [v1.0.25] - 2025-10-02 +- 【人脸检测】新增6个模型(MTCNN、YOLOV5、RetinaFace小尺寸版),大幅提升性能 +- 【人脸识别】新增Seetaface6轻量模型 +- 【目标检测】支持视频流目标检测(rtsp、视频文件等) +- 【目标检测】支持tensorflow2目标检测模型 +- 【目标检测】新增行人检测模型(yolo-person) +- 【通用视觉】新增4个动作识别模型 +- 【通用视觉】新增语义分割模型 +- 【通用视觉】新增5个实例分割模型(含yolov8-seg、yolov11-seg) +- 【通用视觉】新增yolo-obb11旋转框检测(含yolov11-obb) +- 【通用视觉】新增5个姿态估计模型(含yolov8-pose、yolov11-pose) + ## [v1.0.24] - 2025-09-07 - 【人脸检测】新增6个模型(MTCNN、YOLOV5、RetinaFace小尺寸版),大幅提升性能 - 【人脸识别】新增Seetaface6轻量模型 diff --git a/all/pom.xml b/all/pom.xml index 0e06207..5bb04f3 100644 --- a/all/pom.xml +++ b/all/pom.xml @@ -6,11 +6,11 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 all - 1.0.24 + 1.0.25 ${project.artifactId} SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/bom/pom.xml b/bom/pom.xml index c73693c..697dc70 100644 --- a/bom/pom.xml +++ b/bom/pom.xml @@ -6,12 +6,12 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 - 1.0.24 + 1.0.25 bom - sbom + bom 统一版本管理的 BOM 包,同时支持 import 和全量依赖 diff --git a/common/pom.xml b/common/pom.xml index 24a57e4..06eabfa 100644 --- a/common/pom.xml +++ b/common/pom.xml @@ -6,7 +6,7 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 common diff --git a/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java b/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java index 9f456a1..2d0a184 100644 --- a/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java +++ b/common/src/main/java/cn/smartjavaai/common/cv/SmartImageFactory.java @@ -21,45 +21,132 @@ import org.opencv.imgproc.Imgproc; import java.awt.image.BufferedImage; import java.io.ByteArrayInputStream; +import java.io.File; import java.io.IOException; import java.io.InputStream; import java.nio.ByteBuffer; import java.nio.IntBuffer; import java.nio.file.Path; +import java.nio.file.Paths; /** * 图片处理工厂类 * @author dwj */ -public class SmartImageFactory extends BufferedImageFactory { +public class SmartImageFactory { + public enum Engine { + BUFFEREDIMAGE, + OPENCV + } + + private static volatile Engine currentEngine = Engine.BUFFEREDIMAGE; private static volatile SmartImageFactory instance; - public static SmartImageFactory newInstance() { + public static synchronized void setEngine(Engine engine) { + if (engine == null || engine == currentEngine) { + return; + } + currentEngine = engine; + // 只在切换时注册全局 + switch (currentEngine) { + case OPENCV: + ImageFactory.setImageFactory(new OpenCVImageFactory()); + break; + case BUFFEREDIMAGE: + default: + ImageFactory.setImageFactory(new BufferedImageFactory()); + } + } + + public static synchronized SmartImageFactory getInstance() { if (instance == null) { - synchronized (SmartImageFactory.class) { - if (instance == null) { - instance = new SmartImageFactory(); - } + instance = new SmartImageFactory(); + // 初始化全局 Engine + switch (currentEngine) { + case OPENCV: + ImageFactory.setImageFactory(new OpenCVImageFactory()); + break; + case BUFFEREDIMAGE: + default: + ImageFactory.setImageFactory(new BufferedImageFactory()); } } return instance; } - public static SmartImageFactory getInstance(){ - return newInstance(); - } public Image fromBufferedImage(BufferedImage sourceImage){ - return fromImage(OpenCVUtils.image2Mat(sourceImage)); + if (sourceImage == null) { + throw new IllegalArgumentException("BufferedImage 不能为空"); + } + Image image = null; + switch (currentEngine) { + case BUFFEREDIMAGE: + image = ImageFactory.getInstance().fromImage(sourceImage); + break; + case OPENCV: + // 先转 Mat + Mat mat = OpenCVUtils.image2Mat(sourceImage); + image = ImageFactory.getInstance().fromImage(mat); + break; + default: + throw new IllegalStateException("未知 Engine: " + currentEngine); + } + return image; + } + + public Image fromMat(Mat mat){ + if (mat == null) { + throw new IllegalArgumentException("mat 不能为空"); + } + Image image = null; + switch (currentEngine) { + case OPENCV: + image = ImageFactory.getInstance().fromImage(mat); + break; + case BUFFEREDIMAGE: + // 先转 Mat + BufferedImage sourceImage = OpenCVUtils.mat2Image(mat); + image = ImageFactory.getInstance().fromImage(sourceImage); + break; + default: + throw new IllegalStateException("未知 Engine: " + currentEngine); + } + return image; } public Image fromBase64(String base64Image) throws IOException { - return fromUrl(base64Image); + return ImageFactory.getInstance().fromUrl(base64Image); } - public Image fromBytes(byte[] imageData){ - return fromImage(new ByteArrayInputStream(imageData)); + public Image fromBytes(byte[] imageData) throws IOException { + return ImageFactory.getInstance().fromInputStream(new ByteArrayInputStream(imageData)); } + public Image fromFile(File file) throws IOException { + return ImageFactory.getInstance().fromFile(file.toPath()); + } + + public Image fromFile(Path path) throws IOException { + return ImageFactory.getInstance().fromFile(path); + } + + public Image fromFile(String filePath) throws IOException { + if (filePath == null || filePath.trim().isEmpty()) { + throw new IllegalArgumentException("filePath 不能为空"); + } + return fromFile(Paths.get(filePath)); + } + + public Image fromPixels(int[] pixels, int width, int height){ + return ImageFactory.getInstance().fromPixels(pixels, width, height); + } + + public Image fromInputStream(InputStream inputStream) throws IOException { + return ImageFactory.getInstance().fromInputStream(inputStream); + } + + + } diff --git a/common/src/main/java/cn/smartjavaai/common/entity/PolygonLabel.java b/common/src/main/java/cn/smartjavaai/common/entity/PolygonLabel.java new file mode 100644 index 0000000..24d07a7 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/entity/PolygonLabel.java @@ -0,0 +1,28 @@ +package cn.smartjavaai.common.entity; + +import lombok.Data; + +import java.util.List; + +/** + * 多边形 + * @author dwj + */ +@Data +public class PolygonLabel { + + private List points; + private String text; + + public PolygonLabel(List points, String text) { + this.points = points; + this.text = text; + } + + public PolygonLabel() { + } + + public PolygonLabel(List points) { + this.points = points; + } +} diff --git a/common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java b/common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java index ec0b109..e4ce4e8 100644 --- a/common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java +++ b/common/src/main/java/cn/smartjavaai/common/preprocess/BufferedImagePreprocessor.java @@ -10,7 +10,7 @@ import java.awt.image.BufferedImage; * @author dwj * @date 2025/6/27 */ -public class BufferedImagePreprocessor { +public class BufferedImagePreprocessor implements ImagePreprocessor{ private BufferedImage image; private DetectionRectangle rect; diff --git a/common/src/main/java/cn/smartjavaai/common/preprocess/DJLImagePreprocessor.java b/common/src/main/java/cn/smartjavaai/common/preprocess/DJLImagePreprocessor.java new file mode 100644 index 0000000..8cf1d30 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/preprocess/DJLImagePreprocessor.java @@ -0,0 +1,84 @@ +package cn.smartjavaai.common.preprocess; + +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.cv.SmartImageFactory; +import cn.smartjavaai.common.entity.DetectionRectangle; +import org.opencv.core.Mat; + +import java.awt.*; +import java.awt.image.BufferedImage; + +/** + * @author dwj + */ +public class DJLImagePreprocessor implements ImagePreprocessor{ + + + private final ImagePreprocessor delegate; + private Image input = null; + + public DJLImagePreprocessor(Image image, DetectionRectangle rect) { + this.input = image; + if (input.getWrappedImage() instanceof BufferedImage) { + this.delegate = new BufferedImagePreprocessor((BufferedImage) input.getWrappedImage(), rect); + } else if (input.getWrappedImage() instanceof Mat) { + this.delegate = new OpenCVPreprocessor((Mat) input.getWrappedImage(), rect); + } else { + throw new IllegalArgumentException("Unsupported input type"); + } + } + + @Override + public DJLImagePreprocessor setExtendRatio(float ratio) { + delegate.setExtendRatio(ratio); + return this; + } + + @Override + public DJLImagePreprocessor setTargetSize(int size) { + delegate.setTargetSize(size); + return this; + } + + @Override + public DJLImagePreprocessor setCenterCropSize(int size) { + delegate.setCenterCropSize(size); + return this; + } + + @Override + public DJLImagePreprocessor enableSquarePadding(boolean enable) { + delegate.enableSquarePadding(enable); + return this; + } + + @Override + public DJLImagePreprocessor enableScaling(boolean enable) { + delegate.enableScaling(enable); + return this; + } + + @Override + public DJLImagePreprocessor enableCenterCrop(boolean enable) { + delegate.enableCenterCrop(enable); + return this; + } + + @Override + public ImagePreprocessor setPaddingColor(Color color) { + return delegate.setPaddingColor(color); + } + + + @Override + public Image process() { + Object result = delegate.process(); + if (result instanceof BufferedImage) { + return SmartImageFactory.getInstance().fromBufferedImage((BufferedImage) result); + } else if (result instanceof Mat) { + return SmartImageFactory.getInstance().fromMat((Mat) result); + } + throw new IllegalStateException("Unsupported process result: " + result.getClass()); + } + +} diff --git a/common/src/main/java/cn/smartjavaai/common/preprocess/ImagePreprocessor.java b/common/src/main/java/cn/smartjavaai/common/preprocess/ImagePreprocessor.java new file mode 100644 index 0000000..396e080 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/preprocess/ImagePreprocessor.java @@ -0,0 +1,28 @@ +package cn.smartjavaai.common.preprocess; + +import java.awt.*; + +/** + * 图片预处理 + * @author dwj + */ +public interface ImagePreprocessor { + + + ImagePreprocessor setExtendRatio(float ratio); + + ImagePreprocessor setTargetSize(int size); + + ImagePreprocessor setCenterCropSize(int size); + + ImagePreprocessor enableSquarePadding(boolean enable); + + ImagePreprocessor enableScaling(boolean enable); + + ImagePreprocessor enableCenterCrop(boolean enable); + + ImagePreprocessor setPaddingColor(Color color); + + T process(); + +} diff --git a/common/src/main/java/cn/smartjavaai/common/preprocess/OpenCVPreprocessor.java b/common/src/main/java/cn/smartjavaai/common/preprocess/OpenCVPreprocessor.java new file mode 100644 index 0000000..44dddf9 --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/preprocess/OpenCVPreprocessor.java @@ -0,0 +1,154 @@ +package cn.smartjavaai.common.preprocess; + +import cn.smartjavaai.common.entity.DetectionRectangle; +import org.opencv.core.Mat; +import org.opencv.core.*; +import org.opencv.imgproc.Imgproc; + +/** + * @author dwj + */ +public class OpenCVPreprocessor implements ImagePreprocessor { + + private Mat image; + private DetectionRectangle rect; + private float extendRatio = 1; + private int targetSize = 128; + private int centerCropSize = 80; + + private Scalar paddingColor = new Scalar(127, 127, 127); // 默认灰色 + private boolean enableSquarePadding = true; + private boolean enableScaling = true; + private boolean enableCenterCrop = false; + + public OpenCVPreprocessor(Mat image, DetectionRectangle rect) { + this.image = image; + this.rect = rect; + } + + @Override + public OpenCVPreprocessor setExtendRatio(float ratio) { + this.extendRatio = ratio; + return this; + } + + @Override + public OpenCVPreprocessor setTargetSize(int size) { + this.targetSize = size; + return this; + } + + @Override + public OpenCVPreprocessor setCenterCropSize(int size) { + this.centerCropSize = size; + return this; + } + + @Override + public OpenCVPreprocessor enableSquarePadding(boolean enable) { + this.enableSquarePadding = enable; + return this; + } + + @Override + public OpenCVPreprocessor enableScaling(boolean enable) { + this.enableScaling = enable; + return this; + } + + @Override + public OpenCVPreprocessor enableCenterCrop(boolean enable) { + this.enableCenterCrop = enable; + return this; + } + + @Override + public OpenCVPreprocessor setPaddingColor(java.awt.Color color) { + this.paddingColor = new Scalar(color.getBlue(), color.getGreen(), color.getRed()); + return this; + } + + @Override + public Mat process() { + // Step 1: 裁剪 + 扩展 + Mat cropped = cropAndExtend(); + + // Step 2: 填充正方形 + Mat squared = enableSquarePadding ? squarePadding(cropped) : cropped; + + // Step 3: 缩放 + Mat scaled = enableScaling ? scaleToTarget(squared) : squared; + + // Step 4: CenterCrop + Mat finalResult = enableCenterCrop ? centerCrop(scaled) : scaled; + + return finalResult; + } + + /** + * 检测框扩展及裁剪 + */ + private Mat cropAndExtend() { + int x = rect.x; + int y = rect.y; + int width = rect.width; + int height = rect.height; + + int extendX = Math.round(width * extendRatio); + int extendY = Math.round(height * extendRatio); + + int left = Math.max(0, x - extendX); + int right = Math.min(image.width(), x + width + extendX); + int top = Math.max(0, y - extendY); + int bottom = Math.min(image.height(), y + height + extendY); + + int origRoiWidth = right - left; + int origRoiHeight = bottom - top; + int longSide = Math.max(origRoiWidth, origRoiHeight); + + // 计算可扩展空间(不超出原图边界) + int extendLeft = Math.min(left, (longSide - origRoiWidth) / 2); + int extendRight = Math.min(image.width() - right, (longSide - origRoiWidth + 1) / 2); + int extendTop = Math.min(top, (longSide - origRoiHeight) / 2); + int extendBottom = Math.min(image.height() - bottom, (longSide - origRoiHeight + 1) / 2); + + int expandedLeft = left - extendLeft; + int expandedRight = right + extendRight; + int expandedTop = top - extendTop; + int expandedBottom = bottom + extendBottom; + + Rect roi = new Rect(expandedLeft, expandedTop, expandedRight - expandedLeft, expandedBottom - expandedTop); + return new Mat(image, roi).clone(); // clone 避免与原图共享内存 + } + + /** + * 填充为正方形 + */ + private Mat squarePadding(Mat src) { + int longSide = Math.max(src.width(), src.height()); + Mat squared = new Mat(new Size(longSide, longSide), src.type(), paddingColor); + int xOffset = (longSide - src.width()) / 2; + int yOffset = (longSide - src.height()) / 2; + src.copyTo(squared.submat(yOffset, yOffset + src.height(), xOffset, xOffset + src.width())); + return squared; + } + + /** + * 缩放到目标大小 + */ + private Mat scaleToTarget(Mat src) { + Mat result = new Mat(); + Imgproc.resize(src, result, new Size(targetSize, targetSize), 0, 0, Imgproc.INTER_AREA); + return result; + } + + /** + * CenterCrop + */ + private Mat centerCrop(Mat src) { + int startX = (src.width() - centerCropSize) / 2; + int startY = (src.height() - centerCropSize) / 2; + Rect roi = new Rect(startX, startY, centerCropSize, centerCropSize); + return new Mat(src, roi).clone(); + } +} diff --git a/common/src/main/java/cn/smartjavaai/common/utils/BufferedImageUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/BufferedImageUtils.java new file mode 100644 index 0000000..fd9cb8f --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/utils/BufferedImageUtils.java @@ -0,0 +1,576 @@ +package cn.smartjavaai.common.utils; + +import ai.djl.ndarray.NDArray; +import ai.djl.util.RandomUtils; +import cn.smartjavaai.common.entity.*; +import cn.smartjavaai.common.entity.Point; +import cn.smartjavaai.common.entity.face.FaceAttribute; +import cn.smartjavaai.common.entity.face.FaceSearchResult; +import cn.smartjavaai.common.entity.face.HeadPose; +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import org.apache.commons.lang3.StringUtils; +import org.opencv.core.Mat; +import org.opencv.core.Scalar; +import org.opencv.imgproc.Imgproc; + +import javax.imageio.ImageIO; +import java.awt.*; +import java.awt.image.BufferedImage; +import java.awt.image.ComponentSampleModel; +import java.io.File; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Objects; + +/** + * @author dwj + */ +public class BufferedImageUtils { + + /** + * 拷贝图片 + * @param src + * @return + */ + public static BufferedImage copyBufferedImage(BufferedImage src) { + BufferedImage copy = new BufferedImage(src.getWidth(), src.getHeight(), src.getType()); + Graphics2D g = copy.createGraphics(); + g.drawImage(src, 0, 0, null); + g.dispose(); + return copy; + } + + /** + * 对图像解码返回BGR格式矩阵数据 + * + * @param image + * @return + */ + public static byte[] getMatrixBGR(BufferedImage image) { + byte[] matrixBGR; + if (isBGR3Byte(image)) { + matrixBGR = (byte[]) image.getData().getDataElements(0, 0, image.getWidth(), image.getHeight(), null); + } else { + // ARGB格式图像数据 + int intrgb[] = image.getRGB(0, 0, image.getWidth(), image.getHeight(), null, 0, image.getWidth()); + matrixBGR = new byte[image.getWidth() * image.getHeight() * 3]; + // ARGB转BGR格式 + for (int i = 0, j = 0; i < intrgb.length; ++i, j += 3) { + matrixBGR[j] = (byte) (intrgb[i] & 0xff); + matrixBGR[j + 1] = (byte) ((intrgb[i] >> 8) & 0xff); + matrixBGR[j + 2] = (byte) ((intrgb[i] >> 16) & 0xff); + } + } + return matrixBGR; + } + + /** + * 推断图像是否为BGR格式 + * + * @return + */ + public static boolean isBGR3Byte(BufferedImage image) { + return equalBandOffsetWith3Byte(image, new int[]{0, 1, 2}); + } + + /** + * @param image + * @param bandOffset 用于推断通道顺序 + * @return + */ + private static boolean equalBandOffsetWith3Byte(BufferedImage image, int[] bandOffset) { + if (image.getType() == BufferedImage.TYPE_3BYTE_BGR) { + if (image.getData().getSampleModel() instanceof ComponentSampleModel) { + ComponentSampleModel sampleModel = (ComponentSampleModel) image.getData().getSampleModel(); + if (Arrays.equals(sampleModel.getBandOffsets(), bandOffset)) { + return true; + } + } + } + return false; + } + + public static BufferedImage bgrToBufferedImage(byte[] data, int width, int height) { + int type = BufferedImage.TYPE_3BYTE_BGR; + // bgr to rgb + byte b; + for (int i = 0; i < data.length; i = i + 3) { + b = data[i]; + data[i] = data[i + 2]; + data[i + 2] = b; + } + BufferedImage image = new BufferedImage(width, height, type); + image.getRaster().setDataElements(0, 0, width, height, data); + return image; + } + + /** + * 检查图像是否有效 + * @param image + * @return + */ + public static boolean isImageValid(BufferedImage image) { + // 检查是否为 null 或尺寸异常(如宽高为0) + return image != null && image.getWidth() > 0 && image.getHeight() > 0; + } + + /** + * 画检测框 + * + * @param image + * @param x + * @param y + * @param width + * @param height + */ + public static void drawRect(BufferedImage image, int x, int y, int width, int height) { + // 将绘制图像转换为Graphics2D + Graphics2D g = (Graphics2D) image.getGraphics(); + try { + g.setColor(new Color(0, 255, 0)); + // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 + BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); + g.setStroke(bStroke); + g.drawRect(x, y, width, height); + } finally { + g.dispose(); + } + } + + /** + * 保存BufferedImage图片 + * @param image + * @param outputPath + * @param formatName + * @throws IOException + */ + public static void saveBufferedImage(BufferedImage image, String outputPath, String formatName) throws IOException { + if (image == null) { + throw new IllegalArgumentException("BufferedImage 不能为空"); + } + if (outputPath == null || outputPath.isEmpty()) { + throw new IllegalArgumentException("输出路径不能为空"); + } + if (formatName == null || formatName.isEmpty()) { + throw new IllegalArgumentException("格式不能为空"); + } + + Path path = Paths.get(outputPath); + Path parent = path.getParent(); + if (parent != null && !Files.exists(parent)) { + Files.createDirectories(parent); // 自动创建父目录 + } + File outFile = path.toFile(); + boolean result = ImageIO.write(image, formatName, outFile); + if (!result) { + throw new IOException("保存图片失败,不支持的格式: " + formatName); + } + } + + + /** + * 默认保存图片格式为png + * @param image + * @param outputPath + * @throws IOException + */ + public static void saveImage(BufferedImage image, String outputPath) throws IOException { + saveBufferedImage(image, outputPath, "png"); + } + + + + /** + * 画检测框(有倾斜角) + * + * @param image + * @param box + */ + public static void drawRect(BufferedImage image, NDArray box) { + float[] points = box.toFloatArray(); + int[] xPoints = new int[5]; + int[] yPoints = new int[5]; + + for (int i = 0; i < 4; i++) { + xPoints[i] = (int) points[2 * i]; + yPoints[i] = (int) points[2 * i + 1]; + } + xPoints[4] = xPoints[0]; + yPoints[4] = yPoints[0]; + + // 将绘制图像转换为Graphics2D + Graphics2D g = (Graphics2D) image.getGraphics(); + try { + g.setColor(new Color(0, 255, 0)); + // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 + BasicStroke bStroke = new BasicStroke(4, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); + g.setStroke(bStroke); + g.drawPolyline(xPoints, yPoints, 5); // xPoints, yPoints, nPoints + } finally { + g.dispose(); + } + } + + /** + * 画检测框(有倾斜角)和文本 + * + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(BufferedImage image, NDArray box, String text) { + float[] points = box.toFloatArray(); + int[] xPoints = new int[5]; + int[] yPoints = new int[5]; + + for (int i = 0; i < 4; i++) { + xPoints[i] = (int) points[2 * i]; + yPoints[i] = (int) points[2 * i + 1]; + } + xPoints[4] = xPoints[0]; + yPoints[4] = yPoints[0]; + + // 将绘制图像转换为Graphics2D + Graphics2D g = (Graphics2D) image.getGraphics(); + try { + int fontSize = 32; + Font font = new Font("楷体", Font.PLAIN, fontSize); + g.setFont(font); + g.setColor(new Color(0, 0, 255)); + // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 + BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); + g.setStroke(bStroke); + g.drawPolyline(xPoints, yPoints, 5); // xPoints, yPoints, nPoints + g.drawString(text, xPoints[0], yPoints[0]); + } finally { + g.dispose(); + } + } + + + /** + * 显示文字 + * + * @param image + * @param text + * @param x + * @param y + */ + public static void drawImageText(BufferedImage image, String text, int x, int y) { + Graphics graphics = image.getGraphics(); + int fontSize = 32; + Font font = new Font("楷体", Font.PLAIN, fontSize); + try { + graphics.setFont(font); + graphics.setColor(new Color(0, 0, 255)); + int strWidth = graphics.getFontMetrics().stringWidth(text); + graphics.drawString(text, x, y); + } finally { + graphics.dispose(); + } + } + + + /** + * 画检测框(有倾斜角)和文本 + * + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(BufferedImage image, DetectionRectangle box, String text, Color color) { + // 将绘制图像转换为Graphics2D + Graphics2D graphics = (Graphics2D) image.getGraphics(); + try { + drawRectAndText(graphics, box, text, color); + } finally { + graphics.dispose(); + } + } + + + /** + * 画检测框(有倾斜角)和文本 + * + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(BufferedImage image, DetectionRectangle box, String text, int fontSize) { + Color color = new Color(255, 0, 0); + // 将绘制图像转换为Graphics2D + Graphics2D graphics = (Graphics2D) image.getGraphics(); + graphics.setFont(new Font("楷体", Font.PLAIN, fontSize)); + try { + drawRectAndText(graphics, box, text, color); + } finally { + graphics.dispose(); + } + } + + /** + * 画检测框(有倾斜角)和文本 + * + * @param graphics + * @param box + * @param text + */ + public static void drawRectAndText(Graphics2D graphics, DetectionRectangle box, String text, Color color) { + graphics.setStroke(new BasicStroke(2)); // 线宽2像素 + graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, + RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 + int stroke = 2; + graphics.setColor(color);// 边框颜色 + graphics.drawRect(box.getX(), box.getY(), box.getWidth(), box.getHeight()); + Graphics2DUtils.drawText(graphics, text, box.getX(), box.getY(), stroke, 4); + } + + /** + * 绘制检测框 + * @param sourceImage + * @param detectionResponse + * @throws IOException + */ + public static void drawFaceSearchResult(BufferedImage sourceImage, DetectionResponse detectionResponse, String displayField) { + if(!BufferedImageUtils.isImageValid(sourceImage)){ + throw new IllegalArgumentException("图像无效"); + } + if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getDetectionInfoList()) || detectionResponse.getDetectionInfoList().isEmpty()){ + throw new IllegalArgumentException("无目标数据"); + } + Graphics2D graphics = sourceImage.createGraphics(); + graphics.setColor(Color.RED);// 边框颜色 + graphics.setStroke(new BasicStroke(2)); // 线宽2像素 + graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, + RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 + int stroke = 2; + for(DetectionInfo detectionInfo : detectionResponse.getDetectionInfoList()){ + DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); + graphics.setColor(Color.RED);// 边框颜色 + graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); + //绘制人脸关键点 + if(detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getKeyPoints() != null && + !detectionInfo.getFaceInfo().getKeyPoints().isEmpty()){ + Graphics2DUtils.drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); + //人脸查询结果 + if(detectionInfo.getFaceInfo().getFaceSearchResults() != null){ + for (FaceSearchResult faceSearchResult : detectionInfo.getFaceInfo().getFaceSearchResults()){ + if(StringUtils.isNotBlank(faceSearchResult.getMetadata())){ + JsonObject metadata = GsonUtils.parseToJsonObject(faceSearchResult.getMetadata()); + JsonElement nameElement = metadata.get("name"); + if(metadata.has("name")){ + Graphics2DUtils.drawText(graphics, nameElement.getAsString(), rectangle.getX(), rectangle.getY(), stroke, 4); + } + } + } + } + } + } + graphics.dispose(); + } + + + + + + /** + * 绘制检测框 + * @param sourceImage + * @param detectionResponse + * @throws IOException + */ + public static void drawBoundingBoxes(BufferedImage sourceImage, DetectionResponse detectionResponse) { + if(!BufferedImageUtils.isImageValid(sourceImage)){ + throw new IllegalArgumentException("图像无效"); + } + if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getDetectionInfoList()) || detectionResponse.getDetectionInfoList().isEmpty()){ + throw new IllegalArgumentException("无目标数据"); + } + Graphics2D graphics = sourceImage.createGraphics(); + graphics.setColor(Color.RED);// 边框颜色 + graphics.setStroke(new BasicStroke(2)); // 线宽2像素 + graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, + RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 + int stroke = 2; + for(DetectionInfo detectionInfo : detectionResponse.getDetectionInfoList()){ + DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); + graphics.setColor(Color.RED);// 边框颜色 + graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); + //绘制人脸关键点 + if(detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getKeyPoints() != null && + !detectionInfo.getFaceInfo().getKeyPoints().isEmpty()){ + Graphics2DUtils.drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); +// Graphics2DUtils.drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4); + //人脸属性 + if(detectionInfo.getFaceInfo().getFaceAttribute() != null){ + FaceAttribute faceAttribute = detectionInfo.getFaceInfo().getFaceAttribute(); + drawFaceAttribute(faceAttribute, rectangle, graphics); + } + } + //绘制目标检测信息 + if(detectionInfo.getObjectDetInfo() != null){ + String className = detectionInfo.getObjectDetInfo().getClassName(); + Graphics2DUtils.drawText(graphics, className, rectangle.getX(), rectangle.getY(), stroke, 4); + } + } + graphics.dispose(); + } + + /** + * 绘制矩形框和文字 + * + * @param sourceImage + * @param detectionInfo + */ + public static void drawRectAndText(BufferedImage sourceImage, DetectionInfo detectionInfo) { + if(!BufferedImageUtils.isImageValid(sourceImage)){ + throw new IllegalArgumentException("图像无效"); + } + if(Objects.isNull(detectionInfo)){ + throw new IllegalArgumentException("无目标数据"); + } + Graphics2D graphics = sourceImage.createGraphics(); + graphics.setColor(Color.RED);// 边框颜色 + graphics.setStroke(new BasicStroke(2)); // 线宽2像素 + graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, + RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 + int stroke = 2; + DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); + graphics.setColor(Color.RED);// 边框颜色 + graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); + //绘制人脸关键点 + if(detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getKeyPoints() != null && + !detectionInfo.getFaceInfo().getKeyPoints().isEmpty()){ + Graphics2DUtils.drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); +// Graphics2DUtils.drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4); + //人脸属性 + if(detectionInfo.getFaceInfo().getFaceAttribute() != null){ + FaceAttribute faceAttribute = detectionInfo.getFaceInfo().getFaceAttribute(); + drawFaceAttribute(faceAttribute, rectangle, graphics); + } + } + //绘制目标检测信息 + if(detectionInfo.getObjectDetInfo() != null){ + String className = detectionInfo.getObjectDetInfo().getClassName(); + Graphics2DUtils.drawText(graphics, className, rectangle.getX(), rectangle.getY(), stroke, 4); + } + graphics.dispose(); + } + + + /** + * 绘制人脸属性 + * @param faceAttribute + * @param rectangle + * @param graphics + */ + public static void drawFaceAttribute(FaceAttribute faceAttribute, DetectionRectangle rectangle, Graphics2D graphics){ + List lines = new ArrayList<>(); + if (faceAttribute.getGenderType() != null) { + lines.add("性别: " + faceAttribute.getGenderType().name()); + } + if (faceAttribute.getAge() != null) { + lines.add("年龄: " + faceAttribute.getAge()); + } + if (faceAttribute.getWearingMask() != null) { + lines.add("口罩: " + (faceAttribute.getWearingMask() ? "是" : "否")); + } + if (faceAttribute.getLeftEyeStatus() != null && faceAttribute.getRightEyeStatus() != null) { + lines.add("眼睛: " + faceAttribute.getLeftEyeStatus().name() + "/" + faceAttribute.getRightEyeStatus().name()); + } + if (faceAttribute.getHeadPose() != null) { + HeadPose pose = faceAttribute.getHeadPose(); + String pitch = pose.getPitch() != null ? String.valueOf(pose.getPitch().intValue()) : "-"; + String yaw = pose.getYaw() != null ? String.valueOf(pose.getYaw().intValue()) : "-"; + String roll = pose.getRoll() != null ? String.valueOf(pose.getRoll().intValue()) : "-"; + lines.add("姿态: P=" + pitch + " Y=" + yaw + " R=" + roll); + } + if (!lines.isEmpty()) { + Graphics2DUtils.drawMultilineTextWithBackground(graphics, lines, rectangle.getX(), rectangle.getY()); // 适当偏移 + } + } + + + + public static void drawPolygonWithText(BufferedImage image,List polygonLabelList, int fontSize) { + Font font = new Font("楷体", Font.PLAIN, fontSize); + Stroke stroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); + Color color = new Color(0, 0, 255); + Graphics2D g = (Graphics2D) image.getGraphics(); + for (PolygonLabel polygonLabel : polygonLabelList){ + drawPolygonWithText(g, polygonLabel.getPoints(), polygonLabel.getText(), font, color, stroke); + } + } + + /** + * 绘制多边形及文字 + * @param g Graphics2D + * @param points 多边形顶点 + * @param text 绘制的文字(可为空) + * @param font 字体 + * @param color 颜色 + * @param stroke 画笔样式 + */ + public static void drawPolygonWithText(Graphics2D g, List points, + String text, int fontSize) { + Font font = new Font("楷体", Font.PLAIN, fontSize); + Stroke stroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); + Color color = new Color(0, 0, 255); + drawPolygonWithText(g, points, text, font, color, stroke); + } + + + /** + * 绘制多边形及文字 + * @param g Graphics2D + * @param points 多边形顶点 + * @param text 绘制的文字(可为空) + * @param font 字体 + * @param color 颜色 + * @param stroke 画笔样式 + */ + public static void drawPolygonWithText(Graphics2D g, List points, + String text, Font font, + Color color, Stroke stroke) { + if (points == null || points.size() < 3) { + return; + } + + int[] xPoints = points.stream().mapToInt(p -> (int) p.getX()).toArray(); + int[] yPoints = points.stream().mapToInt(p -> (int) p.getY()).toArray(); + + g.setFont(font); + g.setColor(color); + g.setStroke(stroke); + // 绘制多边形 + g.drawPolygon(xPoints, yPoints, points.size()); + // 绘制文字(默认放在第一个点) + if (text != null && !text.isEmpty()) { + g.drawString(text, xPoints[0], yPoints[0]); + } + } + + /** + * 绘制关键点 + * @param graphics + * @param points + * @param color + */ + public static void drawKeyPoints(Graphics2D graphics, List points, Color color){ + if(points == null || points.isEmpty()){ + return; + } + for (Point point : points){ + //绘制关键点 + graphics.setColor(color); + graphics.drawRect((int)point.getX(), (int)point.getY(), 2, 2); + } + } + +} diff --git a/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java index 61caac2..98d1559 100644 --- a/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java +++ b/common/src/main/java/cn/smartjavaai/common/utils/DJLCommonUtils.java @@ -1,10 +1,17 @@ package cn.smartjavaai.common.utils; +import ai.djl.modality.cv.output.BoundingBox; +import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.output.Point; import ai.djl.ndarray.NDArray; import cn.smartjavaai.common.entity.R; +import org.opencv.core.CvType; +import org.opencv.core.Mat; import java.nio.file.Files; import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; import java.util.Objects; /** @@ -43,4 +50,113 @@ public class DJLCommonUtils { } + /** + * float NDArray To float[][] Array + * @param ndArray + * @return + */ + public static float[][] floatNDArrayToArray(NDArray ndArray) { + int rows = (int) (ndArray.getShape().get(0)); + int cols = (int) (ndArray.getShape().get(1)); + float[][] arr = new float[rows][cols]; + + float[] arrs = ndArray.toFloatArray(); + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + arr[i][j] = arrs[i * cols + j]; + } + } + return arr; + } + + + /** + * float NDArray To float[][] Array + * @param ndArray + * @param cvType + * @return + */ + public static Mat floatNDArrayToMat(NDArray ndArray, int cvType) { + int rows = (int) (ndArray.getShape().get(0)); + int cols = (int) (ndArray.getShape().get(1)); + Mat mat = new Mat(rows, cols, cvType); + + float[] arrs = ndArray.toFloatArray(); + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + mat.put(i, j, arrs[i * cols + j]); + } + } + return mat; + } + + /** + * float NDArray To Mat + * @param ndArray + * @return + */ + public static Mat floatNDArrayToMat(NDArray ndArray) { + int rows = (int) (ndArray.getShape().get(0)); + int cols = (int) (ndArray.getShape().get(1)); + Mat mat = new Mat(rows, cols, CvType.CV_32F); + + float[] arrs = ndArray.toFloatArray(); + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + mat.put(i, j, arrs[i * cols + j]); + } + } + + return mat; + + } + + /** + * uint8 NDArray To Mat + * @param ndArray + * @return + */ + public static Mat uint8NDArrayToMat(NDArray ndArray) { + int rows = (int) (ndArray.getShape().get(0)); + int cols = (int) (ndArray.getShape().get(1)); + Mat mat = new Mat(rows, cols, CvType.CV_8U); + + byte[] arrs = ndArray.toByteArray(); + + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + mat.put(i, j, arrs[i * cols + j]); + } + } + return mat; + } + + + /** + * List To Mat + * @param points + * @return + */ + public static Mat toMat(List points) { + Mat mat = new Mat(points.size(), 2, CvType.CV_32F); + for (int i = 0; i < points.size(); i++) { + ai.djl.modality.cv.output.Point point = points.get(i); + mat.put(i, 0, (float) point.getX()); + mat.put(i, 1, (float) point.getY()); + } + return mat; + } + + /** + * 构建一个空的 DetectedObjects 对象 + * @return + */ + public static DetectedObjects buildEmptyDetectedObjects(){ + List classNames = new ArrayList<>(); + List probabilities = new ArrayList<>(); + List boxes = new ArrayList<>(); + return new DetectedObjects(classNames, probabilities, boxes); + } + + } diff --git a/common/src/main/java/cn/smartjavaai/common/utils/Graphics2DUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/Graphics2DUtils.java new file mode 100644 index 0000000..4c50a1f --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/utils/Graphics2DUtils.java @@ -0,0 +1,74 @@ +package cn.smartjavaai.common.utils; + +import cn.smartjavaai.common.entity.Point; + +import java.awt.*; +import java.util.List; + +/** + * @author dwj + */ +public class Graphics2DUtils { + + + /** + * 绘制文本 + * @param g + * @param text + * @param x + * @param y + * @param stroke + * @param padding + */ + public static void drawText(Graphics2D g, String text, int x, int y, int stroke, int padding) { + FontMetrics metrics = g.getFontMetrics(); + x += stroke / 2; + y += stroke / 2; + int width = metrics.stringWidth(text) + padding * 2 - stroke / 2; + int height = metrics.getHeight() + metrics.getDescent(); + int ascent = metrics.getAscent(); + y = Math.max(0, y - height); + java.awt.Rectangle background = new java.awt.Rectangle(x, y, width, height); + g.fill(background); + g.setPaint(Color.WHITE); + g.drawString(text, x + padding, y + ascent); + } + + /** + * 绘制人脸关键点 + * @param g + * @param keyPoints + */ + public static void drawLandmarks(Graphics2D g, List keyPoints) { + g.setColor(new Color(246, 96, 0)); + BasicStroke bStroke = new BasicStroke(4.0F, 0, 0); + g.setStroke(bStroke); + for (Point point : keyPoints){ + g.drawRect((int)point.getX(), (int)point.getY(), 2, 2); + } + } + + + public static void drawMultilineTextWithBackground(Graphics2D g, List lines, int x, int y) { + Font font = new Font("SansSerif", Font.PLAIN, 14); + g.setFont(font); + FontMetrics fm = g.getFontMetrics(); + int lineHeight = fm.getHeight(); + int maxWidth = lines.stream().mapToInt(fm::stringWidth).max().orElse(0); + + int padding = 4; + int boxWidth = maxWidth + padding * 2; + int boxHeight = lineHeight * lines.size() + padding * 2; + + // 背景矩形 + g.setColor(new Color(0, 0, 0, 128)); + g.fillRoundRect(x, y, boxWidth, boxHeight, 8, 8); + + // 绘制每一行文字 + g.setColor(Color.WHITE); + for (int i = 0; i < lines.size(); i++) { + g.drawString(lines.get(i), x + padding, y + padding + (i + 1) * lineHeight - 4); + } + } + +} diff --git a/common/src/main/java/cn/smartjavaai/common/utils/GsonUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/GsonUtils.java new file mode 100644 index 0000000..3fa22ca --- /dev/null +++ b/common/src/main/java/cn/smartjavaai/common/utils/GsonUtils.java @@ -0,0 +1,69 @@ +package cn.smartjavaai.common.utils; + +import com.google.gson.Gson; +import com.google.gson.JsonObject; +import com.google.gson.JsonParser; +import com.google.gson.JsonSyntaxException; + +/** + * @author dwj + */ +public class GsonUtils { + + private static final Gson GSON = new Gson(); + + private GsonUtils() { + // 私有构造,防止实例化 + } + + /** + * 将 JSON 字符串安全转换为 JsonObject + * + * @param jsonStr JSON 字符串 + * @return JsonObject,如果解析失败返回 null + */ + public static JsonObject parseToJsonObject(String jsonStr) { + if (jsonStr == null || jsonStr.isEmpty()) { + return null; + } + try { + return JsonParser.parseString(jsonStr).getAsJsonObject(); + } catch (JsonSyntaxException | IllegalStateException e) { + // 解析失败返回 null + return null; + } + } + + /** + * 将 JSON 字符串转换为指定类型对象 + * + * @param jsonStr JSON 字符串 + * @param clazz 目标类型 + * @param 类型参数 + * @return 对象实例,如果解析失败返回 null + */ + public static T fromJson(String jsonStr, Class clazz) { + if (jsonStr == null || jsonStr.isEmpty()) { + return null; + } + try { + return GSON.fromJson(jsonStr, clazz); + } catch (JsonSyntaxException e) { + return null; + } + } + + /** + * 将对象转换为 JSON 字符串 + * + * @param obj 对象 + * @return JSON 字符串,如果对象为 null 返回 null + */ + public static String toJson(Object obj) { + if (obj == null) { + return null; + } + return GSON.toJson(obj); + } + +} diff --git a/common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java index e925453..2eb206a 100644 --- a/common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java +++ b/common/src/main/java/cn/smartjavaai/common/utils/ImageUtils.java @@ -1,209 +1,38 @@ package cn.smartjavaai.common.utils; -import ai.djl.modality.cv.BufferedImageFactory; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.ImageFactory; import ai.djl.modality.cv.output.CategoryMask; import ai.djl.modality.cv.output.DetectedObjects; +import ai.djl.modality.cv.util.NDImageUtils; import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDManager; +import ai.djl.opencv.OpenCVImageFactory; import ai.djl.util.RandomUtils; +import cn.hutool.core.codec.Base64; import cn.smartjavaai.common.cv.SmartImageFactory; +import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.PolygonLabel; import org.opencv.core.Mat; -import org.opencv.core.Scalar; -import org.opencv.core.Size; -import org.opencv.imgproc.Imgproc; -import javax.imageio.ImageIO; -import java.awt.*; import java.awt.image.BufferedImage; //import java.awt.image.ColorConvertOp; -import java.awt.image.ComponentSampleModel; -import java.awt.image.DataBufferByte; -import java.awt.image.ImageObserver; +import java.io.ByteArrayOutputStream; import java.io.File; import java.io.IOException; +import java.io.OutputStream; import java.nio.file.Files; import java.nio.file.Path; import java.nio.file.Paths; import java.util.ArrayList; -import java.util.Arrays; import java.util.List; -import java.util.Objects; /** * 图片处理工具类 */ public class ImageUtils { - /** - * @param image - * @param bandOffset 用于推断通道顺序 - * @return - */ - private static boolean equalBandOffsetWith3Byte(BufferedImage image, int[] bandOffset) { - if (image.getType() == BufferedImage.TYPE_3BYTE_BGR) { - if (image.getData().getSampleModel() instanceof ComponentSampleModel) { - ComponentSampleModel sampleModel = (ComponentSampleModel) image.getData().getSampleModel(); - if (Arrays.equals(sampleModel.getBandOffsets(), bandOffset)) { - return true; - } - } - } - return false; - } - - /** - * 推断图像是否为BGR格式 - * - * @return - */ - public static boolean isBGR3Byte(BufferedImage image) { - return equalBandOffsetWith3Byte(image, new int[]{0, 1, 2}); - } - - /** - * 对图像解码返回BGR格式矩阵数据 - * - * @param image - * @return - */ - public static byte[] getMatrixBGR(BufferedImage image) { - byte[] matrixBGR; - if (isBGR3Byte(image)) { - matrixBGR = (byte[]) image.getData().getDataElements(0, 0, image.getWidth(), image.getHeight(), null); - } else { - // ARGB格式图像数据 - int intrgb[] = image.getRGB(0, 0, image.getWidth(), image.getHeight(), null, 0, image.getWidth()); - matrixBGR = new byte[image.getWidth() * image.getHeight() * 3]; - // ARGB转BGR格式 - for (int i = 0, j = 0; i < intrgb.length; ++i, j += 3) { - matrixBGR[j] = (byte) (intrgb[i] & 0xff); - matrixBGR[j + 1] = (byte) ((intrgb[i] >> 8) & 0xff); - matrixBGR[j + 2] = (byte) ((intrgb[i] >> 16) & 0xff); - } - } - return matrixBGR; - } - - - public static BufferedImage bgrToBufferedImage(byte[] data, int width, int height) { - int type = BufferedImage.TYPE_3BYTE_BGR; - // bgr to rgb - byte b; - for (int i = 0; i < data.length; i = i + 3) { - b = data[i]; - data[i] = data[i + 2]; - data[i + 2] = b; - } - BufferedImage image = new BufferedImage(width, height, type); - image.getRaster().setDataElements(0, 0, width, height, data); - return image; - } - - /** - * 检查图像是否有效 - * @param image - * @return - */ - public static boolean isImageValid(BufferedImage image) { - // 检查是否为 null 或尺寸异常(如宽高为0) - return image != null && image.getWidth() > 0 && image.getHeight() > 0; - } - - /** - * 画检测框 - * - * @param image - * @param x - * @param y - * @param width - * @param height - */ - public static void drawImageRect(BufferedImage image, int x, int y, int width, int height) { - // 将绘制图像转换为Graphics2D - Graphics2D g = (Graphics2D) image.getGraphics(); - try { - g.setColor(new Color(0, 255, 0)); - // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 - BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); - g.setStroke(bStroke); - g.drawRect(x, y, width, height); - } finally { - g.dispose(); - } - } - - /** - * 画检测框 - * - * @param image - * @param x - * @param y - * @param width - * @param 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(); - try { - g.setColor(new Color(0, 255, 0)); - // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 - BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); - g.setStroke(bStroke); - g.drawRect(x, y, width, height); - } finally { - g.dispose(); - } - } - - - /** - * 保存BufferedImage图片 - * - * @param img - * @param name - * @param path - */ - public static void saveImage(BufferedImage img, String name, String path) { - Mat mat = OpenCVUtils.image2Mat(img); - Image djlImg = ImageFactory.getInstance().fromImage(mat); // 支持多种图片格式,自动适配 - Path outputDir = Paths.get(path); - Path imagePath = outputDir.resolve(name); - // OpenJDK 不能保存 jpg 图片的 alpha channel - try { - djlImg.save(Files.newOutputStream(imagePath), "png"); - } catch (IOException e) { - e.printStackTrace(); - } - mat.release(); - } - - - /** - * 保存BufferedImage图片 - * - * @param img - * @param path - */ - public static void saveImage(BufferedImage img, String path) { - Mat mat = OpenCVUtils.image2Mat(img); - Image djlImg = ImageFactory.getInstance().fromImage(mat); // 支持多种图片格式,自动适配 - Path outputDir = Paths.get(path); - // OpenJDK 不能保存 jpg 图片的 alpha channel - try { - djlImg.save(Files.newOutputStream(outputDir), "png"); - } catch (IOException e) { - e.printStackTrace(); - } - mat.release(); - } - - - - - /** @@ -213,7 +42,7 @@ public class ImageUtils { * @param name * @param path */ - public static void saveImage(Image img, String name, String path) { + public static void save(Image img, String name, String path) { Path outputDir = Paths.get(path); Path imagePath = outputDir.resolve(name); // OpenJDK 不能保存 jpg 图片的 alpha channel @@ -224,6 +53,23 @@ public class ImageUtils { } } + /** + * 获取图片矩阵BGR + * + * @param img + * @return + */ + public static byte[] getMatrixBGR(Image img) { + if (img.getWrappedImage() instanceof BufferedImage){ + return BufferedImageUtils.getMatrixBGR((BufferedImage)img.getWrappedImage()); + }else if (img.getWrappedImage() instanceof Mat){ + return OpenCVUtils.getMatrixBGR((Mat)img.getWrappedImage()); + }else { + throw new RuntimeException("不支持的图片类型"); + } + } + + /** * 保存图片,含检测框 * @@ -246,138 +92,6 @@ public class ImageUtils { - - /** - * 画检测框(有倾斜角) - * - * @param image - * @param box - */ - public static void drawImageRect(BufferedImage image, NDArray box) { - float[] points = box.toFloatArray(); - int[] xPoints = new int[5]; - int[] yPoints = new int[5]; - - for (int i = 0; i < 4; i++) { - xPoints[i] = (int) points[2 * i]; - yPoints[i] = (int) points[2 * i + 1]; - } - xPoints[4] = xPoints[0]; - yPoints[4] = yPoints[0]; - - // 将绘制图像转换为Graphics2D - Graphics2D g = (Graphics2D) image.getGraphics(); - try { - g.setColor(new Color(0, 255, 0)); - // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 - BasicStroke bStroke = new BasicStroke(4, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); - g.setStroke(bStroke); - g.drawPolyline(xPoints, yPoints, 5); // xPoints, yPoints, nPoints - } finally { - g.dispose(); - } - } - - /** - * 画检测框(有倾斜角)和文本 - * - * @param image - * @param box - * @param text - */ - public static void drawImageRectWithText(BufferedImage image, NDArray box, String text) { - float[] points = box.toFloatArray(); - int[] xPoints = new int[5]; - int[] yPoints = new int[5]; - - for (int i = 0; i < 4; i++) { - xPoints[i] = (int) points[2 * i]; - yPoints[i] = (int) points[2 * i + 1]; - } - xPoints[4] = xPoints[0]; - yPoints[4] = yPoints[0]; - - // 将绘制图像转换为Graphics2D - Graphics2D g = (Graphics2D) image.getGraphics(); - try { - int fontSize = 32; - Font font = new Font("楷体", Font.PLAIN, fontSize); - g.setFont(font); - g.setColor(new Color(0, 0, 255)); - // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 - BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); - g.setStroke(bStroke); - g.drawPolyline(xPoints, yPoints, 5); // xPoints, yPoints, nPoints - g.drawString(text, xPoints[0], yPoints[0]); - } finally { - g.dispose(); - } - } - - - - /** - * 显示文字 - * - * @param image - * @param text - * @param x - * @param y - */ - public static void drawImageText(BufferedImage image, String text, int x, int y) { - Graphics graphics = image.getGraphics(); - int fontSize = 32; - Font font = new Font("楷体", Font.PLAIN, fontSize); - try { - graphics.setFont(font); - graphics.setColor(new Color(0, 0, 255)); - int strWidth = graphics.getFontMetrics().stringWidth(text); - graphics.drawString(text, x, y); - } finally { - graphics.dispose(); - } - } - - - /** - * 画检测框(有倾斜角)和文本 - * - * @param image - * @param box - * @param text - */ - public static void drawImageRectWithText(BufferedImage image, DetectionRectangle box, String text, Color color) { - // 将绘制图像转换为Graphics2D - Graphics2D graphics = (Graphics2D) image.getGraphics(); - try { - graphics.setColor(Color.RED);// 边框颜色 - graphics.setStroke(new BasicStroke(2)); // 线宽2像素 - graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, - RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 - int stroke = 2; - graphics.setColor(color);// 边框颜色 - graphics.drawRect(box.getX(), box.getY(), box.getWidth(), box.getHeight()); - drawText(graphics, text, box.getX(), box.getY(), stroke, 4); - graphics.dispose(); - } finally { - graphics.dispose(); - } - } - - public static void drawText(Graphics2D g, String text, int x, int y, int stroke, int padding) { - FontMetrics metrics = g.getFontMetrics(); - x += stroke / 2; - y += stroke / 2; - int width = metrics.stringWidth(text) + padding * 2 - stroke / 2; - int height = metrics.getHeight() + metrics.getDescent(); - int ascent = metrics.getAscent(); - y = Math.max(0, y - height); - java.awt.Rectangle background = new java.awt.Rectangle(x, y, width, height); - g.fill(background); - g.setPaint(Color.WHITE); - g.drawString(text, x + padding, y + ascent); - } - /** * 计算左上角,右下角坐标 (x0,y0,x1,y1) * Get absolute coordinations @@ -448,7 +162,7 @@ public class ImageUtils { name.endsWith(".gif") || name.endsWith(".tiff") || name.endsWith(".webp")) { - Image img = ImageFactory.getInstance().fromInputStream(Files.newInputStream(file.toPath())); + Image img = SmartImageFactory.getInstance().fromInputStream(Files.newInputStream(file.toPath())); imageList.add(img); } } @@ -476,46 +190,6 @@ public class ImageUtils { return true; } - /** - * 在图像上绘制带白色背景、黑色文字的文本 - */ - public static void putTextWithBackground(Mat image, String text, org.opencv.core.Point origin, Scalar textColor, Scalar backgroundColor, int padding) { - // 默认字体 - int font = Imgproc.FONT_HERSHEY_SCRIPT_SIMPLEX; - // 默认字体缩放大小 - double fontScale = 1.0; - //线条粗细 - int thickness = 2; - //获取文字大小 - int[] baseLine = new int[1]; - Size textSize = Imgproc.getTextSize(text, font, fontScale, thickness, baseLine); - int textWidth = (int) textSize.width; - int textHeight = (int) textSize.height; - - //计算带padding的背景框 - org.opencv.core.Point bgTopLeft = new org.opencv.core.Point(origin.x - padding, origin.y - textHeight - padding); - org.opencv.core.Point bgBottomRight = new org.opencv.core.Point(origin.x + textWidth + padding, origin.y + baseLine[0] + padding); - - //绘制背景矩形 - Imgproc.rectangle(image, bgTopLeft, bgBottomRight, backgroundColor, Imgproc.FILLED); - - //绘制文字(黑色) - Imgproc.putText(image, text, origin, font, fontScale, textColor, thickness); - } - - - /** - * 拷贝图片 - * @param src - * @return - */ - public static BufferedImage copyBufferedImage(BufferedImage src) { - BufferedImage copy = new BufferedImage(src.getWidth(), src.getHeight(), src.getType()); - Graphics2D g = copy.createGraphics(); - g.drawImage(src, 0, 0, null); - g.dispose(); - return copy; - } /** @@ -525,9 +199,9 @@ public class ImageUtils { */ public static Image copy(Image src) { Object srcData = src.getWrappedImage(); - //当图片未BufferedImage,DJL的duplicate会有问题 + //当图片是BufferedImage,DJL的duplicate会有问题 if (srcData instanceof BufferedImage) { - return SmartImageFactory.getInstance().fromImage(copyBufferedImage((BufferedImage) srcData)); + return SmartImageFactory.getInstance().fromBufferedImage(BufferedImageUtils.copyBufferedImage((BufferedImage) srcData)); }else{ return src.duplicate(); } @@ -584,8 +258,241 @@ public class ImageUtils { image.drawImage(maskImage, true); } + /** + * 保存 Image 到指定路径,格式根据后缀自动推断 + */ + public static void save(Image image, Path path) throws IOException { + String fileName = path.getFileName().toString().toLowerCase(); + String format = "png"; // 默认 png + if (fileName.endsWith(".jpg") || fileName.endsWith(".jpeg")) { + format = "jpg"; + } else if (fileName.endsWith(".bmp")) { + format = "bmp"; + } else if (fileName.endsWith(".webp")) { + format = "webp"; + } + Files.createDirectories(path.getParent()); + try (OutputStream os = Files.newOutputStream(path)) { + image.save(os, format); + } + } + + /** + * 保存 Image 到指定路径,格式根据后缀自动推断 + */ + public static void save(Image image, Path path, String format) throws IOException { + Files.createDirectories(path.getParent()); + try (OutputStream os = Files.newOutputStream(path)) { + image.save(os, format); + } + } + + /** + * 保存 Image 到指定路径 + */ + public static void save(Image image, String imagePath) throws IOException { + Path path = Paths.get(imagePath); + Files.createDirectories(path.getParent()); + try (OutputStream os = Files.newOutputStream(path)) { + image.save(os, "png"); + } + } + + /** + * 转换为 BufferedImage + */ + public static BufferedImage toBufferedImage(Image image) { + Object wrapped = image.getWrappedImage(); + if (wrapped instanceof BufferedImage) { + return (BufferedImage) wrapped; + } else if (wrapped instanceof Mat) { + Mat mat = (Mat) wrapped; + return OpenCVUtils.mat2Image(mat); + } else { + throw new IllegalArgumentException("Unsupported wrapped image type: " + wrapped.getClass()); + } + } + + /** + * 转换为 Mat + */ + public static Mat toMat(Image image) { + Object wrapped = image.getWrappedImage(); + if (wrapped instanceof BufferedImage) { + return OpenCVUtils.image2Mat((BufferedImage) wrapped); + } else if (wrapped instanceof Mat) { + return (Mat) wrapped; + } else { + throw new IllegalArgumentException("Unsupported wrapped image type: " + wrapped.getClass()); + } + } + + /** + * Image 转 byte[] (默认 png 格式) + */ + public static byte[] toBytes(Image image, String format) { + try (ByteArrayOutputStream baos = new ByteArrayOutputStream()) { + image.save(baos, format); + return baos.toByteArray(); + } catch (IOException e) { + throw new RuntimeException("Failed to convert Image to byte[]", e); + } + } + + /** + * 保存 Image 到 OutputStream + * + * @param image 图像对象 + * @param os 输出流(需要调用方负责关闭) + * @param format 保存格式(png/jpg/webp) + */ + public static void toOutputStream(Image image, OutputStream os, String format) { + try { + image.save(os, format); + } catch (IOException e) { + throw new RuntimeException("Failed to write image to OutputStream", e); + } + } + + /** + * 转换 Image 为 Base64 字符串 + * + * @param image 图像对象 + * @param format 输出格式(png/jpg/webp) + * @return Base64 编码的字符串 + */ + public static String toBase64(Image image, String format) { + try (ByteArrayOutputStream baos = new ByteArrayOutputStream()) { + image.save(baos, format); + return Base64.encode(baos.toByteArray()); + } catch (IOException e) { + throw new RuntimeException("Failed to convert image to Base64", e); + } + } + + /** + * 释放 OpenCV Mat + * @param image + */ + public static void releaseOpenCVMat(Image image){ + if (image != null && image.getWrappedImage() instanceof Mat){ + ((Mat)image.getWrappedImage()).release(); + } + } + + /** + * 绘制检测结果 + * @param sourceImage + * @param detectionResponse + * @return + */ + public static Image drawBoundingBoxes(Image sourceImage, DetectionResponse detectionResponse){ + Object srcData = sourceImage.getWrappedImage(); + if (srcData instanceof BufferedImage) { + BufferedImage copyBufferedImage = BufferedImageUtils.copyBufferedImage((BufferedImage) srcData); + BufferedImageUtils.drawBoundingBoxes(copyBufferedImage, detectionResponse); + return SmartImageFactory.getInstance().fromBufferedImage(copyBufferedImage); + }else if (srcData instanceof Mat) { + Mat srcMat = ((Mat) srcData).clone(); + OpenCVUtils.drawBoundingBoxes(srcMat, detectionResponse); + return SmartImageFactory.getInstance().fromMat(srcMat); + }else { + throw new IllegalArgumentException("Unsupported wrapped image type: " + srcData.getClass()); + } + } + + /** + * 逆时针旋转图片 + * + * @param image + * @param times + * @return + */ + public static Image rotateImg(Image image, int times) { + try (NDManager manager = NDManager.newBaseManager()) { + NDArray rotated = NDImageUtils.rotate90(image.toNDArray(manager), times); + return OpenCVImageFactory.getInstance().fromNDArray(rotated); + } + } + + /** + * 图片旋转 + * + * @param manager + * @param image + * @return + */ + public static Image rotateImg(NDManager manager, Image image) { + NDArray rotated = NDImageUtils.rotate90(image.toNDArray(manager), 1); + return ImageFactory.getInstance().fromNDArray(rotated); + } + + + public static void drawPolygonWithText(Image image, List polygonLabelList, int fontSize) { + Object srcData = image.getWrappedImage(); + if (srcData instanceof BufferedImage) { + BufferedImageUtils.drawPolygonWithText((BufferedImage) srcData, polygonLabelList, fontSize); + }else if (srcData instanceof Mat) { + Mat srcMat = (Mat) srcData; + OpenCVUtils.drawPolygonWithText(srcMat, polygonLabelList, fontSize); + }else { + throw new IllegalArgumentException("Unsupported wrapped image type: " + srcData.getClass()); + } + } + /** + * 绘制矩形框 + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(Image image, DetectionRectangle box, String text) { + Object srcData = image.getWrappedImage(); + if (srcData instanceof BufferedImage) { + BufferedImageUtils.drawRectAndText((BufferedImage) srcData, box, text,12); + }else if (srcData instanceof Mat) { + Mat srcMat = (Mat) srcData; + OpenCVUtils.drawRectAndText(srcMat, box, text, 0.5); + }else { + throw new IllegalArgumentException(); + } + } + /** + * 绘制矩形框 + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(Image image, DetectionRectangle box, String text, double fontSize) { + Object srcData = image.getWrappedImage(); + if (srcData instanceof BufferedImage) { + BufferedImageUtils.drawRectAndText((BufferedImage) srcData, box, text, (int)fontSize); + }else if (srcData instanceof Mat) { + Mat srcMat = (Mat) srcData; + OpenCVUtils.drawRectAndText(srcMat, box, text, fontSize); + }else { + throw new IllegalArgumentException(); + } + } + + public static void drawRectAndText(Image image, DetectionInfo detectionInfo){ + Object srcData = image.getWrappedImage(); + if (srcData instanceof BufferedImage) { + BufferedImageUtils.drawRectAndText((BufferedImage) srcData, detectionInfo); + }else if (srcData instanceof Mat) { + Mat srcMat = (Mat) srcData; + OpenCVUtils.drawRectAndText(srcMat, detectionInfo); + }else { + throw new IllegalArgumentException(); + } + } + + public static void drawRectAndText(Image image, List detectionInfoList){ + for(DetectionInfo detectionInfo : detectionInfoList){ + drawRectAndText(image, detectionInfo); + } + } } diff --git a/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java b/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java index 909cf6a..bec5fb8 100644 --- a/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java +++ b/common/src/main/java/cn/smartjavaai/common/utils/OpenCVUtils.java @@ -4,24 +4,36 @@ 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.opencv.OpenCVImageFactory; import ai.djl.util.RandomUtils; import cn.smartjavaai.common.entity.DetectionInfo; +import cn.smartjavaai.common.entity.DetectionRectangle; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.PolygonLabel; +import cn.smartjavaai.common.entity.face.FaceAttribute; +import cn.smartjavaai.common.entity.face.HeadPose; import org.apache.commons.collections.CollectionUtils; import org.apache.commons.lang3.StringUtils; import org.opencv.core.*; import org.opencv.core.Point; +import org.opencv.imgcodecs.Imgcodecs; import org.opencv.imgproc.Imgproc; import java.awt.*; import java.awt.image.BufferedImage; import java.awt.image.DataBufferByte; +import java.util.ArrayList; +import java.util.Arrays; import java.util.List; import java.util.Objects; +import java.util.stream.Collectors; /** * OpenCV 工具类 */ public class OpenCVUtils { + + /** * canny算法,边缘检测 * @@ -153,7 +165,7 @@ public class OpenCVUtils { * @param image * @param detectionInfoList */ - public static void drawRectAndText(Image image, List detectionInfoList) { + public static void drawRectAndText(Mat image, List detectionInfoList) { if(CollectionUtils.isEmpty(detectionInfoList)) return; for(DetectionInfo detectionInfo : detectionInfoList){ @@ -162,42 +174,90 @@ public class OpenCVUtils { } + /** + * 绘制矩形框和文字 + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(Mat image, DetectionRectangle box, String text) { + if(image.empty()) + return; + int x = box.getX(); + int y = box.getY(); + int width = box.getWidth(); + int height = box.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(image, pt1, pt2, rectangleColor, 2); + Scalar textColor = new Scalar(255.0, 255.0, 255.0); + putTextWithBackground(image, text, pt1, textColor, rectangleColor, 1); + } + + /** + * 绘制矩形框和文字 + * @param image + * @param box + * @param text + */ + public static void drawRectAndText(Mat image, DetectionRectangle box, String text, double fontSize) { + if(image.empty()) + return; + int x = box.getX(); + int y = box.getY(); + int width = box.getWidth(); + int height = box.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(image, pt1, pt2, rectangleColor, 2); + Scalar textColor = new Scalar(255.0, 255.0, 255.0); + putTextWithBackground(image, text, pt1, textColor, rectangleColor, 1, fontSize); + } + + /** * 绘制矩形框和文字 * * @param image * @param detectionInfo */ - public static void drawRectAndText(Image image, DetectionInfo detectionInfo) { - - - Mat mat = (Mat)image.getWrappedImage(); + public static void drawRectAndText(Mat image, DetectionInfo detectionInfo) { 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); - + Imgproc.rectangle(image, 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); + String className = null; + Scalar textColor = new Scalar(255.0, 255.0, 255.0); + //目标检测信息 + if(detectionInfo.getObjectDetInfo() != null){ + className = detectionInfo.getObjectDetInfo().getClassName(); + putTextWithBackground(image, className, pt1, textColor, rectangleColor, 1); + } + //人脸 + if(detectionInfo.getFaceInfo() != null){ + className = "face"; + putTextWithBackground(image, className, pt1, textColor, rectangleColor, 1); + //绘制关键点 + drawLandmarks(image, detectionInfo.getFaceInfo().getKeyPoints()); + //绘制人脸属性 + if(detectionInfo.getFaceInfo().getFaceAttribute() != null){ + drawFaceAttribute(detectionInfo.getFaceInfo().getFaceAttribute(), detectionInfo.getDetectionRectangle(), image); + } } - image = ImageFactory.getInstance().fromImage(mat); } + /** * 在Mat上绘制矩形框和文字 * @@ -262,4 +322,430 @@ public class OpenCVUtils { // } // return null; // } + + + /** + * 从 OpenCV Mat 中获取 BGR 格式矩阵数据 + * + * @param mat OpenCV Mat,需为 CV_8UC3 或可转换为 BGR 格式 + * @return BGR 格式字节数组,按行连续存储 + */ + public static byte[] getMatrixBGR(Mat mat) { + if (mat == null || mat.empty()) { + throw new IllegalArgumentException("Mat 不能为空"); + } + + // 确保是三通道 BGR 格式 + Mat bgrMat = new Mat(); + if (mat.channels() == 3) { + mat.copyTo(bgrMat); + } else if (mat.channels() == 4) { + // RGBA 转 BGR + Imgproc.cvtColor(mat, bgrMat, Imgproc.COLOR_RGBA2BGR); + } else if (mat.channels() == 1) { + // 灰度转 BGR + Imgproc.cvtColor(mat, bgrMat, Imgproc.COLOR_GRAY2BGR); + } else { + throw new IllegalArgumentException("不支持的通道数: " + mat.channels()); + } + + int size = (int) (bgrMat.total() * bgrMat.channels()); + byte[] data = new byte[size]; + bgrMat.get(0, 0, data); + bgrMat.release(); // 释放临时 Mat + return data; + } + + /** + * 在图像上绘制带白色背景、黑色文字的文本 + */ + public static void putTextWithBackground(Mat image, String text, org.opencv.core.Point origin, Scalar textColor, Scalar backgroundColor, int padding) { + // 默认字体 + int font = Imgproc.FONT_HERSHEY_SIMPLEX; + // 默认字体缩放大小 + double fontScale = 0.6; + //线条粗细 + int thickness = 2; + //获取文字大小 + int[] baseLine = new int[1]; + Size textSize = Imgproc.getTextSize(text, font, fontScale, thickness, baseLine); + int textWidth = (int) textSize.width; + int textHeight = (int) textSize.height; + + //计算带padding的背景框 + org.opencv.core.Point bgTopLeft = new org.opencv.core.Point(origin.x - padding, origin.y - textHeight - padding); + org.opencv.core.Point bgBottomRight = new org.opencv.core.Point(origin.x + textWidth + padding, origin.y + baseLine[0] + padding); + + //绘制背景矩形 + Imgproc.rectangle(image, bgTopLeft, bgBottomRight, backgroundColor, Imgproc.FILLED); + + //绘制文字(黑色) + Imgproc.putText(image, text, origin, font, fontScale, textColor, thickness); + } + + /** + * 在图像上绘制带白色背景、黑色文字的文本 + */ + public static void putTextWithBackground(Mat image, String text, org.opencv.core.Point origin, Scalar textColor, Scalar backgroundColor, int padding, double fontScale) { + // 默认字体 + int font = Imgproc.FONT_HERSHEY_SIMPLEX; + //线条粗细 + int thickness = 1; + //获取文字大小 + int[] baseLine = new int[1]; + Size textSize = Imgproc.getTextSize(text, font, fontScale, thickness, baseLine); + int textWidth = (int) textSize.width; + int textHeight = (int) textSize.height; + + //计算带padding的背景框 + org.opencv.core.Point bgTopLeft = new org.opencv.core.Point(origin.x - padding, origin.y - textHeight - padding); + org.opencv.core.Point bgBottomRight = new org.opencv.core.Point(origin.x + textWidth + padding, origin.y + baseLine[0] + padding); + + //绘制背景矩形 + Imgproc.rectangle(image, bgTopLeft, bgBottomRight, backgroundColor, Imgproc.FILLED); + + //绘制文字(黑色) + Imgproc.putText(image, text, origin, font, fontScale, textColor, thickness); + } + + /** + * 绘制检测结果 + * @param image 待绘制的图片 + * @param detectionResponse 检测结果 + * @return 绘制后的图片 + */ + public static void drawBoundingBoxes(Mat image, DetectionResponse detectionResponse){ + if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getDetectionInfoList()) || detectionResponse.getDetectionInfoList().isEmpty()){ + throw new IllegalArgumentException("无目标数据"); + } + drawRectAndText(image, detectionResponse.getDetectionInfoList()); + } + + /** + * 在 Mat 上绘制关键点 + * @param mat OpenCV 图像 + * @param keyPoints 人脸关键点列表 + */ + public static void drawLandmarks(Mat mat, List keyPoints) { + // 设置颜色 (BGR 格式),这里是橙色 (0,96,246) + Scalar color = new Scalar(0, 96, 246); + // 遍历关键点,用圆来表示 (比矩形更自然) + for (cn.smartjavaai.common.entity.Point p : keyPoints) { + Point cvPoint = new Point(p.getX(), p.getY()); + Imgproc.circle(mat, cvPoint, 2, color, 2, Imgproc.LINE_AA, 0); + } + } + + /** + * 在 Mat 上绘制多行文字,并带背景 + * @param mat OpenCV Mat + * @param lines 文字行 + * @param x 起始 x + * @param y 起始 y + */ + public static void drawMultilineTextWithBackground(Mat mat, List lines, int x, int y) { + int fontFace = Imgproc.FONT_HERSHEY_SIMPLEX; + double fontScale = 0.5; // 字体大小 + int thickness = 1; + int baseline[] = {0}; + + // 逐行计算最大宽度 & 总高度 + int maxWidth = 0; + int lineHeight = 0; + for (String line : lines) { + Size textSize = Imgproc.getTextSize(line, fontFace, fontScale, thickness, baseline); + maxWidth = Math.max(maxWidth, (int) textSize.width); + lineHeight = Math.max(lineHeight, (int) (textSize.height + baseline[0])); + } + + int padding = 4; + int boxWidth = maxWidth + padding * 2; + int boxHeight = lineHeight * lines.size() + padding * 2; + + // 绘制背景矩形 (半透明黑色在 OpenCV 里不好直接实现,只能画实色或用 addWeighted 合成) + Scalar bgColor = new Scalar(0, 0, 0); // BGR = 黑色 + Point topLeft = new Point(x, y); + Point bottomRight = new Point(x + boxWidth, y + boxHeight); + Imgproc.rectangle(mat, topLeft, bottomRight, bgColor, -1); // -1 表示填充 + + // 逐行绘制文字 (白色) + Scalar textColor = new Scalar(255, 255, 255); + for (int i = 0; i < lines.size(); i++) { + int textY = y + padding + (i + 1) * lineHeight; + Imgproc.putText(mat, lines.get(i), + new Point(x + padding, textY), + fontFace, fontScale, textColor, thickness, Imgproc.LINE_AA, false); + } + } + + /** + * 绘制人脸属性 + * @param faceAttribute + * @param rectangle + * @param mat + */ + public static void drawFaceAttribute(FaceAttribute faceAttribute, DetectionRectangle rectangle, Mat mat){ + List lines = new ArrayList<>(); + if (faceAttribute.getGenderType() != null) { + lines.add("gender: " + faceAttribute.getGenderType()); + } + if (faceAttribute.getAge() != null) { + lines.add("age: " + faceAttribute.getAge()); + } + if (faceAttribute.getWearingMask() != null) { + lines.add("mask: " + (faceAttribute.getWearingMask() ? "yes" : "no")); + } + if (faceAttribute.getLeftEyeStatus() != null && faceAttribute.getRightEyeStatus() != null) { + lines.add("eyes: " + faceAttribute.getLeftEyeStatus() + "/" + faceAttribute.getRightEyeStatus()); + } + if (faceAttribute.getHeadPose() != null) { + HeadPose pose = faceAttribute.getHeadPose(); + String pitch = pose.getPitch() != null ? String.valueOf(pose.getPitch().intValue()) : "-"; + String yaw = pose.getYaw() != null ? String.valueOf(pose.getYaw().intValue()) : "-"; + String roll = pose.getRoll() != null ? String.valueOf(pose.getRoll().intValue()) : "-"; + lines.add("head pose: P=" + pitch + " Y=" + yaw + " R=" + roll); + } + if (!lines.isEmpty()) { + drawMultilineTextWithBackground(mat, lines, rectangle.getX(), rectangle.getY()); // 适当偏移 + } + } + + + /** + * 透视变换 + 裁剪 + * @param srcMat + * @param landMarks + * @return + */ + public static Image transformAndCrop(Mat srcMat, List landMarks){ + if (landMarks == null || landMarks.size() != 4) { + throw new IllegalArgumentException("必须提供4个关键点"); + } + + // 步骤 1:排序为 左上、右上、右下、左下 + List ordered = PointUtils.orderPoints(landMarks); + + ai.djl.modality.cv.output.Point lt = ordered.get(0); + ai.djl.modality.cv.output.Point rt = ordered.get(1); + ai.djl.modality.cv.output.Point rb = ordered.get(2); + ai.djl.modality.cv.output.Point lb = ordered.get(3); + + // 步骤 2:计算目标图像尺寸(宽、高) + int img_crop_width = (int) Math.max( + PointUtils.distance(lt, rt), + PointUtils.distance(rb, lb) + ); + int img_crop_height = (int) Math.max( + PointUtils.distance(lt, lb), + PointUtils.distance(rt, rb) + ); + + // 步骤 3:构造目标坐标点 + List dstPoints = Arrays.asList( + new ai.djl.modality.cv.output.Point(0, 0), + new ai.djl.modality.cv.output.Point(img_crop_width, 0), + new ai.djl.modality.cv.output.Point(img_crop_width, img_crop_height), + new ai.djl.modality.cv.output.Point(0, img_crop_height) + ); + + // 步骤 4:透视变换 + Mat srcPoint2f = DJLCommonUtils.toMat(ordered); + Mat dstPoint2f = DJLCommonUtils.toMat(dstPoints); + Mat cvMat = OpenCVUtils.perspectiveTransform(srcMat, srcPoint2f, dstPoint2f); + + // 步骤 5:转为 DJL Image + 裁剪 + Image subImg = OpenCVImageFactory.getInstance().fromImage(cvMat); + subImg = subImg.getSubImage(0, 0, img_crop_width, img_crop_height); + + // 释放资源 + cvMat.release(); + srcPoint2f.release(); + dstPoint2f.release(); + return subImg; + } + + /** + * Mat To MatOfPoint + * @param mat + * @return + */ + public static MatOfPoint matToMatOfPoint(Mat mat) { + int rows = mat.rows(); + MatOfPoint matOfPoint = new MatOfPoint(); + + List list = new ArrayList<>(); + for (int i = 0; i < rows; i++) { + Point point = new Point((float) mat.get(i, 0)[0], (float) mat.get(i, 1)[0]); + list.add(point); + } + matOfPoint.fromList(list); + + return matOfPoint; + } + + /** + * Mat To double[][] Array + * @param mat + * @return + */ + public static double[][] matToDoubleArray(Mat mat) { + int rows = mat.rows(); + int cols = mat.cols(); + double[][] doubles = new double[rows][cols]; + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + doubles[i][j] = mat.get(i, j)[0]; + } + } + return doubles; + } + + /** + * Mat To float[][] Array + * @param mat + * @return + */ + public static float[][] matToFloatArray(Mat mat) { + int rows = mat.rows(); + int cols = mat.cols(); + + float[][] floats = new float[rows][cols]; + + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + floats[i][j] = (float) mat.get(i, j)[0]; + } + } + + return floats; + } + + /** + * Mat To byte[][] Array + * @param mat + * @return + */ + public static byte[][] matToUint8Array(Mat mat) { + int rows = mat.rows(); + int cols = mat.cols(); + + byte[][] bytes = new byte[rows][cols]; + + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + bytes[i][j] = (byte) mat.get(i, j)[0]; + } + } + + return bytes; + } + + /** + * float[][] Array To Mat + * @param arr + * @return + */ + public static Mat floatArrayToMat(float[][] arr) { + int rows = arr.length; + int cols = arr[0].length; + Mat mat = new Mat(rows, cols, CvType.CV_32F); + + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + mat.put(i, j, arr[i][j]); + } + } + + return mat; + } + + /** + * byte[][] Array To Mat + * @param arr + * @return + */ + public static Mat uint8ArrayToMat(byte[][] arr) { + int rows = arr.length; + int cols = arr[0].length; + Mat mat = new Mat(rows, cols, CvType.CV_8U); + + for (int i = 0; i < rows; i++) { + for (int j = 0; j < cols; j++) { + mat.put(i, j, arr[i][j]); + } + } + + return mat; + } + + + /** + * 将自定义 Point 列表转换为 OpenCV Point 列表 + * @param pointList 自定义 Point 列表 + * @return OpenCV Point 列表 + */ + public static List toCvPointList(List pointList) { + if (pointList == null) { + return null; + } + return pointList.stream() + .map(p -> new Point(p.getX(), p.getY())) + .collect(Collectors.toList()); + } + + public static void drawPolygonWithText(Mat mat, List polygonLabelList, int fontSize) { + for (PolygonLabel polygonLabel : polygonLabelList){ + List cvPointList = toCvPointList(polygonLabel.getPoints()); + drawPolygonWithText(mat, cvPointList, polygonLabel.getText(), new Scalar(0, 255, 0), 2); + } + } + + + /** + * 在图像上绘制多边形(任意边数) + * + * @param mat 图像 + * @param points 点的列表(至少2个点) + * @param color 颜色 + * @param thickness 线宽 + */ + public static void drawPolygonWithText(Mat mat, List points, String text, Scalar color, int thickness) { + if (points == null || points.size() < 2) { + return; + } + // 连线 + for (int i = 0; i < points.size(); i++) { + Point p1 = points.get(i); + Point p2 = points.get((i + 1) % points.size()); // 最后一个点连回第一个 + Imgproc.line(mat, p1, p2, color, thickness); + } + if(StringUtils.isNotBlank(text)){ + Scalar textScalar = new Scalar(0,0,0); + // 保证文字不超出矩形 + Imgproc.putText(mat, text, points.get(0), Imgproc.FONT_HERSHEY_SIMPLEX, 1, textScalar, thickness); + } + } + + public static Mat getSubImage(Mat image, int x, int y, int w, int h) { + return image.submat(new Rect(x, y, w, h)); + } + + /** + * 从本地路径读取图片并转为 Mat + * + * @param path 图片路径 + * @return Mat 对象 + */ + public static Mat loadImage(String path) { + // 使用 imread 读取 + Mat mat = Imgcodecs.imread(path); + // 判空,避免后续处理时报错 + if (mat.empty()) { + throw new IllegalArgumentException("无法加载图片: " + path); + } + return mat; + } + + + } diff --git a/examples/face-example/pom.xml b/examples/face-example/pom.xml index 08db1dd..9191c33 100644 --- a/examples/face-example/pom.xml +++ b/examples/face-example/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.24 + 1.0.25 smartai.examples.face.facedet.FaceDetDemo @@ -220,35 +220,6 @@ - - - org.bytedeco - javacpp - ${javacv.version} - ${javacv.platform.linux-arm64} - - - - org.bytedeco - ffmpeg - 6.1.1-1.5.10 - ${javacv.platform.linux-arm64} - - - - org.bytedeco - openblas - 0.3.26-1.5.10 - ${javacv.platform.linux-arm64} - - - - org.bytedeco - opencv - 4.9.0-1.5.10 - ${javacv.platform.linux-arm64} - - @@ -278,7 +249,21 @@ + + + + + + + + + + + + + + central https://repo1.maven.org/maven2/ diff --git a/examples/face-example/src/main/java/smartai/examples/face/attribute/FaceAttributeDetDemo.java b/examples/face-example/src/main/java/smartai/examples/face/attribute/FaceAttributeDetDemo.java index d240f90..b019d93 100644 --- a/examples/face-example/src/main/java/smartai/examples/face/attribute/FaceAttributeDetDemo.java +++ b/examples/face-example/src/main/java/smartai/examples/face/attribute/FaceAttributeDetDemo.java @@ -1,11 +1,14 @@ package smartai.examples.face.attribute; +import ai.djl.modality.cv.Image; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.entity.face.FaceAttribute; import cn.smartjavaai.common.entity.face.FaceInfo; +import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.face.config.FaceAttributeConfig; import cn.smartjavaai.face.config.FaceDetConfig; import cn.smartjavaai.face.enums.FaceAttributeModelEnum; @@ -38,6 +41,8 @@ public class FaceAttributeDetDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -47,13 +52,13 @@ public class FaceAttributeDetDemo { FaceAttributeConfig config = new FaceAttributeConfig(); config.setModelEnum(FaceAttributeModelEnum.SEETA_FACE6_MODEL); //需替换为实际模型存储路径 - config.setModelPath("C:/Users/Administrator/Downloads/sf3.0_models/sf3.0_models"); + config.setModelPath("C:/Users/DengWenJie/Downloads/sf3.0_models/sf3.0_models"); return FaceAttributeModelFactory.getInstance().getModel(config); } public FaceDetModel getFaceDetModel() { //需替换为实际模型存储路径 - String modelPath = "C:/Users/Administrator/Downloads/sf3.0_models/sf3.0_models"; + String modelPath = "C:/Users/DengWenJie/Downloads/sf3.0_models/sf3.0_models"; FaceDetConfig faceDetectModelConfig = new FaceDetConfig(); faceDetectModelConfig.setModelEnum(FaceDetModelEnum.SEETA_FACE6_MODEL); faceDetectModelConfig.setModelPath(modelPath); @@ -68,10 +73,12 @@ public class FaceAttributeDetDemo { public void testFaceAttributeDetect(){ try { FaceAttributeModel faceAttributeModel = getFaceAttributeModel(); - DetectionResponse detectionResponse = faceAttributeModel.detect("src/main/resources/iu_1.jpg"); + ////创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + DetectionResponse detectionResponse = faceAttributeModel.detect(image); //绘制并导出人脸属性图片,小人脸仅有人脸框 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); - FaceUtils.drawBoxesWithFaceAttribute(image, detectionResponse,"C:/Users/Administrator/Downloads/double_person_.png"); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(image); + FaceUtils.drawBoxesWithFaceAttribute(bufferedImage, detectionResponse,"C:/Users/Administrator/Downloads/double_person_.png"); log.info("人脸属性检测结果:{}", JSONObject.toJSONString(detectionResponse)); } catch (Exception e) { e.printStackTrace(); @@ -85,7 +92,9 @@ public class FaceAttributeDetDemo { public void testFaceAttributeDetect2(){ try { FaceAttributeModel faceAttributeModel = getFaceAttributeModel(); - FaceAttribute faceAttribute = faceAttributeModel.detectTopFace("src/main/resources/iu_1.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + FaceAttribute faceAttribute = faceAttributeModel.detectTopFace(image); log.info("人脸属性检测结果:{}", JSONObject.toJSONString(faceAttribute)); } catch (Exception e) { e.printStackTrace(); @@ -101,8 +110,8 @@ public class FaceAttributeDetDemo { try { FaceDetModel faceDetModel = getFaceDetModel(); FaceAttributeModel faceAttributeModel = getFaceAttributeModel(); - //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); diff --git a/examples/face-example/src/main/java/smartai/examples/face/expression/ExpressionRecDemo.java b/examples/face-example/src/main/java/smartai/examples/face/expression/ExpressionRecDemo.java index b112dba..00435a4 100644 --- a/examples/face-example/src/main/java/smartai/examples/face/expression/ExpressionRecDemo.java +++ b/examples/face-example/src/main/java/smartai/examples/face/expression/ExpressionRecDemo.java @@ -3,6 +3,7 @@ package smartai.examples.face.expression; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.ImageFactory; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; @@ -11,6 +12,7 @@ import cn.smartjavaai.common.entity.face.ExpressionResult; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.face.FacialExpression; import cn.smartjavaai.common.enums.face.LivenessStatus; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.face.config.FaceDetConfig; @@ -60,6 +62,8 @@ public class ExpressionRecDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -75,7 +79,7 @@ public class ExpressionRecDemo { //人脸检测模型,SmartJavaAI提供了多种模型选择(更多模型,请查看文档),切换模型需要同时修改modelEnum及modelPath config.setModelEnum(FaceDetModelEnum.MTCNN); //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/face_model"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/mtcnn"); //只返回相似度大于该值的人脸,需要根据实际情况调整,分值越大越严格容易漏检,分值越小越宽松容易误识别 config.setConfidenceThreshold(0.5f); //用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 @@ -106,7 +110,9 @@ public class ExpressionRecDemo { public void testExpressionDetect() { try { ExpressionModel model = getExpressionModel(); - R result = model.detectTopFace("src/main/resources/emotion/happy.png"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/emotion/happy.png"); + R result = model.detectTopFace(image); if(result.isSuccess()){ log.info("识别结果:{}", JSONObject.toJSONString(result.getData().getExpression().getDescription())); }else{ @@ -125,7 +131,9 @@ public class ExpressionRecDemo { public void testExpressionDetect2() { try { ExpressionModel model = getExpressionModel(); - R result = model.detect("src/main/resources/emotion/happy.png"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/emotion/happy.png"); + R result = model.detect(image); if(result.isSuccess()){ //log.info("识别结果:{}", JSONObject.toJSONString(result.getData())); for (DetectionInfo detectionInfo : result.getData().getDetectionInfoList()) { @@ -149,7 +157,8 @@ public class ExpressionRecDemo { try { FaceDetModel faceDetModel = getFaceDetModel(); ExpressionModel model = getExpressionModel(); - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/emotion/happy.png").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/emotion/happy.png"); R detResult = faceDetModel.detect(image); if(detResult.isSuccess()){ R> result = model.detect(image, detResult.getData()); @@ -178,7 +187,8 @@ public class ExpressionRecDemo { try { FaceDetModel faceDetModel = getFaceDetModel(); ExpressionModel model = getExpressionModel(); - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/emotion/happy.png").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/emotion/happy.png"); R detResult = faceDetModel.detect(image); if(detResult.isSuccess()){ for (DetectionInfo detectionInfo : detResult.getData().getDetectionInfoList()) { @@ -204,15 +214,16 @@ public class ExpressionRecDemo { public void testExpressionDetectAndDraw(){ try { ExpressionModel model = getExpressionModel(); - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/emotion/surprise.png").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/emotion/surprise.png"); R result = model.detect(image); if(result.isSuccess()){ //log.info("识别结果:{}", JSONObject.toJSONString(result.getData())); for (DetectionInfo detectionInfo : result.getData().getDetectionInfoList()) { log.info("识别结果:{}", JSONObject.toJSONString(detectionInfo.getFaceInfo().getExpressionResult().getExpression().getDescription())); - ImageUtils.drawImageRectWithText(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getExpressionResult().getExpression().getDescription(), Color.red); + ImageUtils.drawRectAndText(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getExpressionResult().getExpression().getDescription()); } - ImageUtils.saveImage(image, "output/detect.jpg"); + ImageUtils.save(image, "output/detect.jpg"); }else{ log.info("识别失败:{}", result.getMessage()); } @@ -225,7 +236,7 @@ public class ExpressionRecDemo { * 摄像头表情识别 * 注意事项:如果视频比较卡,可以使用轻量的人脸检测模型 */ - @Test +// @Test public void testExpressionDetectCamera(){ try { ExpressionModel expressionModel = getExpressionModel(); @@ -264,7 +275,7 @@ public class ExpressionRecDemo { JOptionPane.showConfirmDialog(null, "Failed to capture image from WebCam."); } ViewerFrame frame = new ViewerFrame(width, height); - ImageFactory factory = ImageFactory.getInstance(); + SmartImageFactory factory = SmartImageFactory.getInstance(); Size size = new Size(width, height); while (capture.isOpened()) { @@ -273,19 +284,18 @@ public class ExpressionRecDemo { } Mat resizeImage = new Mat(); Imgproc.resize(image, resizeImage, size); - Image img = factory.fromImage(resizeImage); - BufferedImage bufferedImage = OpenCVUtils.mat2Image(resizeImage); - R detectedResult = expressionModel.detect(bufferedImage); + Image img = factory.fromMat(resizeImage); + R detectedResult = expressionModel.detect(img); if(!detectedResult.isSuccess()){ log.debug("识别失败:{}", detectedResult.getMessage()); continue; } for(DetectionInfo detectionInfo : detectedResult.getData().getDetectionInfoList()){ DetectionRectangle detectionRectangle = detectionInfo.getDetectionRectangle(); - String text = detectionInfo.getFaceInfo().getExpressionResult().getExpression().getDescription() + ":" + detectionInfo.getFaceInfo().getExpressionResult().getScore(); - ImageUtils.drawImageRectWithText(bufferedImage, detectionRectangle, text, Color.red); + String text = detectionInfo.getFaceInfo().getExpressionResult().getExpression().getLabel() + ":" + detectionInfo.getFaceInfo().getExpressionResult().getScore(); + ImageUtils.drawRectAndText(img, detectionRectangle, text); } - frame.showImage(bufferedImage); + frame.showImage(ImageUtils.toBufferedImage(img)); } capture.release(); 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 8fe5a89..2593894 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 @@ -2,7 +2,9 @@ package smartai.examples.face.facedet; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.ImageFactory; +import ai.djl.util.JsonUtils; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; @@ -17,6 +19,7 @@ import cn.smartjavaai.face.enums.FaceDetModelEnum; import cn.smartjavaai.face.factory.FaceDetModelFactory; import cn.smartjavaai.face.model.facedect.FaceDetModel; import cn.smartjavaai.face.model.liveness.LivenessDetModel; +import cn.smartjavaai.face.utils.FaceUtils; import com.alibaba.fastjson.JSONObject; import lombok.extern.slf4j.Slf4j; import nu.pattern.OpenCV; @@ -51,6 +54,8 @@ public class FaceDetDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -67,7 +72,7 @@ public class FaceDetDemo { //人脸检测模型,SmartJavaAI提供了多种模型选择(更多模型,请查看文档),切换模型需要同时修改modelEnum及modelPath config.setModelEnum(FaceDetModelEnum.MTCNN); //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/face_model"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/mtcnn"); //只返回相似度大于该值的人脸,需要根据实际情况调整,分值越大越严格容易漏检,分值越小越宽松容易误识别 config.setConfidenceThreshold(0.5f); //用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 @@ -104,7 +109,7 @@ public class FaceDetDemo { //人脸检测模型,SmartJavaAI提供了多种模型选择(更多模型,请查看文档),切换模型需要同时修改modelEnum及modelPath config.setModelEnum(FaceDetModelEnum.YOLOV5_FACE_320); //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/face_model/yolo-face/yolov5face-n-0.5-320x320.onnx"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/yolo-face/yolov5face-n-0.5-320x320.onnx"); //只返回相似度大于该值的人脸,需要根据实际情况调整,分值越大越严格容易漏检,分值越小越宽松容易误识别 config.setConfidenceThreshold(0.5f); //用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 @@ -137,9 +142,16 @@ public class FaceDetDemo { public void testFaceDetect(){ try { FaceDetModel faceModel = getFaceDetModel(); - R detectedResult = faceModel.detect(imgPath); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(imgPath); + R detectedResult = faceModel.detect(image); if(detectedResult.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectedResult.getData())); + //裁剪人脸保存 + for (DetectionInfo detectionInfo : detectedResult.getData().getDetectionInfoList()) { + Image faceImage = FaceUtils.cropFace(image, detectionInfo.getDetectionRectangle()); + ImageUtils.save(faceImage, "output/face_" + detectionInfo.getDetectionRectangle().getX() + "_" + detectionInfo.getDetectionRectangle().getY() + ".jpg"); + } }else{ log.info("人脸检测失败:{}", detectedResult.getMessage()); } @@ -156,7 +168,12 @@ public class FaceDetDemo { public void testFaceDetectAndDraw(){ try { FaceDetModel faceModel = getFaceDetModel(); - faceModel.detectAndDraw("src/main/resources/largest_selfie.jpg","output/largest_selfie_detected.png"); + R detectedResult = faceModel.detectAndDraw("src/main/resources/largest_selfie.jpg","output/largest_selfie_detected.png"); + if(detectedResult.isSuccess()){ + log.info("人脸检测成功:{}", JsonUtils.toJson(detectedResult.getData())); + }else{ + log.info("人脸检测失败:{}", detectedResult.getMessage()); + } } catch (Exception e) { throw new RuntimeException(e); } @@ -170,15 +187,14 @@ public class FaceDetDemo { public void testFaceDetectAndDraw2(){ try { FaceDetModel faceModel = getFaceDetModel(); - BufferedImage image = null; - String imagePath = "src/main/resources/largest_selfie.jpg"; - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - //可以根据后续业务场景使用detectedImage - R detectedImage = faceModel.detectAndDraw(image); - if(detectedImage.isSuccess()){ - log.info("人脸检测成功"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(imgPath); + R detectionResponseR = faceModel.detectAndDraw(image); + if(detectionResponseR.isSuccess()){ + log.info("人脸检测成功:{}", JsonUtils.toJson(detectionResponseR.getData())); + ImageUtils.save(detectionResponseR.getData().getDrawnImage(), "output/iu_1_detect.png"); }else{ - log.info("人脸检测失败:{}", detectedImage.getMessage()); + log.info("人脸检测失败:{}", detectionResponseR.getMessage()); } } catch (Exception e) { throw new RuntimeException(e); @@ -187,31 +203,6 @@ public class FaceDetDemo { } - /** - * 人脸检测(GPU模式) - */ - @Test - public void testDetectFaceGPU(){ - try { - FaceDetConfig config = new FaceDetConfig(); - //人脸检测模型,SmartJavaAI提供了多种模型选择(更多模型,请查看文档),切换模型需要同时修改modelEnum及modelPath - config.setModelEnum(FaceDetModelEnum.MTCNN); - //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/face_model"); - //只返回相似度大于该值的人脸,需要根据实际情况调整,分值越大越严格容易漏检,分值越小越宽松容易误识别 - config.setConfidenceThreshold(0.5f); - config.setDevice(DeviceEnum.GPU); - FaceDetModel faceModel = FaceDetModelFactory.getInstance().getModel(config); - R detectedResult = faceModel.detect(imgPath); - if(detectedResult.isSuccess()){ - log.info("人脸检测结果:{}", JSONObject.toJSONString(detectedResult.getData())); - }else{ - log.info("人脸检测失败:{}", detectedResult.getMessage()); - } - } catch (Exception e) { - e.printStackTrace(); - } - } /** * 人脸检测(Seetaface6) @@ -221,7 +212,9 @@ public class FaceDetDemo { public void testFaceDetectSeetaface6(){ try { FaceDetModel faceModel = getSeetaface6DetModel(); - R detectedResult = faceModel.detect(imgPath); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(imgPath); + R detectedResult = faceModel.detect(image); if(detectedResult.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectedResult.getData())); }else{ @@ -276,7 +269,7 @@ public class FaceDetDemo { JOptionPane.showConfirmDialog(null, "Failed to capture image from WebCam."); } ViewerFrame frame = new ViewerFrame(width, height); - ImageFactory factory = ImageFactory.getInstance(); + SmartImageFactory factory = SmartImageFactory.getInstance(); Size size = new Size(width, height); while (capture.isOpened()) { @@ -285,9 +278,8 @@ public class FaceDetDemo { } Mat resizeImage = new Mat(); Imgproc.resize(image, resizeImage, size); - Image img = factory.fromImage(resizeImage); - BufferedImage bufferedImage = OpenCVUtils.mat2Image(resizeImage); - R detectedResult = faceModel.detect(bufferedImage); + Image img = factory.fromMat(resizeImage); + R detectedResult = faceModel.detect(img); if(!detectedResult.isSuccess()){ log.debug("识别失败:{}", detectedResult.getMessage()); continue; @@ -298,11 +290,10 @@ public class FaceDetDemo { if(detectionInfo.getScore() > 0){ text = detectionInfo.getScore() + ""; } - ImageUtils.drawImageRectWithText(bufferedImage, detectionRectangle, text, Color.red); + ImageUtils.drawRectAndText(img, detectionRectangle, text); } - frame.showImage(bufferedImage); + frame.showImage(ImageUtils.toBufferedImage(img)); } - capture.release(); System.exit(0); } catch (Exception e) { 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 8cf5c5d..6c6ead1 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 @@ -1,10 +1,13 @@ package smartai.examples.face.facerec; -import cn.smartjavaai.common.config.Config; +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.entity.face.FaceSearchResult; import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.utils.BufferedImageUtils; +import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.face.config.FaceDetConfig; import cn.smartjavaai.face.config.FaceRecConfig; import cn.smartjavaai.face.constant.FaceDetectConstant; @@ -18,7 +21,6 @@ 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 cn.smartjavaai.face.vector.config.MilvusConfig; import cn.smartjavaai.face.vector.config.SQLiteConfig; import cn.smartjavaai.face.vector.entity.FaceVector; @@ -28,6 +30,7 @@ import lombok.extern.slf4j.Slf4j; import org.junit.BeforeClass; import org.junit.Test; +import java.awt.image.BufferedImage; import java.io.IOException; import java.util.List; @@ -45,6 +48,8 @@ public class FaceRecDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -61,7 +66,7 @@ public class FaceRecDemo { //人脸检测模型,SmartJavaAI提供了多种模型选择(更多模型,请查看文档),切换模型需要同时修改modelEnum及modelPath config.setModelEnum(FaceDetModelEnum.MTCNN); //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/face_model"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/mtcnn"); //只返回相似度大于该值的人脸,需要根据实际情况调整,分值越大越严格容易漏检,分值越小越宽松容易误识别 config.setConfidenceThreshold(0.5f); //用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 @@ -116,19 +121,19 @@ public class FaceRecDemo { * 也可以使用其他模型,具体其他模型参数可以查看文档:http://doc.smartjavaai.cn/face.html * @return */ - public FaceRecModel getHighAccuracyFaceRecModel(){ + public FaceRecModel getFaceRecModel(){ FaceRecConfig config = new FaceRecConfig(); //高精度模型,速度慢 - config.setModelEnum(FaceRecModelEnum.ELASTIC_FACE_MODEL); + config.setModelEnum(FaceRecModelEnum.INSIGHT_FACE_IRSE50_MODEL); //模型路径,请下载模型并替换为本地路径:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/model/elasticface.pt"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/recognition/InsightFace/model_ir_se50.pt"); //裁剪人脸:如果图片已经是裁剪过的,则请将此参数设置为false config.setCropFace(true); //开启人脸对齐:适用于人脸不正的场景,开启将提升人脸特征准确度,关闭可以提升性能 config.setAlign(true); config.setDevice(device); //指定人脸检测模型 - config.setDetectModel(getProFaceDetModel()); + config.setDetectModel(getFaceDetModel()); return FaceRecModelFactory.getInstance().getModel(config); } @@ -141,9 +146,9 @@ public class FaceRecDemo { public FaceRecModel getHighSpeedFaceRecModel(){ FaceRecConfig config = new FaceRecConfig(); //模型枚举 - config.setModelEnum(FaceRecModelEnum.SEETA_FACE6_LIGHT_MODEL); + config.setModelEnum(FaceRecModelEnum.INSIGHT_FACE_MOBILE_FACENET_MODEL); //模型路径,请下载模型并替换为本地路径:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/xxx/Documents/develop/model/sf3.0_models"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/recognition/InsightFace/model_mobilefacenet.pt"); //裁剪人脸:如果图片已经是裁剪过的,则请将此参数设置为false config.setCropFace(true); //开启人脸对齐:适用于人脸不正的场景,开启将提升人脸特征准确度,关闭可以提升性能 @@ -161,8 +166,9 @@ public class FaceRecDemo { public FaceRecModel getFaceRecModelWithDbConfig(){ 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.setModelEnum(FaceRecModelEnum.INSIGHT_FACE_IRSE50_MODEL); + //模型路径,请下载模型并替换为本地路径:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/recognition/InsightFace/model_ir_se50.pt"); //裁剪人脸:如果图片已经是裁剪过的,则请将此参数设置为false config.setCropFace(true); //开启人脸对齐:适用于人脸不正的场景,开启将提升人脸特征准确度,关闭可以提升性能 @@ -175,9 +181,9 @@ public class FaceRecDemo { MilvusConfig vectorDBConfig = new MilvusConfig(); vectorDBConfig.setHost("127.0.0.1"); vectorDBConfig.setPort(19530); - //vectorDBConfig.setUsername("root"); - //vectorDBConfig.setPassword("Milvus"); - //vectorDBConfig.setCollectionName("face5"); +// vectorDBConfig.setUsername("root"); +// vectorDBConfig.setPassword("Milvus"); +// vectorDBConfig.setCollectionName("face6"); //ID策略:自动生成 vectorDBConfig.setIdStrategy(IdStrategy.AUTO); //索引类型:内积 (Inner Product) 不建议修改 @@ -193,8 +199,9 @@ public class FaceRecDemo { public FaceRecModel getFaceRecModelWithSQLiteConfig(){ FaceRecConfig config = new FaceRecConfig(); //高精度模型,速度慢, 追求速度请更换高速模型,具体其他模型参数可以查看文档:http://doc.smartjavaai.cn/face.html - config.setModelEnum(FaceRecModelEnum.ELASTIC_FACE_MODEL);//人脸检测模型 - config.setModelPath("/Users/wenjie/Documents/develop/model/elasticface.pt"); + config.setModelEnum(FaceRecModelEnum.INSIGHT_FACE_IRSE50_MODEL); + //模型路径,请下载模型并替换为本地路径:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/recognition/InsightFace/model_ir_se50.pt"); //裁剪人脸:如果图片已经是裁剪过的,则请将此参数设置为false config.setCropFace(true); //开启人脸对齐:适用于人脸不正的场景,开启将提升人脸特征准确度,关闭可以提升性能 @@ -221,9 +228,11 @@ public class FaceRecDemo { public void testExtractFeatures(){ try { //高精度模型,速度慢, 追求速度请更换高速模型: getHighSpeedFaceRecModel - FaceRecModel faceRecModel = getHighAccuracyFaceRecModel(); + FaceRecModel faceRecModel = getFaceRecModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //提取图片中所有人脸特征 - R faceResult = faceRecModel.extractFeatures("src/main/resources/iu_1.jpg"); + R faceResult = faceRecModel.extractFeatures(image); if(faceResult.isSuccess()){ log.info("人脸特征提取成功:{}", JSONObject.toJSONString(faceResult.getData())); }else{ @@ -246,15 +255,59 @@ public class FaceRecDemo { public void featureComparison(){ try { //高精度模型,速度慢, 追求速度请更换高速模型: getHighSpeedFaceRecModel - FaceRecModel faceRecModel = getHighAccuracyFaceRecModel(); + FaceRecModel faceRecModel = getFaceRecModel(); //基于图像直接比对人脸特征 R similarResult = faceRecModel.featureComparison("src/main/resources/iu_1.jpg","src/main/resources/iu_2.jpg"); if(similarResult.isSuccess()){ //相似度阈值不同模型不同,具体参看文档 log.info("人脸比对相似度:{}", JSONObject.toJSONString(similarResult.getData())); + //不同模型的相似度标准不同。当前阈值仅适用于 insight_face 模型,切换模型时请相应调整阈值,详情请参考文档。 + if(similarResult.getData() >= 0.62f){ + log.info("识别为同一人"); + }else{ + log.info("识别为不同人"); + } }else{ log.info("人脸比对失败:{}", similarResult.getMessage()); } + + } + catch (Exception e){ + e.printStackTrace(); + } + } + + /** + * 人脸比对1:1(基于图像直接比对) + * 流程:从输入图像中裁剪分数最高的人脸 → 提取其人脸特征 → 比对两张图片中提取的人脸特征。(接口内自动完成) + * 注意事项: + * 1、首次调用接口,可能会较慢。只要不关闭程序,后续调用会明显加快。若每次重启程序,则每次首次调用都将重新加载,仍会较慢。 + * 2、若人脸朝向不正,可开启人脸对齐以提升特征提取准确度。(方法参考自定义配置人脸特征提取) + * @throws Exception + */ + @Test + public void featureComparison3(){ + try { + //高精度模型,速度慢, 追求速度请更换高速模型: getHighSpeedFaceRecModel + FaceRecModel faceRecModel = getFaceRecModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image1 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + Image image2 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_2.jpg"); + //基于图像直接比对人脸特征 + R similarResult = faceRecModel.featureComparison(image1, image2); + if(similarResult.isSuccess()){ + //相似度阈值不同模型不同,具体参看文档 + log.info("人脸比对相似度:{}", JSONObject.toJSONString(similarResult.getData())); + //不同模型的相似度标准不同。当前阈值仅适用于 insight_face 模型,切换模型时请相应调整阈值,详情请参考文档。 + if(similarResult.getData() >= 0.62f){ + log.info("识别为同一人"); + }else{ + log.info("识别为不同人"); + } + }else{ + log.info("人脸比对失败:{}", similarResult.getMessage()); + } + } catch (Exception e){ e.printStackTrace(); @@ -273,9 +326,11 @@ public class FaceRecDemo { public void featureComparison2(){ try { //高精度模型,速度慢, 追求速度请更换高速模型: getHighSpeedFaceRecModel - FaceRecModel faceRecModel = getHighAccuracyFaceRecModel(); + FaceRecModel faceRecModel = getFaceRecModel(); //特征提取(提取分数最高人脸特征),适用于单人脸场景 - R featureResult1 = faceRecModel.extractTopFaceFeature("src/main/resources/iu_1.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image1 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + R featureResult1 = faceRecModel.extractTopFaceFeature(image1); if(featureResult1.isSuccess()){ log.info("图片1人脸特征提取成功:{}", JSONObject.toJSONString(featureResult1.getData())); }else{ @@ -283,7 +338,8 @@ public class FaceRecDemo { return; } //特征提取(提取分数最高人脸特征),适用于单人脸场景 - R featureResult2 = faceRecModel.extractTopFaceFeature("src/main/resources/iu_2.jpg"); + Image image2 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_2.jpg"); + R featureResult2 = faceRecModel.extractTopFaceFeature(image2); if(featureResult2.isSuccess()){ log.info("图片2人脸特征提取成功:{}", JSONObject.toJSONString(featureResult2.getData())); }else{ @@ -293,6 +349,12 @@ public class FaceRecDemo { //计算相似度 float similar = faceRecModel.calculSimilar(featureResult1.getData(), featureResult2.getData()); log.info("相似度:{}", similar); + //不同模型的相似度标准不同。当前阈值仅适用于 insight_face 模型,切换模型时请相应调整阈值,详情请参考文档。 + if(similar >= 0.62f){ + log.info("识别为同一人"); + }else{ + log.info("识别为不同人"); + } } catch (Exception e){ e.printStackTrace(); @@ -318,8 +380,10 @@ public class FaceRecDemo { Thread.sleep(100); } log.info("====================人脸注册=========================="); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //特征提取(提取分数最高人脸特征),适用于单人脸场景 - R featureResult = faceRecModel.extractTopFaceFeature("src/main/resources/iu_1.jpg"); + R featureResult = faceRecModel.extractTopFaceFeature(image); if(featureResult.isSuccess()){ log.info("人脸特征提取成功:{}", JSONObject.toJSONString(featureResult.getData())); }else{ @@ -341,21 +405,23 @@ public class FaceRecDemo { }else{ log.info("注册失败:{}", registerResult.getMessage()); } - /*log.info("====================人脸更新=========================="); + log.info("====================人脸更新=========================="); //更新人脸 只支持自定义ID:vectorDBConfig.setIdStrategy(IdStrategy.CUSTOM); - FaceRegisterInfo updateInfo = new FaceRegisterInfo(); - //设置人脸注册的自定义元数据,本例中使用 JSON 格式存储用户信息 - JSONObject metadataJsonUpdate = new JSONObject(); - metadataJsonUpdate.put("name", "iu_update"); - metadataJsonUpdate.put("age", "25"); - updateInfo.setMetadata(metadataJsonUpdate.toJSONString()); - //更新必须设置ID,只有 - updateInfo.setId(registerResult.getData()); - faceRecModel.upsertFace(updateInfo, "src/main/resources/iu_2.jpg"); - log.info("更新人脸成功");*/ +// FaceRegisterInfo updateInfo = new FaceRegisterInfo(); +// //设置人脸注册的自定义元数据,本例中使用 JSON 格式存储用户信息 +// JSONObject metadataJsonUpdate = new JSONObject(); +// metadataJsonUpdate.put("name", "iu_update"); +// metadataJsonUpdate.put("age", "25"); +// updateInfo.setMetadata(metadataJsonUpdate.toJSONString()); +// //更新必须设置ID,只有 +// updateInfo.setId(registerResult.getData()); +// Image image2 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_2.jpg"); +// faceRecModel.upsertFace(updateInfo, image2); +// log.info("更新人脸成功"); log.info("====================人脸查询=========================="); //特征提取(提取分数最高人脸特征),适用于单人脸场景 - R featureResult2 = faceRecModel.extractTopFaceFeature("src/main/resources/iu_3.jpg"); + Image image3 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_3.jpg"); + R featureResult2 = faceRecModel.extractTopFaceFeature(image3); if(featureResult2.isSuccess()){ log.info("人脸特征提取成功:{}", JSONObject.toJSONString(featureResult2.getData())); }else{ @@ -364,8 +430,7 @@ public class FaceRecDemo { } FaceSearchParams faceSearchParams = new FaceSearchParams(); faceSearchParams.setTopK(1); - faceSearchParams.setThreshold(0.8f); - +// faceSearchParams.setThreshold(0.62f); List faceSearchResults = faceRecModel.search(featureResult2.getData(), faceSearchParams); // R faceSearchResults = faceModel.search("src/main/resources/face/iu_3.jpg", faceSearchParams); log.info("人脸查询结果:{}", JSONArray.toJSONString(faceSearchResults)); @@ -397,7 +462,9 @@ public class FaceRecDemo { } log.info("====================人脸注册=========================="); //特征提取(提取分数最高人脸特征),适用于单人脸场景 - R featureResult = faceRecModel.extractTopFaceFeature("src/main/resources/iu_1.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + R featureResult = faceRecModel.extractTopFaceFeature(image); if(featureResult.isSuccess()){ log.info("人脸特征提取成功:{}", JSONObject.toJSONString(featureResult.getData())); }else{ @@ -429,11 +496,13 @@ public class FaceRecDemo { updateInfo.setMetadata(metadataJsonUpdate.toJSONString()); //更新必须设置ID,只有 updateInfo.setId(registerResult.getData()); - faceRecModel.upsertFace(updateInfo, "src/main/resources/iu_2.jpg"); + Image image2 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_2.jpg"); + faceRecModel.upsertFace(updateInfo, image2); log.info("更新人脸成功"); log.info("====================人脸查询=========================="); //特征提取(提取分数最高人脸特征),适用于单人脸场景 - R featureResult2 = faceRecModel.extractTopFaceFeature("src/main/resources/iu_3.jpg"); + Image image3 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_3.jpg"); + R featureResult2 = faceRecModel.extractTopFaceFeature(image3); if(featureResult2.isSuccess()){ log.info("人脸特征提取成功:{}", JSONObject.toJSONString(featureResult2.getData())); }else{ @@ -442,7 +511,7 @@ public class FaceRecDemo { } FaceSearchParams faceSearchParams = new FaceSearchParams(); faceSearchParams.setTopK(1); - faceSearchParams.setThreshold(0.8f); + //faceSearchParams.setThreshold(0.62f); List faceSearchResults = faceRecModel.search(featureResult2.getData(), faceSearchParams); log.info("人脸查询结果:{}", JSONArray.toJSONString(faceSearchResults)); log.info("====================人脸删除=========================="); @@ -454,6 +523,57 @@ public class FaceRecDemo { } } + /** + * 人脸查询及绘制 + * + * @throws Exception + */ + @Test + public void searchFace3(){ + try { + //高精度模型,速度慢, 追求速度请更换高速模型 + FaceRecModel faceRecModel = getFaceRecModelWithSQLiteConfig(); + //等待加载人脸库结束 + while (!faceRecModel.isLoadFaceCompleted()){ + Thread.sleep(100); + } + log.info("====================人脸注册=========================="); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + //人脸注册信息 + FaceRegisterInfo faceRegisterInfo = new FaceRegisterInfo(); + //设置人脸注册的自定义元数据,本例中使用 JSON 格式存储用户信息 + JSONObject metadataJson = new JSONObject(); + metadataJson.put("name", "iu"); + metadataJson.put("age", "25"); + faceRegisterInfo.setMetadata(metadataJson.toJSONString()); + //可自定义 ID,若未设置则自动生成。 + //faceRegisterInfo.setId("00001"); + //人脸注册,返回人脸库ID + R registerResult = faceRecModel.register(faceRegisterInfo, image); + if(registerResult.isSuccess()){ + log.info("注册成功:ID-{}", registerResult.getData()); + }else{ + log.info("注册失败:{}", registerResult.getMessage()); + } + log.info("====================人脸查询=========================="); + //特征提取(提取分数最高人脸特征),适用于单人脸场景 + Image image3 = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_3.jpg"); + FaceSearchParams faceSearchParams = new FaceSearchParams(); + faceSearchParams.setTopK(1); + //faceSearchParams.setThreshold(0.62f); + //图片中只会显示Metadata信息中name的字段 + Image drawSearchResult = faceRecModel.drawSearchResult(image3, faceSearchParams, "name"); + ImageUtils.save(drawSearchResult, "output/search_result.jpg"); + log.info("====================人脸删除=========================="); + faceRecModel.removeRegister(registerResult.getData()); + log.info("人脸删除成功"); + } + catch (Exception e){ + e.printStackTrace(); + } + } + /** * 获取人脸信息 @@ -486,7 +606,7 @@ public class FaceRecDemo { public void listFaces(){ //使用ID获取人脸信息 try { - FaceRecModel faceRecModel = getFaceRecModelWithDbConfig(); + FaceRecModel faceRecModel = getFaceRecModelWithSQLiteConfig(); //等待加载人脸库结束 while (!faceRecModel.isLoadFaceCompleted()){ Thread.sleep(100); diff --git a/examples/face-example/src/main/java/smartai/examples/face/liveness/LivenessDetDemo.java b/examples/face-example/src/main/java/smartai/examples/face/liveness/LivenessDetDemo.java index 8950c9d..e6854e2 100644 --- a/examples/face-example/src/main/java/smartai/examples/face/liveness/LivenessDetDemo.java +++ b/examples/face-example/src/main/java/smartai/examples/face/liveness/LivenessDetDemo.java @@ -4,6 +4,7 @@ import ai.djl.modality.cv.Image; import ai.djl.modality.cv.ImageFactory; import cn.hutool.core.lang.UUID; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; @@ -65,6 +66,8 @@ public class LivenessDetDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -104,9 +107,9 @@ public class LivenessDetDemo { config.setModelEnum(LivenessModelEnum.MINI_VISION_MODEL); config.setDevice(device); //模型1路径:需替换为实际模型存储路径 - config.setModelPath("/Users/xxx/Documents/develop/model/live/2.7_80x80_MiniFASNetV2.onnx"); + config.setModelPath("/Users/wenjie/Documents/develop/model/live/2.7_80x80_MiniFASNetV2.onnx"); //SE模型路径:需替换为实际模型存储路径 - config.putCustomParam("seModelPath", "/Users/xxx/Documents/develop/model/live/4_0_0_80x80_MiniFASNetV1SE.onnx"); + config.putCustomParam("seModelPath", "/Users/wenjie/Documents/develop/model/live/4_0_0_80x80_MiniFASNetV1SE.onnx"); //人脸活体阈值,可选,超过阈值则认为是真人,低于阈值是非活体 config.setRealityThreshold(0.5f); /*视频检测帧数,可选,默认10,输出帧数超过这个number之后,就可以输出识别结果。 @@ -132,7 +135,7 @@ public class LivenessDetDemo { //人脸检测模型,SmartJavaAI提供了多种模型选择(更多模型,请查看文档),切换模型需要同时修改modelEnum及modelPath config.setModelEnum(FaceDetModelEnum.MTCNN); //下载模型并替换本地路径,下载地址:https://pan.baidu.com/s/10l22x5fRz_gwLr8EAHa1Jg?pwd=1234 提取码: 1234 - config.setModelPath("/Users/wenjie/Documents/develop/face_model"); + config.setModelPath("/Users/wenjie/Documents/develop/model/face_model/mtcnn"); //只返回相似度大于该值的人脸,需要根据实际情况调整,分值越大越严格容易漏检,分值越小越宽松容易误识别 config.setConfidenceThreshold(0.5f); //用于去除重复的人脸框,当两个框的重叠度超过该值时,只保留一个 @@ -149,10 +152,12 @@ public class LivenessDetDemo { public void testLivenessDetect(){ try { LivenessDetModel livenessDetModel = getLivenessDetModel(); - R response = livenessDetModel.detect("src/main/resources/liveness/1.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/liveness/1.jpg"); + R response = livenessDetModel.detect(image); if(response.isSuccess()){ for (DetectionInfo detectionInfo : response.getData().getDetectionInfoList()){ - log.info("活体检测结果:{}", JSONObject.toJSONString(detectionInfo.getFaceInfo().getLivenessStatus().getStatus().getDescription())); + log.info("活体检测结果:{}", JSONObject.toJSONString(detectionInfo)); } }else{ log.info("活体检测失败:{}", response.getMessage()); @@ -169,18 +174,18 @@ public class LivenessDetDemo { public void testLivenessDetectAndDraw(){ try { LivenessDetModel livenessDetModel = getLivenessDetModel(); - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/liveness/1.jpg").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/liveness/1.jpg"); R response = livenessDetModel.detect(image); if(response.isSuccess()){ for (DetectionInfo detectionInfo : response.getData().getDetectionInfoList()){ log.info("活体检测结果:{}", JSONObject.toJSONString(detectionInfo.getFaceInfo().getLivenessStatus().getStatus().getDescription())); - Color color = detectionInfo.getFaceInfo().getLivenessStatus().getStatus() == LivenessStatus.LIVE ? Color.GREEN : Color.RED; - ImageUtils.drawImageRectWithText(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getLivenessStatus().getStatus().getDescription(), color); + ImageUtils.drawRectAndText(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getLivenessStatus().getStatus().toString()); + ImageUtils.save(image, "output/detect.jpg"); } }else{ log.info("活体检测失败:{}", response.getMessage()); } - ImageUtils.saveImage(image, "output/detect.jpg"); } catch (Exception e) { throw new RuntimeException(e); } @@ -194,10 +199,12 @@ public class LivenessDetDemo { try { LivenessDetModel livenessDetModel = getLivenessDetModel(); //指定文件夹路径 - File dir = new File("face-example/src/main/resources/liveness"); + File dir = new File("src/main/resources/liveness"); File[] files = dir.listFiles(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + SmartImageFactory imageFactory = SmartImageFactory.getInstance(); for (File file : files) { - R response = livenessDetModel.detectTopFace(ImageIO.read(file)); + R response = livenessDetModel.detectTopFace(imageFactory.fromFile(file)); if(response.isSuccess()){ log.info("{}活体检测结果:{},分数:{}", file.getName(), response.getData().getStatus().getDescription(), response.getData().getScore()); }else{ @@ -218,8 +225,8 @@ public class LivenessDetDemo { try { FaceDetModel faceDetectModel = getFaceDetModel(); LivenessDetModel livenessDetModel = getLivenessDetModel(); - // 将图片路径转换为 BufferedImage - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/liveness/1.jpg").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/liveness/1.jpg"); //人脸检测 R detectionResponse = faceDetectModel.detect(image); if(detectionResponse.isSuccess()){ @@ -251,8 +258,8 @@ public class LivenessDetDemo { try { FaceDetModel faceDetModel = getFaceDetModel(); LivenessDetModel livenessDetModel = getMiniVisionLivenessDetModel(); - // 将图片路径转换为 BufferedImage - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/liveness/1.jpg").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/liveness/1.jpg"); R detResult = faceDetModel.detect(image); if(detResult.isSuccess()){ for (DetectionInfo detectionInfo : detResult.getData().getDetectionInfoList()) { @@ -281,7 +288,7 @@ public class LivenessDetDemo { try { LivenessDetModel livenessDetModel = getLivenessDetModel(); //视频路径 - R livenessStatus = livenessDetModel.detectVideo("video.mp4"); + R livenessStatus = livenessDetModel.detectVideo("/Users/wenjie/Documents/idea_workplace/SmartJavaAI-Demo/src/main/resources/girl.mp4"); if (livenessStatus.isSuccess()){ log.info("识别结果:{}", JSONObject.toJSONString(livenessStatus.getData())); }else{ @@ -296,7 +303,7 @@ public class LivenessDetDemo { * 摄像头活体检测 * 注意事项:如果视频比较卡,可以使用轻量的人脸检测模型 */ - @Test +// @Test public void testLivenessDetectCamera(){ try { LivenessDetModel livenessDetModel = getLivenessDetModel(); @@ -335,7 +342,7 @@ public class LivenessDetDemo { JOptionPane.showConfirmDialog(null, "Failed to capture image from WebCam."); } ViewerFrame frame = new ViewerFrame(width, height); - ImageFactory factory = ImageFactory.getInstance(); + SmartImageFactory factory = SmartImageFactory.getInstance(); Size size = new Size(width, height); while (capture.isOpened()) { @@ -344,9 +351,8 @@ public class LivenessDetDemo { } Mat resizeImage = new Mat(); Imgproc.resize(image, resizeImage, size); - Image img = factory.fromImage(resizeImage); - BufferedImage bufferedImage = OpenCVUtils.mat2Image(resizeImage); - R detectedResult = livenessDetModel.detect(bufferedImage); + Image img = factory.fromMat(resizeImage); + R detectedResult = livenessDetModel.detect(img); if(!detectedResult.isSuccess()){ log.debug("识别失败:{}", detectedResult.getMessage()); continue; @@ -355,11 +361,10 @@ public class LivenessDetDemo { DetectionRectangle detectionRectangle = detectionInfo.getDetectionRectangle(); Color color = detectionInfo.getFaceInfo().getLivenessStatus().getStatus() == LivenessStatus.LIVE ? Color.GREEN : Color.RED; String text = detectionInfo.getFaceInfo().getLivenessStatus().getStatus().getDescription() + ":" + detectionInfo.getFaceInfo().getLivenessStatus().getScore(); - ImageUtils.drawImageRectWithText(bufferedImage, detectionRectangle, text, color); + ImageUtils.drawRectAndText(img, detectionRectangle, text); } - frame.showImage(bufferedImage); + frame.showImage(ImageUtils.toBufferedImage(img)); } - capture.release(); System.exit(0); } catch (Exception e) { diff --git a/examples/face-example/src/main/java/smartai/examples/face/quality/FaceQualityDetDemo.java b/examples/face-example/src/main/java/smartai/examples/face/quality/FaceQualityDetDemo.java index 10b6ed0..a297908 100644 --- a/examples/face-example/src/main/java/smartai/examples/face/quality/FaceQualityDetDemo.java +++ b/examples/face-example/src/main/java/smartai/examples/face/quality/FaceQualityDetDemo.java @@ -1,6 +1,8 @@ package smartai.examples.face.quality; +import ai.djl.modality.cv.Image; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.R; @@ -47,6 +49,8 @@ public class FaceQualityDetDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -62,7 +66,7 @@ public class FaceQualityDetDemo { QualityConfig config = new QualityConfig(); config.setModelEnum(QualityModelEnum.SEETA_FACE6_MODEL); //需替换为实际模型存储路径 - config.setModelPath("C:/Users/Administrator/Downloads/sf3.0_models/sf3.0_models"); + config.setModelPath("C:/Users/DengWenJie/Downloads/sf3.0_models/sf3.0_models"); config.setDevice(device); return FaceQualityModelFactory.getInstance().getModel(config); } @@ -74,7 +78,7 @@ public class FaceQualityDetDemo { */ public FaceDetModel getFaceDetModel() { //需替换为实际模型存储路径 - String modelPath = "C:/Users/Administrator/Downloads/sf3.0_models/sf3.0_models"; + String modelPath = "C:/Users/DengWenJie/Downloads/sf3.0_models/sf3.0_models"; FaceDetConfig faceDetectModelConfig = new FaceDetConfig(); faceDetectModelConfig.setModelEnum(FaceDetModelEnum.SEETA_FACE6_MODEL); faceDetectModelConfig.setModelPath(modelPath); @@ -91,8 +95,9 @@ public class FaceQualityDetDemo { try { FaceQualityModel faceQualityModel = getFaceQualityModel(); FaceDetModel faceDetModel = getFaceDetModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); @@ -124,8 +129,9 @@ public class FaceQualityDetDemo { try { FaceQualityModel faceQualityModel = getFaceQualityModel(); FaceDetModel faceDetModel = getFaceDetModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); @@ -157,8 +163,8 @@ public class FaceQualityDetDemo { try { FaceQualityModel faceQualityModel = getFaceQualityModel(); FaceDetModel faceDetModel = getFaceDetModel(); - //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); @@ -190,8 +196,9 @@ public class FaceQualityDetDemo { try { FaceQualityModel faceQualityModel = getFaceQualityModel(); FaceDetModel faceDetModel = getFaceDetModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); @@ -224,8 +231,9 @@ public class FaceQualityDetDemo { try { FaceQualityModel faceQualityModel = getFaceQualityModel(); FaceDetModel faceDetModel = getFaceDetModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); @@ -258,8 +266,9 @@ public class FaceQualityDetDemo { try { FaceQualityModel faceQualityModel = getFaceQualityModel(); FaceDetModel faceDetModel = getFaceDetModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); //人脸检测 - BufferedImage image = ImageIO.read(new File(Paths.get("src/main/resources/iu_1.jpg").toAbsolutePath().toString())); R detectionResponse = faceDetModel.detect(image); if(detectionResponse.isSuccess()){ log.info("人脸检测结果:{}", JSONObject.toJSONString(detectionResponse.getData())); diff --git a/examples/ocr-examples/output/ocr_4_recognized.jpg b/examples/ocr-examples/output/ocr_4_recognized.jpg index 9d5dfa1..f3de69f 100644 Binary files a/examples/ocr-examples/output/ocr_4_recognized.jpg and b/examples/ocr-examples/output/ocr_4_recognized.jpg differ diff --git a/examples/ocr-examples/output/plate_recognized2.jpg b/examples/ocr-examples/output/plate_recognized2.jpg index 468a416..d0ef832 100644 Binary files a/examples/ocr-examples/output/plate_recognized2.jpg and b/examples/ocr-examples/output/plate_recognized2.jpg differ diff --git a/examples/ocr-examples/output/table_ch2_result.jpg b/examples/ocr-examples/output/table_ch2_result.jpg index fe84cb2..999445d 100644 Binary files a/examples/ocr-examples/output/table_ch2_result.jpg and b/examples/ocr-examples/output/table_ch2_result.jpg differ diff --git a/examples/ocr-examples/pom.xml b/examples/ocr-examples/pom.xml index 646a98f..3266564 100644 --- a/examples/ocr-examples/pom.xml +++ b/examples/ocr-examples/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.24 + 1.0.25 smartai.examples.ocr.common.OcrRecognizeDemo @@ -219,39 +219,6 @@ - - - org.bytedeco - javacpp - ${javacv.version} - ${javacv.platform.linux-arm64} - - - - org.bytedeco - ffmpeg - 6.1.1-1.5.10 - ${javacv.platform.linux-arm64} - - - - org.bytedeco - openblas - 0.3.26-1.5.10 - ${javacv.platform.linux-arm64} - - - - org.bytedeco - opencv - 4.9.0-1.5.10 - ${javacv.platform.linux-arm64} - - - - - - diff --git a/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDetectionDemo.java b/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDetectionDemo.java index 0d05cb2..1b70918 100644 --- a/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDetectionDemo.java +++ b/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDetectionDemo.java @@ -2,6 +2,7 @@ package smartai.examples.ocr.common; import ai.djl.modality.cv.Image; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.utils.ImageUtils; @@ -43,6 +44,7 @@ public class OcrDetectionDemo { @BeforeClass public static void beforeAll() throws IOException { + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -56,7 +58,7 @@ public class OcrDetectionDemo { //指定检测模型,切换模型需要同时修改modelEnum及modelPath config.setModelEnum(CommonDetModelEnum.PP_OCR_V5_MOBILE_DET_MODEL); //指定模型位置,需要更改为自己的模型路径(下载地址请查看文档) - config.setDetModelPath("/Users/xxx/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_mobile_det_infer/PP-OCRv5_mobile_det_infer.onnx"); config.setDevice(device); return OcrModelFactory.getInstance().getDetModel(config); } @@ -73,7 +75,9 @@ public class OcrDetectionDemo { public void detect(){ try { OcrCommonDetModel model = getDetectionModel(); - List boxes = model.detect("src/main/resources/ocr_1.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/ocr_1.jpg"); + List boxes = model.detect(image); log.info("OCR检测结果:{}", JSONObject.toJSONString(boxes)); } catch (Exception e) { e.printStackTrace(); @@ -97,6 +101,26 @@ public class OcrDetectionDemo { } } + /** + * 文本检测并绘制结果 + * 检测图像中的文本区域,仅检测文本框位置,不识别文字内容 + * 注意事项: + * 1、批量检测时,模型应统一放在外层 try 中使用,避免重复加载,自动释放资源更安全。 + * 2、模型文件需要放在单独文件夹 + */ + @Test + public void detectAndDraw2(){ + try { + OcrCommonDetModel model = getDetectionModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/ocr_1.jpg"); + Image resultImage = model.detectAndDraw(image); + ImageUtils.save(resultImage, "output/ocr_1_detected2.jpg"); + } catch (Exception e) { + e.printStackTrace(); + } + } + /** * 批量文本检测:批量检测要求图片宽高一致 @@ -110,7 +134,7 @@ public class OcrDetectionDemo { try { OcrCommonDetModel model = getDetectionModel(); //批量检测要求图片宽高一致 - String folderPath = "/Users/xxx/Downloads/testing33"; + String folderPath = "/Users/wenjie/Downloads/testing33"; //读取文件夹中所有图片 List images = ImageUtils.readImagesFromFolder(folderPath); List> ocrResult = model.batchDetectDJLImage(images); diff --git a/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDirectionDetDemo.java b/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDirectionDetDemo.java index 6175653..e665fae 100644 --- a/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDirectionDetDemo.java +++ b/examples/ocr-examples/src/main/java/smartai/examples/ocr/common/OcrDirectionDetDemo.java @@ -1,7 +1,10 @@ package smartai.examples.ocr.common; +import ai.djl.modality.cv.Image; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.ocr.config.DirectionModelConfig; import cn.smartjavaai.ocr.config.OcrDetModelConfig; import cn.smartjavaai.ocr.entity.OcrBox; @@ -46,7 +49,7 @@ public class OcrDirectionDetDemo { //指定行文本方向检测模型,切换模型需要同时修改modelEnum及modelPath directionModelConfig.setModelEnum(DirectionModelEnum.PP_LCNET_X0_25); //指定行文本方向检测模型路径,需要更改为自己的模型路径(下载地址请查看文档) - directionModelConfig.setModelPath("/Users/xxx/Documents/develop/model/ocr/PP-LCNet_x0_25_textline_ori_infer/PP-LCNet_x0_25_textline_ori_infer.onnx"); + directionModelConfig.setModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-LCNet_x0_25_textline_ori_infer/PP-LCNet_x0_25_textline_ori_infer.onnx"); directionModelConfig.setDevice(device); directionModelConfig.setTextDetModel(getDetectionModel()); return OcrModelFactory.getInstance().getDirectionModel(directionModelConfig); @@ -61,7 +64,7 @@ public class OcrDirectionDetDemo { //指定检测模型,切换模型需要同时修改modelEnum及modelPath config.setModelEnum(CommonDetModelEnum.PP_OCR_V5_MOBILE_DET_MODEL); //指定模型位置,需要更改为自己的模型路径(下载地址请查看文档) - config.setDetModelPath("/Users/xxx/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_mobile_det_infer/PP-OCRv5_mobile_det_infer.onnx"); config.setDevice(device); return OcrModelFactory.getInstance().getDetModel(config); } @@ -78,7 +81,9 @@ public class OcrDirectionDetDemo { public void detect(){ try { OcrDirectionModel directionModel = getDirectionModel(); - List itemList = directionModel.detect("src/main/resources/ocr_1.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/ocr_1.jpg"); + List itemList = directionModel.detect(image); log.info("OCR方向检测结果1:{}", JSONObject.toJSONString(itemList)); } catch (Exception e) { e.printStackTrace(); @@ -102,6 +107,25 @@ public class OcrDirectionDetDemo { } } + /** + * 文本检测并绘制结果 + * 流程:文本检测 -> 方向分类 + * 检测图像中的文本区域,仅检测文本框位置,不识别文字内容 + * 模型需要放在单独文件夹 + */ + @Test + public void detectAndDraw2(){ + try { + OcrDirectionModel directionModel = getDirectionModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/ocr_1.jpg"); + Image resultImage = directionModel.detectAndDraw(image); + ImageUtils.save(resultImage, "output/ocr_1_detected4.jpg"); + } catch (Exception e) { + e.printStackTrace(); + } + } + } 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 e938ac4..45d5b47 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 @@ -5,7 +5,9 @@ import ai.djl.util.JsonUtils; import cn.hutool.core.img.ImgUtil; import cn.hutool.core.io.FileUtil; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.ocr.config.DirectionModelConfig; import cn.smartjavaai.ocr.config.OcrDetModelConfig; @@ -48,34 +50,70 @@ public class OcrRecognizeDemo { @BeforeClass public static void beforeAll() throws IOException { + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 //Config.setCachePath("/Users/xxx/smartjavaai_cache"); } /** - * 获取通用识别模型(不带方向矫正) + * 获取通用识别模型(高精确度模型) + * 注意事项:高精度模型,识别准确度高,速度慢 * @return */ - public OcrCommonRecModel getRecModel(){ + public OcrCommonRecModel getProRecModel(){ OcrRecModelConfig recModelConfig = new OcrRecModelConfig(); //指定文本识别模型,切换模型需要同时修改modelEnum及modelPath - recModelConfig.setRecModelEnum(CommonRecModelEnum.PP_OCR_V5_MOBILE_REC_MODEL); + recModelConfig.setRecModelEnum(CommonRecModelEnum.PP_OCR_V5_SERVER_REC_MODEL); //指定识别模型位置,需要更改为自己的模型路径(下载地址请查看文档) recModelConfig.setRecModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-OCRv5_server_rec_infer/PP-OCRv5_server_rec.onnx"); recModelConfig.setDevice(device); - recModelConfig.setTextDetModel(getDetectionModel()); + recModelConfig.setTextDetModel(getProDetectionModel()); + recModelConfig.setDirectionModel(getDirectionModel()); return OcrModelFactory.getInstance().getRecModel(recModelConfig); } /** - * 获取文本检测模型 + * 获取通用识别模型(极速模型) + * 注意事项:极速模型,识别准确度低,速度快 * @return */ - public OcrCommonDetModel getDetectionModel() { + public OcrCommonRecModel getFastRecModel(){ + OcrRecModelConfig recModelConfig = new OcrRecModelConfig(); + //指定文本识别模型,切换模型需要同时修改modelEnum及modelPath + 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.setDevice(device); + recModelConfig.setTextDetModel(getFastDetectionModel()); + return OcrModelFactory.getInstance().getRecModel(recModelConfig); + } + + + /** + * 获取文本检测模型(极速模型) + * 注意事项:极速模型,识别准确度低,速度快 + * @return + */ + public OcrCommonDetModel getFastDetectionModel() { OcrDetModelConfig config = new OcrDetModelConfig(); //指定检测模型,切换模型需要同时修改modelEnum及modelPath 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.setDevice(device); + return OcrModelFactory.getInstance().getDetModel(config); + } + + /** + * 获取文本检测模型(高精确度模型) + * 注意事项:高精度模型,识别准确度高,速度慢 + * @return + */ + public OcrCommonDetModel getProDetectionModel() { + OcrDetModelConfig config = new OcrDetModelConfig(); + //指定检测模型,切换模型需要同时修改modelEnum及modelPath + config.setModelEnum(CommonDetModelEnum.PP_OCR_V5_SERVER_DET_MODEL); + //指定模型位置,需要更改为自己的模型路径(下载地址请查看文档) 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); @@ -90,28 +128,12 @@ public class OcrRecognizeDemo { //指定行文本方向检测模型,切换模型需要同时修改modelEnum及modelPath directionModelConfig.setModelEnum(DirectionModelEnum.PP_LCNET_X0_25); //指定行文本方向检测模型路径,需要更改为自己的模型路径(下载地址请查看文档) - directionModelConfig.setModelPath("/Users/xxx/Documents/develop/model/ocr/PP-LCNet_x0_25_textline_ori_infer/PP-LCNet_x0_25_textline_ori_infer.onnx"); + directionModelConfig.setModelPath("/Users/wenjie/Documents/develop/model/ocr/PP-LCNet_x0_25_textline_ori_infer/PP-LCNet_x0_25_textline_ori_infer.onnx"); directionModelConfig.setDevice(device); return OcrModelFactory.getInstance().getDirectionModel(directionModelConfig); } - /** - * 获取通用识别模型(带方向矫正) - * @return - */ - public OcrCommonRecModel getRecModelWithDirection() { - OcrRecModelConfig recModelConfig = new OcrRecModelConfig(); - //指定文本识别模型,切换模型需要同时修改modelEnum及modelPath - recModelConfig.setRecModelEnum(CommonRecModelEnum.PP_OCR_V5_MOBILE_REC_MODEL); - //指定识别模型位置,需要更改为自己的模型路径(下载地址请查看文档) - recModelConfig.setRecModelPath("/Users/xxx/Documents/develop/model/ocr/PP-OCRv5_mobile_rec_infer/PP-OCRv5_mobile_rec_infer.onnx"); - recModelConfig.setDevice(device); - recModelConfig.setTextDetModel(getDetectionModel()); - recModelConfig.setDirectionModel(getDirectionModel()); - return OcrModelFactory.getInstance().getRecModel(recModelConfig); - } - /** * 文本识别 @@ -124,10 +146,12 @@ public class OcrRecognizeDemo { @Test public void recognize(){ try { - OcrCommonRecModel recModel = getRecModel(); + OcrCommonRecModel recModel = getFastRecModel(); //不带方向矫正,分行返回文本 OcrRecOptions options = new OcrRecOptions(false, true); - OcrInfo ocrInfo = recModel.recognize("/Users/wenjie/Downloads/49421755855753_.pic_hd.jpg",options); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/ocr_1.jpg"); + OcrInfo ocrInfo = recModel.recognize(image, options); log.info("OCR识别结果:{}", JSONObject.toJSONString(ocrInfo)); } catch (Exception e) { e.printStackTrace(); @@ -146,8 +170,10 @@ public class OcrRecognizeDemo { @Test public void recognizeHandWriting(){ try { - OcrCommonRecModel recModel = getRecModel(); - OcrInfo ocrInfo = recModel.recognize("src/main/resources/handwriting_1.jpg",new OcrRecOptions()); + OcrCommonRecModel recModel = getFastRecModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/handwriting_1.jpg"); + OcrInfo ocrInfo = recModel.recognize(image, new OcrRecOptions()); log.info("OCR识别结果:{}", JSONObject.toJSONString(ocrInfo)); } catch (Exception e) { e.printStackTrace(); @@ -166,10 +192,12 @@ public class OcrRecognizeDemo { @Test public void recognize2(){ try { - OcrCommonRecModel recModel = getRecModelWithDirection(); + OcrCommonRecModel recModel = getFastRecModel(); //带方向矫正,分行返回文本 OcrRecOptions options = new OcrRecOptions(true, true); - OcrInfo ocrInfo = recModel.recognize("src/main/resources/ocr_3.jpg",options); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/ocr_3.jpg"); + OcrInfo ocrInfo = recModel.recognize(image, options); log.info("OCR识别结果:{}", JSONObject.toJSONString(ocrInfo)); } catch (Exception e) { e.printStackTrace(); @@ -189,7 +217,7 @@ public class OcrRecognizeDemo { @Test public void recognizeAndDraw(){ try { - OcrCommonRecModel recModel = getRecModelWithDirection(); + OcrCommonRecModel recModel = getFastRecModel(); int fontSize = 18; recModel.recognizeAndDraw("src/main/resources/general_ocr_002.png", "output/ocr_4_recognized.jpg", fontSize, new OcrRecOptions()); } catch (Exception e) { @@ -200,55 +228,24 @@ public class OcrRecognizeDemo { @Test public void recognizeAndDraw2(){ try { - OcrCommonRecModel recModel = getRecModel(); + OcrCommonRecModel recModel = getFastRecModel(); int fontSize = 18; //创建保存路径 Path inputImagePath = Paths.get("src/main/resources/general_ocr_002.png"); - Path imageOutputPath = Paths.get("output/ocr_4_recognized.jpg"); - BufferedImage image = null; - image = ImageIO.read(new File(inputImagePath.toAbsolutePath().toString())); - BufferedImage resultImage = recModel.recognizeAndDraw(image, fontSize, new OcrRecOptions()); - ImageUtils.saveImage(resultImage, imageOutputPath.toAbsolutePath().toString()); + Path imageOutputPath = Paths.get("output/ocr_5_recognized.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(inputImagePath); + OcrInfo ocrInfo = recModel.recognizeAndDraw(image, fontSize, new OcrRecOptions()); + log.info("OCR识别结果:{}", JSONObject.toJSONString(ocrInfo)); + //保存绘制结果 + if(ocrInfo != null && ocrInfo.getDrawnImage() != null){ + ImageUtils.save(ocrInfo.getDrawnImage(), imageOutputPath.toAbsolutePath().toString()); + } } catch (Exception e) { e.printStackTrace(); } } - /** - * 文本识别并绘制结果(返回base64) - */ - @Test - public void recognizeAndDrawToBase64(){ - try { - OcrCommonRecModel recModel = getRecModel(); - int fontSize = 18; - //创建保存路径 - Path inputImagePath = Paths.get("src/main/resources/general_ocr_002.png"); - byte[] imageBytes = FileUtil.readBytes(inputImagePath); - String base64 = recModel.recognizeAndDrawToBase64(imageBytes, fontSize, new OcrRecOptions()); - log.info("base64:{}", base64); - } catch (Exception e) { - e.printStackTrace(); - } - } - - /** - * 文本识别并绘制结果(返回OcrInfo,OcrInfo中包含base64) - */ - @Test - public void recognizeAndDraw3(){ - try { - OcrCommonRecModel recModel = getRecModel(); - int fontSize = 18; - //创建保存路径 - Path inputImagePath = Paths.get("src/main/resources/general_ocr_002.png"); - byte[] imageBytes = FileUtil.readBytes(inputImagePath); - OcrInfo ocrInfo = recModel.recognizeAndDraw(imageBytes, fontSize, new OcrRecOptions()); - log.info("ocrInfo:{}", JsonUtils.toJson(ocrInfo)); - } catch (Exception e) { - e.printStackTrace(); - } - } /** * 批量识别 @@ -259,7 +256,7 @@ public class OcrRecognizeDemo { @Test public void batchRecognize(){ try { - OcrCommonRecModel recModel = getRecModelWithDirection(); + OcrCommonRecModel recModel = getFastRecModel(); //批量检测要求图片宽高一致 String folderPath = "/Users/xxx/Downloads/testing33"; //读取文件夹中所有图片 diff --git a/examples/ocr-examples/src/main/java/smartai/examples/ocr/plate/PlateRecDemo.java b/examples/ocr-examples/src/main/java/smartai/examples/ocr/plate/PlateRecDemo.java index b9319b7..0dd1947 100644 --- a/examples/ocr-examples/src/main/java/smartai/examples/ocr/plate/PlateRecDemo.java +++ b/examples/ocr-examples/src/main/java/smartai/examples/ocr/plate/PlateRecDemo.java @@ -1,7 +1,9 @@ package smartai.examples.ocr.plate; +import ai.djl.modality.cv.Image; import ai.djl.util.JsonUtils; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.utils.ImageUtils; @@ -38,6 +40,7 @@ public class PlateRecDemo { @BeforeClass public static void beforeAll() throws IOException { + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -68,6 +71,7 @@ public class PlateRecDemo { recModelConfig.setModelPath("/Users/wenjie/Documents/develop/model/plate/plate_rec_color.onnx"); //指定车牌检测模型 recModelConfig.setPlateDetModel(getPlateDetModel()); + recModelConfig.setDevice(device); return PlateModelFactory.getInstance().getRecModel(recModelConfig); } @@ -75,10 +79,12 @@ public class PlateRecDemo { * 车牌识别 */ @Test - public void testDetect() { + public void testDetect() throws IOException { PlateRecModel plateRecModel = getPlateRecModel(); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/plate/Quicker_20220930_180856.png"); //识别车号 - R> result = plateRecModel.recognize("src/main/resources/plate/Quicker_20220930_180856.png"); + R> result = plateRecModel.recognize(image); if(result.isSuccess()){ log.info("车牌识别结果:{}", JsonUtils.toJson(result.getData())); }else{ @@ -109,14 +115,14 @@ public class PlateRecDemo { public void recognizeAndDraw2() { try { PlateRecModel plateRecModel = getPlateRecModel(); - BufferedImage image = null; String imagePath = "src/main/resources/plate/Quicker_20220930_180856.png"; - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(imagePath); //可以根据后续业务场景使用detectedImage - R detectedImage = plateRecModel.recognizeAndDraw(image); + R detectedImage = plateRecModel.recognizeAndDraw(image); if(detectedImage.isSuccess()){ log.info("车牌识别成功"); - ImageUtils.saveImage(detectedImage.getData(), "output/plate_recognized2.jpg"); + ImageUtils.save(detectedImage.getData(), "output/plate_recognized3.jpg"); }else{ log.error("车牌识别失败:{}", detectedImage.getMessage()); } diff --git a/examples/ocr-examples/src/main/java/smartai/examples/ocr/table/TableRecDemo.java b/examples/ocr-examples/src/main/java/smartai/examples/ocr/table/TableRecDemo.java index c50ad68..942d1bc 100644 --- a/examples/ocr-examples/src/main/java/smartai/examples/ocr/table/TableRecDemo.java +++ b/examples/ocr-examples/src/main/java/smartai/examples/ocr/table/TableRecDemo.java @@ -3,6 +3,7 @@ package smartai.examples.ocr.table; import ai.djl.modality.cv.Image; import cn.hutool.core.io.FileUtil; import cn.smartjavaai.common.config.Config; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.utils.ImageUtils; @@ -47,6 +48,7 @@ public class TableRecDemo { @BeforeClass public static void beforeAll() throws IOException { + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -79,7 +81,6 @@ public class TableRecDemo { 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/xxx/Documents/develop/model/ocr/PP-OCRv5_server_det_infer/PP-OCRv5_server_det.onnx"); config.setDevice(device); return OcrModelFactory.getInstance().getDetModel(config); } @@ -115,45 +116,6 @@ public class TableRecDemo { - /** - * 表格识别 - * 仅支持简单表格 - * 流程:表格结构识别 -> 文本检测 -> 文本识别 -> 合成html table - * 注意事项: - * 1、批量检测时,模型应统一放在外层 try 中使用,避免重复加载,自动释放资源更安全。 - * 2、模型文件需要放在单独文件夹 - */ - @Test - public void recognize(){ - try { - TableStructureModel tableStructureModel = getTableStructureModel(); - OcrCommonDetModel detModel = getDetectionModel(); - OcrCommonRecModel recModel = getRecModel(); - OcrDirectionModel directionModel = getDirectionModel(); - //创建表格识别器 - TableRecognizer tableRecognizer = TableRecognizer.builder() - .withStructureModel(tableStructureModel) - .withTextDetModel(detModel) -// .withDirectionModel(getDirectionModel()) //如果表格中存在旋转的文字,可以使用方向分类模型 - .withTextRecModel(recModel).build(); - String imagePath = "src/main/resources/table/table_ch1.png"; - BufferedImage image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - R result = tableRecognizer.recognize(image); - if(result.isSuccess()){ - log.info("result: {}", result.getData().getHtml()); - //导出html内容到文件 - Path outputPath = Paths.get("output/table_ch2_result.html"); - FileUtil.writeUtf8String(result.getData().getHtml(), outputPath.toAbsolutePath().toString()); - //绘制表格结构 - tableRecognizer.drawTable(result.getData(), image, "output/table_ch2_result.jpg"); - //导出excel,如果导出失败,可能是因为表格结果识别的结果是错乱的 - tableRecognizer.exportExcel(result.getData().getHtml(), "output/table_ch2_result.xls"); - } - } catch (Exception e) { - e.printStackTrace(); - } - } - /** * 表格识别 @@ -177,7 +139,8 @@ public class TableRecDemo { // .withDirectionModel(getDirectionModel()) //如果表格中存在旋转的文字,可以使用方向分类模型 .withTextRecModel(recModel).build(); String imagePath = "src/main/resources/table/table_ch1.png"; - BufferedImage image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(imagePath); R result = tableRecognizer.recognize(image); if(result.isSuccess()){ log.info("result: {}", result.getData().getHtml()); @@ -185,8 +148,8 @@ public class TableRecDemo { Path outputPath = Paths.get("output/table_ch2_result.html"); FileUtil.writeUtf8String(result.getData().getHtml(), outputPath.toAbsolutePath().toString()); //绘制表格结构 - BufferedImage resultImage = tableRecognizer.drawTable(result.getData(), image); - ImageUtils.saveImage(resultImage, "output/table_ch2_result.jpg"); + Image resultImage = tableRecognizer.drawTable(result.getData(), image); + ImageUtils.save(resultImage, "output/table_ch2_result.jpg"); //导出excel,如果导出失败,可能是因为表格结果识别的结果是错乱的 try (OutputStream out = Files.newOutputStream(Paths.get("output/table_ch2_result2.xls"))) { tableRecognizer.exportExcel(result.getData().getHtml(), out); diff --git a/examples/speech-examples/pom.xml b/examples/speech-examples/pom.xml index 6fb00b8..6b5d89f 100644 --- a/examples/speech-examples/pom.xml +++ b/examples/speech-examples/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.24 + 1.0.25 smartai.examples.speech.asr.common.OcrRecognizeDemo diff --git a/examples/translation-example/pom.xml b/examples/translation-example/pom.xml index 656484f..2c7e935 100644 --- a/examples/translation-example/pom.xml +++ b/examples/translation-example/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.24 + 1.0.25 smartai.examples.nlp.translation.TranslationDemo diff --git a/examples/vision-example/pom.xml b/examples/vision-example/pom.xml index 79f6a4d..706928a 100644 --- a/examples/vision-example/pom.xml +++ b/examples/vision-example/pom.xml @@ -12,7 +12,7 @@ 11 11 UTF-8 - 1.0.24 + 1.0.25 smartai.examples.vision.ObjectDetectionDemo @@ -272,36 +272,6 @@ - - - org.bytedeco - javacpp - ${javacv.version} - ${javacv.platform.linux-arm64} - - - - org.bytedeco - ffmpeg - 6.1.1-1.5.10 - ${javacv.platform.linux-arm64} - - - - org.bytedeco - openblas - 0.3.26-1.5.10 - ${javacv.platform.linux-arm64} - - - - org.bytedeco - opencv - 4.9.0-1.5.10 - ${javacv.platform.linux-arm64} - - - diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/ActionRecognizeDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/ActionRecognizeDemo.java index 7566d34..e912f1f 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/ActionRecognizeDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/ActionRecognizeDemo.java @@ -32,6 +32,8 @@ public class ActionRecognizeDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/InstanceSegDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/InstanceSegDemo.java index 923252d..dc2caa1 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/InstanceSegDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/InstanceSegDemo.java @@ -41,6 +41,8 @@ public class InstanceSegDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -119,7 +121,7 @@ public class InstanceSegDemo { if(result.isSuccess()){ log.info("实例分割结果:{}", JSONObject.toJSONString(result.getData())); //保存图片 - ImageUtils.saveImage(result.getData().getDrawnImage(), "dog_bike_car_detected.png", "output"); + ImageUtils.save(result.getData().getDrawnImage(), "dog_bike_car_detected2.png", "output"); }else{ log.info("实例分割失败:{}", result.getMessage()); } diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/ObbDetDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/ObbDetDemo.java index 08a2703..18559c2 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/ObbDetDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/ObbDetDemo.java @@ -36,6 +36,8 @@ public class ObbDetDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -115,7 +117,7 @@ public class ObbDetDemo { if(result.isSuccess()){ log.info("旋转框检测结果:{}", JSONObject.toJSONString(result.getData())); //保存图片 - ImageUtils.saveImage(result.getData().getDrawnImage(), "boats_obb_detected.png", "output"); + ImageUtils.save(result.getData().getDrawnImage(), "output/boats_obb_detected2.png"); }else{ log.info("旋转框检测失败:{}", result.getMessage()); } diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/ObjectDetectionDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/ObjectDetectionDemo.java index e70c61f..50fd93c 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/ObjectDetectionDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/ObjectDetectionDemo.java @@ -5,6 +5,7 @@ import ai.djl.modality.cv.ImageFactory; import ai.djl.util.JsonUtils; import cn.hutool.core.date.LocalDateTimeUtil; import cn.hutool.core.lang.UUID; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; @@ -56,6 +57,8 @@ public class ObjectDetectionDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -90,7 +93,9 @@ public class ObjectDetectionDemo { public void objectDetection(){ try { DetectorModel detectorModel = getModel(); - DetectionResponse detectionResponse = detectorModel.detect("src/main/resources/object_detection.jpg"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/object_detection.jpg"); + DetectionResponse detectionResponse = detectorModel.detect(image); log.info("目标检测结果:{}", JSONObject.toJSONString(detectionResponse)); } catch (Exception e) { e.printStackTrace(); @@ -118,10 +123,14 @@ public class ObjectDetectionDemo { try { DetectorModel detectorModel = getModel(); String imagePath = "src/main/resources/object_detection.jpg"; - BufferedImage image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile(imagePath); //可以根据后续业务场景使用detectedImage - BufferedImage detectedImage = detectorModel.detectAndDraw(image); - Assert.assertNotNull("detectedImage null", detectedImage); + DetectionResponse detectionResponse = detectorModel.detectAndDraw(image); + log.info("目标检测结果:{}", JSONObject.toJSONString(detectionResponse)); + if(detectionResponse != null && detectionResponse.getDrawnImage() != null){ + ImageUtils.save(detectionResponse.getDrawnImage(), "output/object_detection_detected2.png"); + } } catch (Exception e) { e.printStackTrace(); } @@ -181,9 +190,9 @@ public class ObjectDetectionDemo { 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"); + //创建Image对象,可以从文件、url、InputStream创建、BufferedImage、Base64创建,具体使用方法可以查看文档 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/dog_bike_car.jpg"); + DetectionResponse detectionResponse = detectorModel.detect(image); log.info("目标检测结果:{}", JSONObject.toJSONString(detectionResponse)); } catch (Exception e) { e.printStackTrace(); @@ -218,11 +227,11 @@ public class ObjectDetectionDemo { log.info("时间:" + LocalDateTimeUtil.now().toString()); log.info("检测结果:{}", JsonUtils.toJson(detectionInfoList)); //绘制检测结果 - OpenCVUtils.drawRectAndText(image, detectionInfoList); + ImageUtils.drawRectAndText(image, detectionInfoList); //保存图片 - ImageUtils.saveImage(image, "test"+ UUID.fastUUID().toString() +".png","/Users/wenjie/Downloads"); + ImageUtils.save(image, "test"+ UUID.fastUUID().toString() +".png","/Users/wenjie/Downloads"); if (image != null){ - ((Mat)image.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(image); } } @@ -268,9 +277,12 @@ public class ObjectDetectionDemo { log.info("时间:" + LocalDateTimeUtil.now().toString()); log.info("检测结果:{}", JsonUtils.toJson(detectionInfoList)); //绘制检测结果 - OpenCVUtils.drawRectAndText(image, detectionInfoList); + ImageUtils.drawRectAndText(image, detectionInfoList); //保存图片 - ImageUtils.saveImage(image, "test"+ UUID.fastUUID().toString() +".png","/Users/wenjie/Downloads"); + ImageUtils.save(image, "test"+ UUID.fastUUID().toString() +".png","/Users/wenjie/Downloads"); + if (image != null){ + ImageUtils.releaseOpenCVMat(image); + } } @Override @@ -316,9 +328,12 @@ public class ObjectDetectionDemo { log.info("时间:" + LocalDateTimeUtil.now().toString()); log.info("检测结果:{}", JsonUtils.toJson(detectionInfoList)); //绘制检测结果 - OpenCVUtils.drawRectAndText(image, detectionInfoList); + ImageUtils.drawRectAndText(image, detectionInfoList); //保存图片 - ImageUtils.saveImage(image, "test"+ UUID.fastUUID().toString() +".png","/Users/wenjie/Downloads"); + ImageUtils.save(image, "test"+ UUID.fastUUID().toString() +".png","/Users/wenjie/Downloads"); + if (image != null){ + ImageUtils.releaseOpenCVMat(image); + } } @Override @@ -385,7 +400,7 @@ public class ObjectDetectionDemo { JOptionPane.showConfirmDialog(null, "Failed to capture image from WebCam."); } ViewerFrame frame = new ViewerFrame(width, height); - ImageFactory factory = ImageFactory.getInstance(); + SmartImageFactory factory = SmartImageFactory.getInstance(); Size size = new Size(width, height); while (capture.isOpened()) { @@ -394,9 +409,8 @@ public class ObjectDetectionDemo { } Mat resizeImage = new Mat(); Imgproc.resize(image, resizeImage, size); - Image img = factory.fromImage(resizeImage); - BufferedImage bufferedImage = OpenCVUtils.mat2Image(resizeImage); - DetectionResponse detectedResult = detectorModel.detect(bufferedImage); + Image img = factory.fromMat(resizeImage); + DetectionResponse detectedResult = detectorModel.detect(img); if (Objects.isNull(detectedResult) || Objects.isNull(detectedResult.getDetectionInfoList()) || detectedResult.getDetectionInfoList().size() == 0){ log.debug("未检测到物体"); continue; @@ -404,11 +418,10 @@ public class ObjectDetectionDemo { for(DetectionInfo detectionInfo : detectedResult.getDetectionInfoList()){ DetectionRectangle detectionRectangle = detectionInfo.getDetectionRectangle(); String text = detectionInfo.getObjectDetInfo().getClassName(); - ImageUtils.drawImageRectWithText(bufferedImage, detectionRectangle, text, Color.RED); + ImageUtils.drawRectAndText(img, detectionRectangle, text); } - frame.showImage(bufferedImage); + frame.showImage(ImageUtils.toBufferedImage(img)); } - capture.release(); System.exit(0); } catch (Exception e) { diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/PersonDetectDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/PersonDetectDemo.java index 26eb4bf..b1125f3 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/PersonDetectDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/PersonDetectDemo.java @@ -32,6 +32,8 @@ public class PersonDetectDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -108,7 +110,7 @@ public class PersonDetectDemo { if(result.isSuccess()){ log.info("行人检测结果:{}", JSONObject.toJSONString(result.getData())); //保存图片 - ImageUtils.saveImage(result.getData().getDrawnImage(), "person_result.png", "output"); + ImageUtils.save(result.getData().getDrawnImage(), "person_result.png", "output"); }else{ log.info("行人检测失败:{}", result.getMessage()); } diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/PoseDetDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/PoseDetDemo.java index 1101868..e49fd33 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/PoseDetDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/PoseDetDemo.java @@ -33,6 +33,8 @@ public class PoseDetDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -105,7 +107,7 @@ public class PoseDetDemo { //可以根据后续业务场景使用detectedImage Image drawImage = detectorModel.detectAndDraw(image); //保存图片 - ImageUtils.saveImage(drawImage, "pose_detected.png", "output"); + ImageUtils.save(drawImage, "pose_detected2.png", "output"); } catch (Exception e) { e.printStackTrace(); } diff --git a/examples/vision-example/src/main/java/smartai/examples/vision/SemSegDemo.java b/examples/vision-example/src/main/java/smartai/examples/vision/SemSegDemo.java index a786085..56cd50c 100644 --- a/examples/vision-example/src/main/java/smartai/examples/vision/SemSegDemo.java +++ b/examples/vision-example/src/main/java/smartai/examples/vision/SemSegDemo.java @@ -36,6 +36,8 @@ public class SemSegDemo { @BeforeClass public static void beforeAll() throws IOException { + //将图片处理的底层引擎切换为 OpenCV + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); //修改缓存路径 // Config.setCachePath("/Users/xxx/smartjavaai_cache"); } @@ -109,7 +111,7 @@ public class SemSegDemo { //可以根据后续业务场景使用detectedImage Image dretectedImage = detectorModel.detectAndDraw(image); //保存 - ImageUtils.saveImage(dretectedImage, "dog_bike_car_detected.png", "output"); + ImageUtils.save(dretectedImage, "dog_bike_car_detected2.png", "output"); } catch (Exception e) { e.printStackTrace(); } diff --git a/face/pom.xml b/face/pom.xml index fac0167..41e0da9 100644 --- a/face/pom.xml +++ b/face/pom.xml @@ -6,11 +6,11 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 face - 1.0.24 + 1.0.25 face SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java b/face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java index 562d1fc..49155d8 100644 --- a/face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java +++ b/face/src/main/java/cn/smartjavaai/face/enums/FaceRecModelEnum.java @@ -6,21 +6,61 @@ package cn.smartjavaai.face.enums; */ 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"); + FACENET_MODEL("PyTorch", 112, 112, 0.7f), + SEETA_FACE6_MODEL("c++", 0, 0, 0.62f), + SEETA_FACE6_LIGHT_MODEL("c++", 0, 0, 0.62f), + INSIGHT_FACE_IRSE50_MODEL("PyTorch", 112, 112, 0.62f), + INSIGHT_FACE_MOBILE_FACENET_MODEL("PyTorch", 112, 112, 0.64f), + ELASTIC_FACE_MODEL("PyTorch", 112, 112, 0.61f), + SPHERE_FACE_20A_ONNX("OnnxRuntime", 96, 112, 0.7f), + SPHERE_FACE_20A_PT("PyTorch", 96, 112, 0.7f), + DREAM_IJBA_RES18_NAIVE("OnnxRuntime", 224, 224, 0.74f), + EVOLVE_FACE_IR50("PyTorch", 112, 112, 0.62f), + EVOLVE_FACE_IR50_ASIA("PyTorch", 112, 112, 0.62f), + EVOLVE_FACE_IR152("PyTorch", 112, 112, 0.62f), + VGG_FACE("PyTorch", 224, 224, 0.75f); - private final String modelClassName; + /** + * 模型输入尺寸:宽 + */ + private final int inputWidth; - FaceRecModelEnum(String modelClassName) { - this.modelClassName = modelClassName; + /** + * 模型输入尺寸:高 + */ + private final int inputHeight; + + /** + * 模型引擎 + */ + private final String engine; + + /** + * 相似度阈值 + */ + private final float threshold; + + FaceRecModelEnum(String engine, int inputWidth, int inputHeight, float threshold) { + this.inputWidth = inputWidth; + this.inputHeight = inputHeight; + this.engine = engine; + this.threshold = threshold; } - public String getModelClassName() { - return modelClassName; + public int getInputWidth() { + return inputWidth; + } + + public int getInputHeight() { + return inputHeight; + } + + public String getEngine() { + return engine; + } + + public float getThreshold() { + return threshold; } /** diff --git a/face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java index d5327dd..ac48b57 100644 --- a/face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/ExpressionModelFactory.java @@ -90,6 +90,7 @@ public class ExpressionModelFactory { throw new FaceException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -101,4 +102,26 @@ public class ExpressionModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(ExpressionModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java index 695bd27..dae9506 100644 --- a/face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/FaceAttributeModelFactory.java @@ -2,6 +2,8 @@ package cn.smartjavaai.face.factory; import cn.smartjavaai.common.config.Config; import cn.smartjavaai.face.config.FaceAttributeConfig; +import cn.smartjavaai.face.enums.ExpressionModelEnum; +import cn.smartjavaai.face.enums.FaceAttributeModelEnum; import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.model.attribute.Seetaface6FaceAttributeModel; import lombok.extern.slf4j.Slf4j; @@ -21,12 +23,12 @@ public class FaceAttributeModelFactory { // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 private static volatile FaceAttributeModelFactory 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<>(); @@ -48,8 +50,8 @@ public class FaceAttributeModelFactory { * @param name * @param clazz */ - private static void registerModel(String name, Class clazz) { - registry.put(name.toLowerCase(), clazz); + private static void registerModel(FaceAttributeModelEnum faceAttributeModelEnum, Class clazz) { + registry.put(faceAttributeModelEnum, clazz); } @@ -62,7 +64,7 @@ public class FaceAttributeModelFactory { 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 createFaceModel(config); }); } @@ -73,9 +75,9 @@ public class FaceAttributeModelFactory { * @return */ private FaceAttributeModel createFaceModel(FaceAttributeConfig config) { - Class clazz = registry.get(config.getModelEnum().getModelClassName().toLowerCase()); + Class clazz = registry.get(config.getModelEnum()); if(clazz == null){ - throw new FaceException("Unsupported algorithm"); + throw new FaceException("Unsupported model"); } FaceAttributeModel model = null; try { @@ -84,14 +86,37 @@ public class FaceAttributeModelFactory { throw new FaceException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } // 初始化默认算法 static { - registerModel("seetaface6model", Seetaface6FaceAttributeModel.class); + registerModel(FaceAttributeModelEnum.SEETA_FACE6_MODEL, Seetaface6FaceAttributeModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(FaceAttributeModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java index 1e0f0a2..2150b53 100644 --- a/face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/FaceDetModelFactory.java @@ -3,6 +3,7 @@ package cn.smartjavaai.face.factory; import cn.smartjavaai.common.config.Config; import cn.smartjavaai.face.config.FaceDetConfig; import cn.smartjavaai.face.constant.FaceDetectConstant; +import cn.smartjavaai.face.enums.FaceAttributeModelEnum; import cn.smartjavaai.face.enums.FaceDetModelEnum; import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.model.facedect.CommonFaceDetModel; @@ -93,16 +94,17 @@ public class FaceDetModelFactory { private FaceDetModel createFaceDetModel(FaceDetConfig config) { Class clazz = registry.get(config.getModelEnum()); if(clazz == null){ - throw new FaceException("Unsupported algorithm"); + throw new FaceException("Unsupported model"); } - FaceDetModel algorithm = null; + FaceDetModel model = null; try { - algorithm = (FaceDetModel) clazz.newInstance(); + model = (FaceDetModel) clazz.newInstance(); } catch (InstantiationException | IllegalAccessException e) { throw new FaceException(e); } - algorithm.loadModel(config); - return algorithm; + model.loadModel(config); + model.setFromFactory(true); + return model; } @@ -133,4 +135,26 @@ public class FaceDetModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(FaceDetModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java index ac0d902..5b030b2 100644 --- a/face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/FaceQualityModelFactory.java @@ -2,6 +2,8 @@ package cn.smartjavaai.face.factory; import cn.smartjavaai.common.config.Config; import cn.smartjavaai.face.config.QualityConfig; +import cn.smartjavaai.face.enums.FaceDetModelEnum; +import cn.smartjavaai.face.enums.QualityModelEnum; import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.model.quality.FaceQualityModel; import cn.smartjavaai.face.model.quality.Seetaface6QualityModel; @@ -21,12 +23,12 @@ public class FaceQualityModelFactory { // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 private static volatile FaceQualityModelFactory 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<>(); @@ -45,11 +47,11 @@ public class FaceQualityModelFactory { /** * 注册模型 - * @param name + * @param qualityModelEnum * @param clazz */ - private static void registerModel(String name, Class clazz) { - registry.put(name.toLowerCase(), clazz); + private static void registerModel(QualityModelEnum qualityModelEnum, Class clazz) { + registry.put(qualityModelEnum, clazz); } @@ -62,7 +64,7 @@ public class FaceQualityModelFactory { 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 createFaceModel(config); }); } @@ -73,7 +75,7 @@ public class FaceQualityModelFactory { * @return */ private FaceQualityModel createFaceModel(QualityConfig config) { - Class clazz = registry.get(config.getModelEnum().getModelClassName().toLowerCase()); + Class clazz = registry.get(config.getModelEnum()); if(clazz == null){ throw new FaceException("Unsupported algorithm"); } @@ -84,14 +86,37 @@ public class FaceQualityModelFactory { throw new FaceException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } // 初始化默认算法 static { - registerModel("seetaface6model", Seetaface6QualityModel.class); + registerModel(QualityModelEnum.SEETA_FACE6_MODEL, Seetaface6QualityModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(QualityModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java index 25cf08c..99c815d 100644 --- a/face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/FaceRecModelFactory.java @@ -4,6 +4,7 @@ import cn.smartjavaai.common.config.Config; import cn.smartjavaai.face.config.FaceRecConfig; import cn.smartjavaai.face.constant.FaceDetectConstant; import cn.smartjavaai.face.enums.FaceRecModelEnum; +import cn.smartjavaai.face.enums.QualityModelEnum; import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.model.facerec.*; import lombok.extern.slf4j.Slf4j; @@ -79,14 +80,15 @@ public class FaceRecModelFactory { if(clazz == null){ throw new FaceException("Unsupported model"); } - FaceRecModel algorithm = null; + FaceRecModel model = null; try { - algorithm = (FaceRecModel) clazz.newInstance(); + model = (FaceRecModel) clazz.newInstance(); } catch (InstantiationException | IllegalAccessException e) { throw new FaceException(e); } - algorithm.loadModel(config); - return algorithm; + model.loadModel(config); + model.setFromFactory(true); + return model; } @@ -98,7 +100,36 @@ public class FaceRecModelFactory { registerAlgorithm(FaceRecModelEnum.ELASTIC_FACE_MODEL, CommonFaceRecModel.class); registerAlgorithm(FaceRecModelEnum.SEETA_FACE6_MODEL, SeetaFace6FaceRecModel.class); registerAlgorithm(FaceRecModelEnum.SEETA_FACE6_LIGHT_MODEL, SeetaFace6FaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.DREAM_IJBA_RES18_NAIVE, CommonFaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.VGG_FACE, CommonFaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.SPHERE_FACE_20A_ONNX, CommonFaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.SPHERE_FACE_20A_PT, CommonFaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.EVOLVE_FACE_IR50, CommonFaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.EVOLVE_FACE_IR50_ASIA, CommonFaceRecModel.class); + registerAlgorithm(FaceRecModelEnum.EVOLVE_FACE_IR152, CommonFaceRecModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(FaceRecModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java b/face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java index 6d2525a..84bcabc 100644 --- a/face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java +++ b/face/src/main/java/cn/smartjavaai/face/factory/LivenessModelFactory.java @@ -2,6 +2,7 @@ package cn.smartjavaai.face.factory; import cn.smartjavaai.common.config.Config; import cn.smartjavaai.face.config.LivenessConfig; +import cn.smartjavaai.face.enums.FaceRecModelEnum; import cn.smartjavaai.face.enums.LivenessModelEnum; import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.model.liveness.CommonLivenessModel; @@ -87,6 +88,7 @@ public class LivenessModelFactory { throw new FaceException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -99,4 +101,26 @@ public class LivenessModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(LivenessModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java b/face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java index df00753..4e8f341 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/attribute/FaceAttributeModel.java @@ -1,5 +1,6 @@ package cn.smartjavaai.face.model.attribute; +import ai.djl.modality.cv.Image; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.face.FaceAttribute; @@ -22,12 +23,12 @@ public interface FaceAttributeModel extends AutoCloseable{ void loadModel(FaceAttributeConfig config); // 加载模型 - /** * 人脸属性识别(多人脸) * @param imagePath 图片路径 * @return */ + @Deprecated default DetectionResponse detect(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -37,6 +38,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default DetectionResponse detect(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -46,6 +48,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param imageData 图片字节流 * @return */ + @Deprecated default DetectionResponse detect(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -56,6 +59,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default List detect(String imagePath, DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -66,6 +70,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default FaceAttribute detect(String imagePath, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -76,6 +81,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default List detect(byte[] imageData,DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -86,6 +92,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default FaceAttribute detect(byte[] imageData, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -97,6 +104,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default List detect(BufferedImage image,DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -107,6 +115,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default FaceAttribute detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -117,6 +126,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param image * @return */ + @Deprecated default FaceAttribute detectTopFace(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -127,6 +137,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param imagePath * @return */ + @Deprecated default FaceAttribute detectTopFace(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -136,6 +147,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default FaceAttribute detectTopFace(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -145,6 +157,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param image * @return */ + @Deprecated default FaceAttribute detectCropedFace(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -154,6 +167,7 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param imagePath * @return */ + @Deprecated default FaceAttribute detectCropedFace(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -163,17 +177,69 @@ public interface FaceAttributeModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default FaceAttribute detectCropedFace(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 人脸属性识别(多人脸) + * @param image + * @return + */ + default DetectionResponse detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 人脸属性识别(多人脸) + * @param image + * @param faceDetectionResponse 人脸检测结果 + * @return + */ + default List detect(Image image, DetectionResponse faceDetectionResponse){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 人脸属性识别(单人脸) + * @param image + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return + */ + default FaceAttribute detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 人脸属性识别(分数最高人脸) + * @param image + * @return + */ + default FaceAttribute detectTopFace(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + + /** + * 人脸属性识别(裁剪后的人脸) + * @param image + * @return + */ + default FaceAttribute detectCropedFace(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + - - - } diff --git a/face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java b/face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java index c2a92c3..29f6bfa 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/attribute/Seetaface6FaceAttributeModel.java @@ -1,12 +1,15 @@ package cn.smartjavaai.face.model.attribute; import ai.djl.engine.Engine; +import ai.djl.modality.cv.Image; import cn.smartjavaai.common.entity.*; +import cn.smartjavaai.common.entity.Point; import cn.smartjavaai.common.entity.face.FaceAttribute; import cn.smartjavaai.common.entity.face.FaceInfo; import cn.smartjavaai.common.entity.face.HeadPose; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.face.EyeStatus; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.PoolUtils; @@ -14,8 +17,10 @@ import cn.smartjavaai.face.config.FaceAttributeConfig; import cn.smartjavaai.common.enums.face.GenderType; import cn.smartjavaai.face.context.PredictorContext; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.FaceAttributeModelFactory; import cn.smartjavaai.face.seetaface.NativeLoader; import cn.smartjavaai.face.utils.FaceUtils; +import cn.smartjavaai.face.utils.Seetaface6Utils; import com.seeta.pool.*; import com.seeta.sdk.*; import lombok.extern.slf4j.Slf4j; @@ -150,7 +155,7 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { @Override public DetectionResponse detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new FaceException("图像无效"); } //创建推力器上下文 @@ -168,7 +173,7 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { predictorContext.eyeStateDetector = config.isEnableEyeStatus() ? eyeStateDetectorPool.borrowObject() : null; predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); + imageData.data = BufferedImageUtils.getMatrixBGR(image); //检测人脸 SeetaRect[] seetaResult = detectPredictor.Detect(imageData); if(Objects.isNull(seetaResult)){ @@ -182,7 +187,7 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { FaceAttribute faceAttribute = detect(imageData, seetaRect, landmarks, predictorContext); faceAttributeList.add(faceAttribute); } - return FaceUtils.convertToFaceAttributeResponse(seetaResult, seetaPointFSList, faceAttributeList); + return Seetaface6Utils.convertToFaceAttributeResponse(seetaResult, seetaPointFSList, faceAttributeList); } catch (Exception e) { throw new FaceException("人脸属性检测错误", e); } finally { @@ -212,15 +217,15 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { if (config.isEnableGender()){ GenderPredictor.GENDER[] gender = new GenderPredictor.GENDER[1]; boolean isSuccess = predictorContext.genderPredictor.PredictGenderWithCrop(imageData, landmarks, gender); - genderType = isSuccess ? FaceUtils.convertToGenderType(gender[0]) : GenderType.UNKNOWN; + genderType = isSuccess ? Seetaface6Utils.convertToGenderType(gender[0]) : GenderType.UNKNOWN; } //眼睛状态检测 EyeStatus leftEyeStatus = null; EyeStatus rightEyeStatus = null; if (config.isEnableEyeStatus()){ EyeStateDetector.EYE_STATE[] eyeState = predictorContext.eyeStateDetector.detect(imageData, landmarks); - leftEyeStatus = FaceUtils.convertToEyeStatus(eyeState[0]); - rightEyeStatus = FaceUtils.convertToEyeStatus(eyeState[1]); + leftEyeStatus = Seetaface6Utils.convertToEyeStatus(eyeState[0]); + rightEyeStatus = Seetaface6Utils.convertToEyeStatus(eyeState[1]); } //年龄检测 Integer age = 0; @@ -278,7 +283,7 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { @Override public List detect(BufferedImage image, DetectionResponse faceDetectionResponse) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new FaceException("图像无效"); } if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ @@ -297,8 +302,8 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; for(DetectionInfo detectionInfo : faceDetectionResponse.getDetectionInfoList()){ SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(detectionInfo.getDetectionRectangle()); + imageData.data = BufferedImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(detectionInfo.getDetectionRectangle()); SeetaPointF[] landmarks = null; FaceInfo faceInfo = detectionInfo.getFaceInfo(); //如果没有人脸标识,则提取人脸标识 @@ -307,7 +312,7 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { landmarks = new SeetaPointF[faceLandmarker.number()]; faceLandmarker.mark(imageData, seetaRect, landmarks); }else{ - landmarks = FaceUtils.convertToSeetaPointF(faceInfo.getKeyPoints()); + landmarks = Seetaface6Utils.convertToSeetaPointF(faceInfo.getKeyPoints()); } //人脸属性检测 FaceAttribute faceAttribute = detect(imageData, seetaRect, landmarks, predictorContext); @@ -365,7 +370,7 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { @Override public FaceAttribute detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new FaceException("图像无效"); } if(Objects.isNull(faceDetectionRectangle)){ @@ -380,13 +385,13 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { predictorContext.eyeStateDetector = config.isEnableEyeStatus() ? eyeStateDetectorPool.borrowObject() : null; predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(faceDetectionRectangle); + imageData.data = BufferedImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(faceDetectionRectangle); SeetaPointF[] landmarks = null; if(keyPoints == null || keyPoints.isEmpty()){ throw new FaceException("人脸关键点keyPoints为空"); } - landmarks = FaceUtils.convertToSeetaPointF(keyPoints); + landmarks = Seetaface6Utils.convertToSeetaPointF(keyPoints); //人脸属性检测 FaceAttribute faceAttribute = detect(imageData, seetaRect, landmarks, predictorContext); return faceAttribute; @@ -431,10 +436,196 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { @Override public FaceAttribute detectTopFace(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new FaceException("图像无效"); } + FaceLandmarker faceLandmarker = null; + FaceDetector detectPredictor = null; + //创建推力器上下文 + PredictorContext predictorContext = new PredictorContext(); + try { + detectPredictor = faceDetectorPool.borrowObject(); + faceLandmarker = faceLandmarkerPool.borrowObject(); + predictorContext.genderPredictor = config.isEnableGender() ? genderPredictorPool.borrowObject() : null; + predictorContext.agePredictor = config.isEnableAge() ? agePredictorPool.borrowObject() : null; + predictorContext.maskDetector = config.isEnableMask() ? maskDetectorPool.borrowObject() : null; + predictorContext.eyeStateDetector = config.isEnableEyeStatus() ? eyeStateDetectorPool.borrowObject() : null; + predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = BufferedImageUtils.getMatrixBGR(image); + //检测人脸 + SeetaRect[] seetaResult = detectPredictor.Detect(imageData); + if(Objects.isNull(seetaResult)){ + throw new FaceException("无人脸数据"); + } + SeetaPointF[] landmarks = new SeetaPointF[faceLandmarker.number()]; + faceLandmarker.mark(imageData, seetaResult[0], landmarks); + //人脸属性检测 + FaceAttribute faceAttribute = detect(imageData, seetaResult[0], landmarks, predictorContext); + return faceAttribute; + } catch (Exception e) { + throw new FaceException("活体检测错误", e); + } finally { + if (detectPredictor != null) { + try { + faceDetectorPool.returnObject(detectPredictor); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + + if (faceLandmarker != null) { + try { + faceLandmarkerPool.returnObject(faceLandmarker); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + PoolUtils.returnToPool(genderPredictorPool, predictorContext.genderPredictor); + PoolUtils.returnToPool(agePredictorPool, predictorContext.agePredictor); + PoolUtils.returnToPool(maskDetectorPool, predictorContext.maskDetector); + PoolUtils.returnToPool(eyeStateDetectorPool, predictorContext.eyeStateDetector); + PoolUtils.returnToPool(poseEstimatorPool, predictorContext.poseEstimator); + } + } + + @Override + public DetectionResponse detect(Image image) { + //创建推力器上下文 + PredictorContext predictorContext = new PredictorContext(); + FaceLandmarker faceLandmarker = null; + FaceDetector detectPredictor = null; + List seetaPointFSList = new ArrayList(); + List faceAttributeList = new ArrayList(); + try { + detectPredictor = faceDetectorPool.borrowObject(); + faceLandmarker = faceLandmarkerPool.borrowObject(); + predictorContext.genderPredictor = config.isEnableGender() ? genderPredictorPool.borrowObject() : null; + predictorContext.agePredictor = config.isEnableAge() ? agePredictorPool.borrowObject() : null; + predictorContext.maskDetector = config.isEnableMask() ? maskDetectorPool.borrowObject() : null; + predictorContext.eyeStateDetector = config.isEnableEyeStatus() ? eyeStateDetectorPool.borrowObject() : null; + predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + //检测人脸 + SeetaRect[] seetaResult = detectPredictor.Detect(imageData); + if(Objects.isNull(seetaResult)){ + throw new FaceException("无人脸数据"); + } + for(SeetaRect seetaRect : seetaResult){ + SeetaPointF[] landmarks = new SeetaPointF[faceLandmarker.number()]; + faceLandmarker.mark(imageData, seetaRect, landmarks); + seetaPointFSList.add(landmarks); + //人脸属性检测 + FaceAttribute faceAttribute = detect(imageData, seetaRect, landmarks, predictorContext); + faceAttributeList.add(faceAttribute); + } + return Seetaface6Utils.convertToFaceAttributeResponse(seetaResult, seetaPointFSList, faceAttributeList); + } catch (Exception e) { + throw new FaceException("人脸属性检测错误", e); + } finally { + // 统一归还所有 Predictor 到池 + PoolUtils.returnToPool(faceDetectorPool, detectPredictor); + PoolUtils.returnToPool(faceLandmarkerPool, faceLandmarker); + PoolUtils.returnToPool(genderPredictorPool, predictorContext.genderPredictor); + PoolUtils.returnToPool(agePredictorPool, predictorContext.agePredictor); + PoolUtils.returnToPool(maskDetectorPool, predictorContext.maskDetector); + PoolUtils.returnToPool(eyeStateDetectorPool, predictorContext.eyeStateDetector); + PoolUtils.returnToPool(poseEstimatorPool, predictorContext.poseEstimator); + } + } + + @Override + public List detect(Image image, DetectionResponse faceDetectionResponse) { + if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ + throw new FaceException("无人脸数据"); + } + //创建推力器上下文 + PredictorContext predictorContext = new PredictorContext(); + FaceLandmarker faceLandmarker = null; + List faceAttributeList = new ArrayList(); + try { + faceLandmarker = faceLandmarkerPool.borrowObject(); + predictorContext.genderPredictor = config.isEnableGender() ? genderPredictorPool.borrowObject() : null; + predictorContext.agePredictor = config.isEnableAge() ? agePredictorPool.borrowObject() : null; + predictorContext.maskDetector = config.isEnableMask() ? maskDetectorPool.borrowObject() : null; + predictorContext.eyeStateDetector = config.isEnableEyeStatus() ? eyeStateDetectorPool.borrowObject() : null; + predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; + for(DetectionInfo detectionInfo : faceDetectionResponse.getDetectionInfoList()){ + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(detectionInfo.getDetectionRectangle()); + SeetaPointF[] landmarks = null; + FaceInfo faceInfo = detectionInfo.getFaceInfo(); + //如果没有人脸标识,则提取人脸标识 + if(faceInfo == null || faceInfo.getKeyPoints() == null || faceInfo.getKeyPoints().isEmpty()){ + //提取人脸的5点人脸标识 + landmarks = new SeetaPointF[faceLandmarker.number()]; + faceLandmarker.mark(imageData, seetaRect, landmarks); + }else{ + landmarks = Seetaface6Utils.convertToSeetaPointF(faceInfo.getKeyPoints()); + } + //人脸属性检测 + FaceAttribute faceAttribute = detect(imageData, seetaRect, landmarks, predictorContext); + faceAttributeList.add(faceAttribute); + } + } catch (Exception e) { + throw new FaceException("活体检测错误", e); + } finally { + if (faceLandmarker != null) { + try { + faceLandmarkerPool.returnObject(faceLandmarker); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + PoolUtils.returnToPool(genderPredictorPool, predictorContext.genderPredictor); + PoolUtils.returnToPool(agePredictorPool, predictorContext.agePredictor); + PoolUtils.returnToPool(maskDetectorPool, predictorContext.maskDetector); + PoolUtils.returnToPool(eyeStateDetectorPool, predictorContext.eyeStateDetector); + PoolUtils.returnToPool(poseEstimatorPool, predictorContext.poseEstimator); + } + return faceAttributeList; + } + + @Override + public FaceAttribute detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints) { + if(Objects.isNull(faceDetectionRectangle)){ + throw new FaceException("无人脸数据"); + } + //创建推力器上下文 + PredictorContext predictorContext = new PredictorContext(); + try { + predictorContext.genderPredictor = config.isEnableGender() ? genderPredictorPool.borrowObject() : null; + predictorContext.agePredictor = config.isEnableAge() ? agePredictorPool.borrowObject() : null; + predictorContext.maskDetector = config.isEnableMask() ? maskDetectorPool.borrowObject() : null; + predictorContext.eyeStateDetector = config.isEnableEyeStatus() ? eyeStateDetectorPool.borrowObject() : null; + predictorContext.poseEstimator = config.isEnableHeadPose() ? poseEstimatorPool.borrowObject() : null; + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(faceDetectionRectangle); + SeetaPointF[] landmarks = null; + if(keyPoints == null || keyPoints.isEmpty()){ + throw new FaceException("人脸关键点keyPoints为空"); + } + landmarks = Seetaface6Utils.convertToSeetaPointF(keyPoints); + //人脸属性检测 + FaceAttribute faceAttribute = detect(imageData, seetaRect, landmarks, predictorContext); + return faceAttribute; + } catch (Exception e) { + throw new FaceException("活体检测错误", e); + } finally { + PoolUtils.returnToPool(genderPredictorPool, predictorContext.genderPredictor); + PoolUtils.returnToPool(agePredictorPool, predictorContext.agePredictor); + PoolUtils.returnToPool(maskDetectorPool, predictorContext.maskDetector); + PoolUtils.returnToPool(eyeStateDetectorPool, predictorContext.eyeStateDetector); + PoolUtils.returnToPool(poseEstimatorPool, predictorContext.poseEstimator); + } + } + + @Override + public FaceAttribute detectTopFace(Image image) { FaceLandmarker faceLandmarker = null; FaceDetector detectPredictor = null; //创建推力器上下文 @@ -485,7 +676,6 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { } } - public FaceDetectorPool getFaceDetectorPool() { return faceDetectorPool; } @@ -514,8 +704,20 @@ public class Seetaface6FaceAttributeModel implements FaceAttributeModel { return poseEstimatorPool; } + private boolean fromFactory = false; + + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + @Override public void close() throws Exception { + if (fromFactory) { + FaceAttributeModelFactory.removeFromCache(config.getModelEnum()); + } if(Objects.nonNull(faceDetectorPool)){ faceDetectorPool.close(); } diff --git a/face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java b/face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java index 06ef818..bcc242f 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/expression/CommonEmotionModel.java @@ -1,47 +1,33 @@ package cn.smartjavaai.face.model.expression; -import ai.djl.Device; 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.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 cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.face.ExpressionResult; import cn.smartjavaai.common.entity.face.FaceInfo; -import cn.smartjavaai.common.entity.face.LivenessResult; -import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.face.FacialExpression; 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.common.utils.*; import cn.smartjavaai.face.config.FaceExpressionConfig; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.ExpressionModelFactory; import cn.smartjavaai.face.model.expression.criterial.EmotionCriteriaFactory; -import cn.smartjavaai.face.model.expression.translator.DenseNetEmotionTranslator; -import cn.smartjavaai.face.preprocess.DJLImagePreprocessor; +import cn.smartjavaai.face.preprocess.DJLImageFacePreprocessor; import cn.smartjavaai.face.utils.FaceUtils; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; -import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.impl.GenericObjectPool; -import org.opencv.core.Mat; -import org.opencv.face.Face; -import javax.imageio.ImageIO; import java.awt.image.BufferedImage; -import java.io.ByteArrayInputStream; -import java.io.File; import java.io.IOException; -import java.nio.file.Paths; import java.util.ArrayList; import java.util.List; import java.util.Objects; @@ -68,6 +54,9 @@ public class CommonEmotionModel implements ExpressionModel{ if(StringUtils.isBlank(config.getModelPath())){ throw new FaceException("modelPath为空"); } + if(Objects.isNull(config.getDetectModel())){ + throw new FaceException("未指定人脸检测模型"); + } this.config = config; @@ -92,7 +81,7 @@ public class CommonEmotionModel implements ExpressionModel{ Predictor predictor = null; try (NDManager manager = model.getNDManager().newSubManager()){ predictor = predictorPool.borrowObject(); - DJLImagePreprocessor imagePreprocessor = new DJLImagePreprocessor(image, manager); + DJLImageFacePreprocessor imagePreprocessor = new DJLImageFacePreprocessor(image, manager); Image faceImg = image; if(config.isAlign()){ //仿射变换 @@ -131,40 +120,24 @@ public class CommonEmotionModel implements ExpressionModel{ if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = detect(image); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return detect(image); } @Override public R detect(BufferedImage image) { - if(Objects.isNull(config.getDetectModel())){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "未指定检测模型"); - } - R faceDetectionResponse = config.getDetectModel().detect(image); - if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ - return R.fail(R.Status.NO_FACE_DETECTED); - } - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); - for(DetectionInfo detectionInfo : faceDetectionResponse.getData().getDetectionInfoList()){ - FaceInfo faceInfo = detectionInfo.getFaceInfo(); - if(Objects.isNull(faceInfo) || Objects.isNull(faceInfo.getKeyPoints())){ - return R.fail(R.Status.Unknown.getCode(), "未检测到人脸关键点"); - } - Classifications classifications = detectCore(djlImage, detectionInfo.getDetectionRectangle(), faceInfo.getKeyPoints()); - Classifications.Classification bestClass = classifications.best(); - FacialExpression expression = FacialExpression.fromLabel(bestClass.getClassName()); - ExpressionResult result = new ExpressionResult(expression, (float)bestClass.getProbability()); - result.setClassifications(classifications); - faceInfo.setExpressionResult(result); - } - ((Mat)djlImage.getWrappedImage()).release(); - return faceDetectionResponse; + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = detect(imageDjl); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -172,11 +145,15 @@ public class CommonEmotionModel implements ExpressionModel{ if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData))); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = detect(imageDjl); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -184,8 +161,16 @@ public class CommonEmotionModel implements ExpressionModel{ if(StringUtils.isBlank(base64Image)){ return R.fail(R.Status.INVALID_IMAGE); } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detect(imageData); + Image image = null; + try { + image = SmartImageFactory.getInstance().fromBase64(base64Image); + R detectionResponseR = detect(image); + return detectionResponseR; + } catch (IOException e) { + throw new FaceException("无效图片", e); + }finally { + ImageUtils.releaseOpenCVMat(image); + } } @Override @@ -193,14 +178,16 @@ public class CommonEmotionModel implements ExpressionModel{ if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R> detectionResponseR = detect(image, faceDetectionResponse); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return detect(image, faceDetectionResponse); } @Override @@ -208,37 +195,23 @@ public class CommonEmotionModel implements ExpressionModel{ if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionResponse); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R> detectionResponseR = detect(imageDjl, faceDetectionResponse); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override public R> detect(BufferedImage image, DetectionResponse faceDetectionResponse) { - if(!ImageUtils.isImageValid(image)){ - R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ - R.fail(R.Status.NO_FACE_DETECTED); - } - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); - List expressionResults = new ArrayList<>(); - for(DetectionInfo detectionInfo : faceDetectionResponse.getDetectionInfoList()){ - FaceInfo faceInfo = detectionInfo.getFaceInfo(); - if(Objects.isNull(faceInfo) || Objects.isNull(faceInfo.getKeyPoints())){ - return R.fail(R.Status.Unknown.getCode(), "未检测到人脸关键点"); - } - Classifications classifications = detectCore(djlImage, detectionInfo.getDetectionRectangle(), faceInfo.getKeyPoints()); - Classifications.Classification bestClass = classifications.best(); - FacialExpression expression = FacialExpression.fromLabel(bestClass.getClassName()); - ExpressionResult result = new ExpressionResult(expression, (float)bestClass.getProbability()); - result.setClassifications(classifications); - expressionResults.add(result); - } - ((Mat)djlImage.getWrappedImage()).release(); - return R.ok(expressionResults); + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R> detectionResponseR = detect(imageDjl, faceDetectionResponse); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -246,8 +219,16 @@ public class CommonEmotionModel implements ExpressionModel{ if(StringUtils.isBlank(base64Image)){ return R.fail(R.Status.INVALID_IMAGE); } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detect(imageData, faceDetectionResponse); + Image image = null; + try { + image = SmartImageFactory.getInstance().fromBase64(base64Image); + R> detectionResponseR = detect(image, faceDetectionResponse); + return detectionResponseR; + } catch (IOException e) { + throw new FaceException("无效图片", e); + }finally { + ImageUtils.releaseOpenCVMat(image); + } } @Override @@ -255,14 +236,16 @@ public class CommonEmotionModel implements ExpressionModel{ if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = detect(image, faceDetectionRectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return detect(image, faceDetectionRectangle, keyPoints); } @Override @@ -270,26 +253,26 @@ public class CommonEmotionModel implements ExpressionModel{ if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionRectangle, keyPoints); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = detect(imageDjl, faceDetectionRectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override public R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); - Classifications classifications = detectCore(djlImage, faceDetectionRectangle, keyPoints); - Classifications.Classification bestClass = classifications.best(); - FacialExpression expression = FacialExpression.fromLabel(bestClass.getClassName()); - ExpressionResult result = new ExpressionResult(expression, (float)bestClass.getProbability()); - result.setClassifications(classifications); - ((Mat)djlImage.getWrappedImage()).release(); - return R.ok(result); + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = detect(imageDjl, faceDetectionRectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @@ -298,15 +281,134 @@ public class CommonEmotionModel implements ExpressionModel{ if(StringUtils.isBlank(base64Image)){ return R.fail(R.Status.INVALID_IMAGE); } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detect(imageData, faceDetectionRectangle, keyPoints); + Image image = null; + try { + image = SmartImageFactory.getInstance().fromBase64(base64Image); + R detectionResponseR = detect(image, faceDetectionRectangle, keyPoints); + return detectionResponseR; + } catch (IOException e) { + throw new FaceException("无效图片", e); + }finally { + ImageUtils.releaseOpenCVMat(image); + } } @Override public R detectTopFace(BufferedImage image) { + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = detectTopFace(imageDjl); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; + } + + @Override + public R detectTopFace(String imagePath) { + if(!FileUtils.isFileExists(imagePath)){ + return R.fail(R.Status.FILE_NOT_FOUND); + } + Image image = null; + try { + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = detectTopFace(image); + return detectionResponseR; + } catch (IOException e) { + throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); + } + } + + @Override + public R detectTopFace(byte[] imageData) { + if(Objects.isNull(imageData)){ + return R.fail(R.Status.INVALID_IMAGE); + } + Image imageDjl = null; + try { + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); + } catch (IOException e) { + throw new RuntimeException(e); + } + R detectionResponseR = detectTopFace(imageDjl); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; + } + + @Override + public R detectTopFaceBase64(String base64Image) { + if(StringUtils.isBlank(base64Image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + Image image = null; + try { + image = SmartImageFactory.getInstance().fromBase64(base64Image); + R detectionResponseR = detectTopFace(image); + return detectionResponseR; + } catch (IOException e) { + throw new FaceException("无效图片", e); + }finally { + ImageUtils.releaseOpenCVMat(image); + } + } + + + @Override + public R detect(Image image) { if(Objects.isNull(config.getDetectModel())){ return R.fail(R.Status.PARAM_ERROR.getCode(), "未指定检测模型"); } + R faceDetectionResponse = config.getDetectModel().detect(image); + if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + for(DetectionInfo detectionInfo : faceDetectionResponse.getData().getDetectionInfoList()){ + FaceInfo faceInfo = detectionInfo.getFaceInfo(); + if(Objects.isNull(faceInfo) || Objects.isNull(faceInfo.getKeyPoints())){ + return R.fail(R.Status.Unknown.getCode(), "未检测到人脸关键点"); + } + Classifications classifications = detectCore(image, detectionInfo.getDetectionRectangle(), faceInfo.getKeyPoints()); + Classifications.Classification bestClass = classifications.best(); + FacialExpression expression = FacialExpression.fromLabel(bestClass.getClassName()); + ExpressionResult result = new ExpressionResult(expression, (float)bestClass.getProbability()); + result.setClassifications(classifications); + faceInfo.setExpressionResult(result); + } + return faceDetectionResponse; + } + + @Override + public R> detect(Image image, DetectionResponse faceDetectionResponse) { + if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ + R.fail(R.Status.NO_FACE_DETECTED); + } + List expressionResults = new ArrayList<>(); + for(DetectionInfo detectionInfo : faceDetectionResponse.getDetectionInfoList()){ + FaceInfo faceInfo = detectionInfo.getFaceInfo(); + if(Objects.isNull(faceInfo) || Objects.isNull(faceInfo.getKeyPoints())){ + return R.fail(R.Status.Unknown.getCode(), "未检测到人脸关键点"); + } + Classifications classifications = detectCore(image, detectionInfo.getDetectionRectangle(), faceInfo.getKeyPoints()); + Classifications.Classification bestClass = classifications.best(); + FacialExpression expression = FacialExpression.fromLabel(bestClass.getClassName()); + ExpressionResult result = new ExpressionResult(expression, (float)bestClass.getProbability()); + result.setClassifications(classifications); + expressionResults.add(result); + } + return R.ok(expressionResults); + } + + @Override + public R detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints) { + Classifications classifications = detectCore(image, faceDetectionRectangle, keyPoints); + Classifications.Classification bestClass = classifications.best(); + FacialExpression expression = FacialExpression.fromLabel(bestClass.getClassName()); + ExpressionResult result = new ExpressionResult(expression, (float)bestClass.getProbability()); + result.setClassifications(classifications); + return R.ok(result); + } + + @Override + public R detectTopFace(Image image) { R faceDetectionResponse = config.getDetectModel().detect(image); if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ return R.fail(R.Status.NO_FACE_DETECTED); @@ -319,49 +421,26 @@ public class CommonEmotionModel implements ExpressionModel{ return detect(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getKeyPoints()); } - @Override - public R detectTopFace(String imagePath) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detectTopFace(image); - } - - @Override - public R detectTopFace(byte[] imageData) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detectTopFace(ImageIO.read(new ByteArrayInputStream(imageData))); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @Override - public R detectTopFaceBase64(String base64Image) { - if(StringUtils.isBlank(base64Image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detectTopFace(imageData); - } - @Override public GenericObjectPool> getPool() { return predictorPool; } + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + @Override public void close() { + if (fromFactory) { + ExpressionModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); diff --git a/face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java b/face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java index d5e4e16..6160b99 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/expression/ExpressionModel.java @@ -35,6 +35,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R detect(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -44,6 +45,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default R detect(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -53,6 +55,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param imageData 图片字节流 * @return */ + @Deprecated default R detect(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -63,6 +66,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param base64Image * @return */ + @Deprecated default R detectBase64(String base64Image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -74,6 +78,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default R> detect(String imagePath, DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -85,6 +90,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default R> detect(byte[] imageData,DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -95,6 +101,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default R> detect(BufferedImage image,DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -105,6 +112,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionResponse 人脸检测结果 * @return */ + @Deprecated default R> detectBase64(String base64Image,DetectionResponse faceDetectionResponse){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -116,6 +124,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default R detect(String imagePath, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -127,6 +136,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default R detect(byte[] imageData, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -140,6 +150,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -150,6 +161,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param faceDetectionRectangle 人脸检测结果-人脸框 * @return */ + @Deprecated default R detectBase64(String base64Image, DetectionRectangle faceDetectionRectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -160,6 +172,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param image * @return */ + @Deprecated default R detectTopFace(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -170,6 +183,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param imagePath * @return */ + @Deprecated default R detectTopFace(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -179,6 +193,7 @@ public interface ExpressionModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default R detectTopFace(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -189,15 +204,59 @@ public interface ExpressionModel extends AutoCloseable{ * @param base64Image * @return */ + @Deprecated default R detectTopFaceBase64(String base64Image){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 表情识别(多人脸) + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + + /** + * 表情识别(多人脸) + * @param image + * @param faceDetectionResponse 人脸检测结果 + * @return + */ + default R> detect(Image image, DetectionResponse faceDetectionResponse){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 表情识别(单人脸) + * @param image + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return + */ + default R detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 表情识别(分数最高人脸) + * @param image + * @return + */ + default R detectTopFace(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + default GenericObjectPool> getPool(){ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java index b79710d..a1dc181 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/CommonFaceDetModel.java @@ -9,15 +9,15 @@ 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.common.utils.Base64ImageUtils; -import cn.smartjavaai.common.utils.FileUtils; -import cn.smartjavaai.common.utils.ImageUtils; -import cn.smartjavaai.common.utils.OpenCVUtils; +import cn.smartjavaai.common.utils.*; import cn.smartjavaai.face.config.FaceDetConfig; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.ExpressionModelFactory; +import cn.smartjavaai.face.factory.FaceDetModelFactory; import cn.smartjavaai.face.model.facedect.criterial.FaceDetCriteriaFactory; import cn.smartjavaai.face.utils.FaceUtils; import lombok.extern.slf4j.Slf4j; @@ -48,6 +48,8 @@ public class CommonFaceDetModel implements FaceDetModel{ private ZooModel model; + private FaceDetConfig config; + /** * 加载模型 @@ -57,6 +59,7 @@ public class CommonFaceDetModel implements FaceDetModel{ public void loadModel(FaceDetConfig config){ Criteria criteria = FaceDetCriteriaFactory.createCriteria(config); try { + this.config = config; model = criteria.loadModel(); this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model)); int predictorPoolSize = config.getPredictorPoolSize(); @@ -72,14 +75,13 @@ public class CommonFaceDetModel implements FaceDetModel{ } } + @Override + public R detect(Image image) { + DetectedObjects detection = detectCore(image); + return R.ok(FaceUtils.convertToDetectionResponse(detection, image)); + } - /** - * 检测人脸 - * @param imagePath 图片路径 - * @return - * @throws Exception - */ @Override public R detect(String imagePath){ if(!FileUtils.isFileExists(imagePath)){ @@ -87,13 +89,17 @@ public class CommonFaceDetModel implements FaceDetModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); - DetectedObjects detection = detect(img); - return R.ok(FaceUtils.convertToDetectionResponse(detection,img)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detection = detectCore(img); + DetectionResponse detectionResponse = FaceUtils.convertToDetectionResponse(detection, img); + if(detectionResponse == null){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + return R.ok(detectionResponse); } catch (IOException e) { throw new FaceException("无效的图片", e); } finally { - if (img != null) { + if (img != null && img.getWrappedImage() instanceof Mat) { ((Mat)img.getWrappedImage()).release(); } } @@ -113,13 +119,17 @@ public class CommonFaceDetModel implements FaceDetModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromInputStream(imageInputStream); - DetectedObjects detection = detect(img); - return R.ok(FaceUtils.convertToDetectionResponse(detection,img)); + img = SmartImageFactory.getInstance().fromInputStream(imageInputStream); + DetectedObjects detection = detectCore(img); + DetectionResponse detectionResponse = FaceUtils.convertToDetectionResponse(detection,img); + if(detectionResponse == null){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + return R.ok(detectionResponse); } catch (IOException e) { throw new FaceException("无效图片输入流", e); } finally { - if (img != null) { + if (img != null && img.getWrappedImage() instanceof Mat) { ((Mat)img.getWrappedImage()).release(); } } @@ -128,18 +138,22 @@ public class CommonFaceDetModel implements FaceDetModel{ @Override public R detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } Image img = null; try { - img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); - DetectedObjects detection = detect(img); - return R.ok(FaceUtils.convertToDetectionResponse(detection,img)); + img = SmartImageFactory.getInstance().fromBufferedImage(image); + DetectedObjects detection = detectCore(img); + DetectionResponse detectionResponse = FaceUtils.convertToDetectionResponse(detection,img); + if(detectionResponse == null){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + return R.ok(detectionResponse); } catch (Exception e) { throw new FaceException(e); } finally { - if (img != null) { + if (img != null && img.getWrappedImage() instanceof Mat) { ((Mat)img.getWrappedImage()).release(); } } @@ -164,14 +178,27 @@ public class CommonFaceDetModel implements FaceDetModel{ } @Override - public R detectAndDraw(String imagePath, String outputPath) { + public R detectAndDraw(Image image) { + DetectedObjects detectedObjects = detectCore(image); + if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + Image drawnImage = ImageUtils.copy(image); + drawnImage.drawBoundingBoxes(detectedObjects); + DetectionResponse detectionResponse = FaceUtils.convertToDetectionResponse(detectedObjects, drawnImage); + detectionResponse.setDrawnImage(drawnImage); + return R.ok(detectionResponse); + } + + @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)); - DetectedObjects detectedObjects = detect(img); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ return R.fail(R.Status.NO_FACE_DETECTED); } @@ -179,11 +206,15 @@ public class CommonFaceDetModel implements FaceDetModel{ Path output = Paths.get(outputPath); log.debug("Saving to {}", output.toAbsolutePath().toString()); img.save(Files.newOutputStream(output), "png"); - return R.ok(); + DetectionResponse detectionResponse = FaceUtils.convertToDetectionResponse(detectedObjects,img); + if(detectionResponse == null){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + return R.ok(detectionResponse); } catch (IOException e) { throw new FaceException(e); } finally { - if (img != null){ + if (img != null && img.getWrappedImage() instanceof Mat) { ((Mat)img.getWrappedImage()).release(); } } @@ -191,11 +222,11 @@ public class CommonFaceDetModel implements FaceDetModel{ @Override public R detectAndDraw(BufferedImage sourceImage) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ return R.fail(R.Status.INVALID_IMAGE); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); - DetectedObjects detectedObjects = detect(img); + Image img = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); + DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ return R.fail(R.Status.NO_FACE_DETECTED); } @@ -210,18 +241,21 @@ public class CommonFaceDetModel implements FaceDetModel{ } catch (IOException e) { throw new FaceException("导出图片失败", e); } finally { - if (img != null){ + if (img != null && img.getWrappedImage() instanceof Mat) { ((Mat)img.getWrappedImage()).release(); } } } + + /** * 人脸检测 * @param image * @return */ - public DetectedObjects detect(Image image){ + @Override + public DetectedObjects detectCore(Image image){ Predictor predictor = null; try { predictor = predictorPool.borrowObject(); @@ -250,8 +284,22 @@ public class CommonFaceDetModel implements FaceDetModel{ return predictorPool; } + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + + @Override public void close() { + if (fromFactory) { + FaceDetModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java index c10aab2..228c4b3 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/FaceDetModel.java @@ -30,6 +30,7 @@ public interface FaceDetModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R detect(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -39,6 +40,7 @@ public interface FaceDetModel extends AutoCloseable{ * @param imageInputStream 图片输入流 * @return */ + @Deprecated default R detect(InputStream imageInputStream){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -48,6 +50,7 @@ public interface FaceDetModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default R detect(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -57,6 +60,7 @@ public interface FaceDetModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default R detect(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -66,16 +70,46 @@ public interface FaceDetModel extends AutoCloseable{ * @param base64Image * @return */ + @Deprecated default R detectBase64(String base64Image) { throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 人脸检测 + * @param image + * @return + */ + default DetectedObjects detectCore(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 人脸检测 + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + + /** + * 检测并绘制人脸 + * @param image + * @return + */ + default R detectAndDraw(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** * 检测并绘制人脸 * @param imagePath 图片输入路径(包含文件名称) * @param outputPath 图片输出路径(包含文件名称) */ - default R detectAndDraw(String imagePath, String outputPath){ + default R detectAndDraw(String imagePath, String outputPath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -84,6 +118,7 @@ public interface FaceDetModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default R detectAndDraw(BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -93,4 +128,7 @@ public interface FaceDetModel extends AutoCloseable{ } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } 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 index 3311d44..e878b12 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facedect/MtcnnFaceDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/MtcnnFaceDetModel.java @@ -5,7 +5,6 @@ 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; @@ -16,6 +15,7 @@ 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.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.Point; import cn.smartjavaai.common.entity.face.FaceInfo; @@ -24,6 +24,7 @@ import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.utils.*; import cn.smartjavaai.face.config.FaceDetConfig; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.FaceDetModelFactory; import cn.smartjavaai.face.model.facedect.criterial.FaceDetCriteriaFactory; import cn.smartjavaai.face.model.facedect.mtcnn.*; import cn.smartjavaai.face.utils.FaceUtils; @@ -53,7 +54,7 @@ import java.util.Objects; * @author dwj */ @Slf4j -public class MtcnnFaceDetModel implements FaceDetModel{ +public class MtcnnFaceDetModel extends CommonFaceDetModel{ public ZooModel pNetModel; @@ -65,6 +66,8 @@ public class MtcnnFaceDetModel implements FaceDetModel{ private GenericObjectPool> rnetPredictorPool; private GenericObjectPool> onetPredictorPool; + private FaceDetConfig config; + /** * 加载模型 @@ -72,6 +75,7 @@ public class MtcnnFaceDetModel implements FaceDetModel{ */ @Override public void loadModel(FaceDetConfig config){ + this.config = config; if(StringUtils.isBlank(config.getModelPath())){ throw new FaceException("modelPath is null"); } @@ -130,146 +134,13 @@ public class MtcnnFaceDetModel implements FaceDetModel{ } - - /** - * 检测人脸 - * @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){ + @Override + public DetectedObjects detectCore(Image image){ Predictor pNetPredictor = null; Predictor rNetPredictor = null; Predictor oNetPredictor = null; @@ -281,35 +152,33 @@ public class MtcnnFaceDetModel implements FaceDetModel{ int w = image.getWidth(); //第一阶段 NDList outputPnet = PNetModel.firstStage(manager, pNetPredictor, image); + if(CollectionUtils.isEmpty(outputPnet)){ - return R.fail(R.Status.NO_FACE_DETECTED); + return DJLCommonUtils.buildEmptyDetectedObjects(); } NDArray boxes = outputPnet.get(0); NDArray image_inds = outputPnet.get(1); NDArray imgs = outputPnet.get(2); if(DJLCommonUtils.isNDArrayEmpty(boxes) || DJLCommonUtils.isNDArrayEmpty(image_inds) || DJLCommonUtils.isNDArrayEmpty(imgs)){ - return R.fail(R.Status.NO_FACE_DETECTED); + return DJLCommonUtils.buildEmptyDetectedObjects(); } NDList pad = MtcnnUtils.pad(boxes, w, h); //第二阶段 - NDList outputRnet = RNetModel.secondStage(manager, rNetPredictor, imgs,boxes,pad, image_inds); + NDList outputRnet = RNetModel.secondStage(manager, rNetPredictor, imgs, boxes, pad, image_inds); if(CollectionUtils.isEmpty(outputRnet)){ - return R.fail(R.Status.NO_FACE_DETECTED); + return DJLCommonUtils.buildEmptyDetectedObjects(); } NDArray image_indsFiltered = outputRnet.get(0); NDArray scoresFiltered = outputRnet.get(1); boxes = outputRnet.get(2); if(DJLCommonUtils.isNDArrayEmpty(boxes) || DJLCommonUtils.isNDArrayEmpty(image_indsFiltered) || DJLCommonUtils.isNDArrayEmpty(scoresFiltered)){ - return R.fail(R.Status.NO_FACE_DETECTED); + return DJLCommonUtils.buildEmptyDetectedObjects(); } //第三阶段 - MtcnnBatchResult oNetResult = ONetModel.thirdStage(manager, oNetPredictor, imgs,boxes, w, h, scoresFiltered, image_indsFiltered); - DetectionResponse detectionResponse = convertToDetectionResponse(oNetResult); - if(Objects.isNull(detectionResponse)){ - return R.fail(R.Status.NO_FACE_DETECTED); - } - return R.ok(detectionResponse); + MtcnnBatchResult oNetResult = ONetModel.thirdStage(manager, oNetPredictor, imgs, boxes, w, h, scoresFiltered, image_indsFiltered); + return FaceUtils.toDetectedObjects(oNetResult, w, h); } catch (Exception e) { + e.printStackTrace(); throw new RuntimeException(e); } finally { if (pNetPredictor != null) { @@ -354,50 +223,50 @@ public class MtcnnFaceDetModel implements FaceDetModel{ - /** - * 转换为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); - - if (DJLCommonUtils.isNDArrayEmpty(boxes) || DJLCommonUtils.isNDArrayEmpty(probs) || DJLCommonUtils.isNDArrayEmpty(points)){ - return null; - } - 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; - } +// /** +// * 转换为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); +// +// if (DJLCommonUtils.isNDArrayEmpty(boxes) || DJLCommonUtils.isNDArrayEmpty(probs) || DJLCommonUtils.isNDArrayEmpty(points)){ +// return null; +// } +// 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; +// } @@ -413,8 +282,22 @@ public class MtcnnFaceDetModel implements FaceDetModel{ return onetPredictorPool; } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + @Override public void close() { + if (fromFactory) { + FaceDetModelFactory.removeFromCache(config.getModelEnum()); + } try { if (pnetPredictorPool != null) { pnetPredictorPool.close(); diff --git a/face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java index bdde257..ef275c4 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facedect/SeetaFace6FaceDetModel.java @@ -1,14 +1,18 @@ package cn.smartjavaai.face.model.facedect; import ai.djl.engine.Engine; +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.utils.Base64ImageUtils; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.face.config.FaceDetConfig; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.FaceDetModelFactory; import cn.smartjavaai.face.seetaface.NativeLoader; import cn.smartjavaai.face.utils.FaceUtils; import com.seeta.pool.*; @@ -81,6 +85,59 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ } + @Override + public R detect(Image image) { + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + FaceDetector predictor = null; + 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(); + for(SeetaRect seetaRect : seetaResult){ + //提取人脸的5点人脸标识 + SeetaPointF[] pointFS = new SeetaPointF[faceLandmarker.number()]; + faceLandmarker.mark(imageData, seetaRect, pointFS); + seetaPointFSList.add(pointFS); + } + return R.ok(FaceUtils.convertToDetectionResponse(seetaResult, seetaPointFSList)); + } catch (Exception e) { + throw new FaceException("目标检测错误", e); + }finally { + if (predictor != null) { + try { + faceDetectorPool.returnObject(predictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + if (faceLandmarker != null) { + try { + faceLandmarkerPool.returnObject(faceLandmarker); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R detectAndDraw(Image image) { + R result = detect(image); + if(result.getCode() != R.Status.SUCCESS.getCode()){ + return R.fail(result.getCode(), result.getMessage()); + } + if(Objects.isNull(result.getData()) || Objects.isNull(result.getData().getDetectionInfoList()) || result.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + Image drawnImage = ImageUtils.drawBoundingBoxes(image, result.getData()); + result.getData().setDrawnImage(drawnImage); + return result; + } + @Override public R detect(String imagePath) { if(!FileUtils.isFileExists(imagePath)){ @@ -112,11 +169,11 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ @Override public R detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); + imageData.data = BufferedImageUtils.getMatrixBGR(image); FaceDetector predictor = null; FaceLandmarker faceLandmarker = null; try { @@ -174,19 +231,14 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ } @Override - public R detectAndDraw(String imagePath, String outputPath) { + public R detectAndDraw(String imagePath, String outputPath) { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } + Image image = null; + Image drawImage = null; try { - //创建保存路径 - Path imageOutputPath = Paths.get(outputPath); - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } + image = SmartImageFactory.getInstance().fromFile(imagePath); R result = detect(image); if(result.getCode() != R.Status.SUCCESS.getCode()){ return R.fail(result.getCode(), result.getMessage()); @@ -195,16 +247,20 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ return R.fail(R.Status.NO_FACE_DETECTED); } //绘制人脸框 - FaceUtils.drawBoundingBoxes(image, result.getData(), imageOutputPath.toAbsolutePath().toString()); - return R.ok(); + drawImage = ImageUtils.drawBoundingBoxes(image, result.getData()); + ImageUtils.save(drawImage, Paths.get(outputPath), "png"); + return result; } catch (IOException e) { - throw new FaceException(e); + throw new FaceException("保存图片失败", e); + } finally { + ImageUtils.releaseOpenCVMat(image); + ImageUtils.releaseOpenCVMat(drawImage); } } @Override public R detectAndDraw(BufferedImage sourceImage) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ return R.fail(R.Status.INVALID_IMAGE); } R result = detect(sourceImage); @@ -215,14 +271,14 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ return R.fail(R.Status.NO_FACE_DETECTED); } //绘制人脸框 - try { - return R.ok(FaceUtils.drawBoundingBoxes(sourceImage, result.getData())); - } catch (IOException e) { - throw new RuntimeException(e); - } + BufferedImage drawnImage = BufferedImageUtils.copyBufferedImage(sourceImage); + BufferedImageUtils.drawBoundingBoxes(drawnImage, result.getData()); + return R.ok(drawnImage); } + + public FaceDetectorPool getFaceDetectorPool() { return faceDetectorPool; } @@ -231,8 +287,23 @@ public class SeetaFace6FaceDetModel implements FaceDetModel{ return faceLandmarkerPool; } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + + @Override public void close() throws Exception { + if (fromFactory) { + FaceDetModelFactory.removeFromCache(config.getModelEnum()); + } try { if (faceDetectorPool != null) { faceDetectorPool.close(); diff --git a/face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java index 38e8941..2e9360a 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/CommonFaceRecModel.java @@ -4,31 +4,27 @@ import ai.djl.MalformedModelException; import ai.djl.engine.Engine; import ai.djl.inference.Predictor; import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.ImageFactory; import ai.djl.ndarray.NDManager; import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ZooModel; -import cn.hutool.core.lang.UUID; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.face.FaceInfo; import cn.smartjavaai.common.entity.face.FaceSearchResult; import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.BufferedImageUtils; 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.config.FaceRecConfig; import cn.smartjavaai.face.constant.FaceDetectConstant; import cn.smartjavaai.face.entity.FaceRegisterInfo; import cn.smartjavaai.face.entity.FaceSearchParams; -import cn.smartjavaai.face.enums.FaceDetModelEnum; import cn.smartjavaai.face.enums.SimilarityType; import cn.smartjavaai.face.exception.FaceException; -import cn.smartjavaai.face.factory.FaceDetModelFactory; -import cn.smartjavaai.face.model.facedect.FaceDetModel; +import cn.smartjavaai.face.factory.FaceRecModelFactory; import cn.smartjavaai.face.model.facerec.criteria.FaceRecCriteriaFactory; -import cn.smartjavaai.face.preprocess.DJLImagePreprocessor; +import cn.smartjavaai.face.preprocess.DJLImageFacePreprocessor; import cn.smartjavaai.face.utils.*; import cn.smartjavaai.face.vector.config.MilvusConfig; import cn.smartjavaai.face.vector.config.SQLiteConfig; @@ -39,8 +35,8 @@ import cn.smartjavaai.face.vector.entity.FaceVector; import cn.smartjavaai.face.vector.exception.VectorDBException; import io.milvus.param.MetricType; 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; @@ -95,7 +91,7 @@ public class CommonFaceRecModel implements FaceRecModel{ throw new FaceException("config为null"); } if(Objects.isNull(config.getDetectModel())){ - config.setDetectModel(getDefaultDetModel()); + throw new FaceException("请指定人脸检测模型"); } this.config = config; Criteria faceFeatureCriteria = FaceRecCriteriaFactory.createCriteria(config); @@ -217,7 +213,7 @@ public class CommonFaceRecModel implements FaceRecModel{ @Override public R featureComparison(BufferedImage sourceImage1, BufferedImage sourceImag2) { - if(!ImageUtils.isImageValid(sourceImage1) || !ImageUtils.isImageValid(sourceImag2)){ + if(!BufferedImageUtils.isImageValid(sourceImage1) || !BufferedImageUtils.isImageValid(sourceImag2)){ throw new FaceException("图像无效"); } R feature1 = extractTopFaceFeature(sourceImage1); @@ -246,19 +242,6 @@ public class CommonFaceRecModel implements FaceRecModel{ } } - /** - * 获取默认人脸检测模型 - * @return - */ - private FaceDetModel getDefaultDetModel() { - FaceDetConfig detectModelConfig = new FaceDetConfig(); - detectModelConfig.setModelEnum(FaceDetModelEnum.ULTRA_LIGHT_FAST_GENERIC_FACE); - detectModelConfig.setConfidenceThreshold(0.98); - log.debug("创建默认人脸检测模型:ULTRA_LIGHT_FAST_GENERIC_FACE"); - FaceDetModel detectModel = FaceDetModelFactory.getInstance().getModel(detectModelConfig); - return detectModel; - } - @Override public R extractFeatures(BufferedImage image) { R detectedResult = config.getDetectModel().detect(image); @@ -268,9 +251,9 @@ public class CommonFaceRecModel implements FaceRecModel{ if(Objects.isNull(detectedResult.getData()) || Objects.isNull(detectedResult.getData().getDetectionInfoList()) || detectedResult.getData().getDetectionInfoList().isEmpty()){ return R.fail(R.Status.NO_FACE_DETECTED); } - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + Image djlImage = SmartImageFactory.getInstance().fromBufferedImage(image); try (NDManager manager = model.getNDManager().newSubManager()) { - DJLImagePreprocessor djlImagePreprocessor = new DJLImagePreprocessor(djlImage, manager); + DJLImageFacePreprocessor djlImagePreprocessor = new DJLImageFacePreprocessor(djlImage, manager); for (DetectionInfo detectionInfo : detectedResult.getData().getDetectionInfoList()){ DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); FaceInfo faceInfo = detectionInfo.getFaceInfo(); @@ -289,12 +272,17 @@ public class CommonFaceRecModel implements FaceRecModel{ subImage = djlImagePreprocessor.process(); } } - features = featureExtraction(subImage); + if (subImage != null && subImage.getWrappedImage() instanceof Mat) { + ((Mat)subImage.getWrappedImage()).release(); + } faceInfo.setFeature(features); } + }finally { + if (djlImage != null && djlImage.getWrappedImage() instanceof Mat) { + ((Mat)djlImage.getWrappedImage()).release(); + } } - ((Mat)djlImage.getWrappedImage()).release(); return detectedResult; } @@ -335,10 +323,10 @@ public class CommonFaceRecModel implements FaceRecModel{ if(Objects.isNull(detectedResult.getData()) || Objects.isNull(detectedResult.getData().getDetectionInfoList()) || detectedResult.getData().getDetectionInfoList().isEmpty()){ return R.fail(R.Status.NO_FACE_DETECTED); } - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + Image djlImage = SmartImageFactory.getInstance().fromBufferedImage(image); float[] features = null; try (NDManager manager = model.getNDManager().newSubManager()) { - DJLImagePreprocessor djlImagePreprocessor = new DJLImagePreprocessor(djlImage, manager); + DJLImageFacePreprocessor djlImagePreprocessor = new DJLImageFacePreprocessor(djlImage, manager); //只取第一个人脸 DetectionInfo detectionInfo = detectedResult.getData().getDetectionInfoList().get(0); DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); @@ -358,8 +346,14 @@ public class CommonFaceRecModel implements FaceRecModel{ } } features = featureExtraction(subImage); + if (subImage != null && subImage.getWrappedImage() instanceof Mat) { + ((Mat)subImage.getWrappedImage()).release(); + } + }finally { + if (djlImage != null && djlImage.getWrappedImage() instanceof Mat) { + ((Mat)djlImage.getWrappedImage()).release(); + } } - ((Mat)djlImage.getWrappedImage()).release(); return Objects.isNull(features) ? R.fail(R.Status.Unknown) : R.ok(features); } @@ -596,7 +590,7 @@ public class CommonFaceRecModel implements FaceRecModel{ throw new FaceException("人脸查询参数为空"); } //设置默认值 - float threshold = Objects.isNull(params.getThreshold()) ? FaceDetectConstant.FACENET_DEFAULT_SIMILARITY_THRESHOLD : params.getThreshold(); + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); @@ -629,13 +623,16 @@ public class CommonFaceRecModel implements FaceRecModel{ return detectionResponse; } //设置默认值 - float threshold = Objects.isNull(params.getThreshold()) ? FaceDetectConstant.FACENET_DEFAULT_SIMILARITY_THRESHOLD : params.getThreshold(); + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); for (DetectionInfo detectionInfo : detectionResponse.getData().getDetectionInfoList()){ if(Objects.nonNull(detectionInfo.getFaceInfo()) && Objects.nonNull(detectionInfo.getFaceInfo().getFeature())){ List searchResults = vectorDBClient.search(detectionInfo.getFaceInfo().getFeature(), searchParams); + if (CollectionUtils.isEmpty(searchResults)){ + return R.fail(1000, "未找到匹配结果"); + } detectionInfo.getFaceInfo().setFaceSearchResults(searchResults); } } @@ -686,9 +683,184 @@ public class CommonFaceRecModel implements FaceRecModel{ vectorDBClient.releaseFaceFeatures(); } + @Override + public R featureComparison(Image image1, Image image2) { + R feature1 = extractTopFaceFeature(image1); + if (!feature1.isSuccess()){ + return R.fail(feature1.getCode(), feature1.getMessage()); + } + R feature2 = extractTopFaceFeature(image2); + if (!feature2.isSuccess()){ + return R.fail(feature2.getCode(), feature2.getMessage()); + } + float ret = calculSimilar(feature1.getData(), feature2.getData()); + return R.ok(ret); + } + + @Override + public R register(FaceRegisterInfo faceRegisterInfo, Image image) { + if(vectorDBClient == null){ + throw new VectorDBException("向量数据库未初始化成功"); + } + //提取最大人脸特征 + R featureResponse = extractTopFaceFeature(image); + if(!featureResponse.isSuccess()){ + return R.fail(featureResponse.getCode(), featureResponse.getMessage()); + } + return register(faceRegisterInfo, featureResponse.getData()); + } + + @Override + public void upsertFace(FaceRegisterInfo faceRegisterInfo, Image image) { + if(vectorDBClient == null){ + throw new VectorDBException("向量数据库未初始化成功"); + } + if(Objects.isNull(faceRegisterInfo)){ + throw new FaceException("注册信息为空"); + } + if(StringUtils.isBlank(faceRegisterInfo.getId())){ + throw new FaceException("注册信息中ID为空"); + } + //提取最大人脸特征 + R featureResponse = extractTopFaceFeature(image); + if(!featureResponse.isSuccess()){ + throw new FaceException(featureResponse.getMessage()); + } + upsertFace(faceRegisterInfo, featureResponse.getData()); + } + + @Override + public R search(Image image, FaceSearchParams params) { + if(vectorDBClient == null){ + throw new VectorDBException("向量数据库未初始化成功"); + } + //提取所有人脸特征 + R detectionResponse = extractFeatures(image); + if(!detectionResponse.isSuccess()){ + return detectionResponse; + } + //设置默认值 + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); + int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); + boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); + FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); + for (DetectionInfo detectionInfo : detectionResponse.getData().getDetectionInfoList()){ + if(Objects.nonNull(detectionInfo.getFaceInfo()) && Objects.nonNull(detectionInfo.getFaceInfo().getFeature())){ + List searchResults = vectorDBClient.search(detectionInfo.getFaceInfo().getFeature(), searchParams); + if (CollectionUtils.isEmpty(searchResults)){ + return R.fail(1000, "未找到匹配结果"); + } + detectionInfo.getFaceInfo().setFaceSearchResults(searchResults); + } + } + return detectionResponse; + } + + @Override + public R> searchByTopFace(Image image, FaceSearchParams params) { + if(vectorDBClient == null){ + return R.fail(1000, "向量数据库未初始化成功"); + } + //提取最大人脸特征 + R featureResponse = extractTopFaceFeature(image); + if(!featureResponse.isSuccess()){ + return R.fail(featureResponse.getCode(), featureResponse.getMessage()); + } + return R.ok(search(featureResponse.getData(), params)); + } + + @Override + public R extractFeatures(Image image) { + R detectedResult = config.getDetectModel().detect(image); + if(!detectedResult.isSuccess()){ + return detectedResult; + } + if(Objects.isNull(detectedResult.getData()) || Objects.isNull(detectedResult.getData().getDetectionInfoList()) || detectedResult.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + try (NDManager manager = model.getNDManager().newSubManager()) { + DJLImageFacePreprocessor djlImagePreprocessor = new DJLImageFacePreprocessor(image, manager); + for (DetectionInfo detectionInfo : detectedResult.getData().getDetectionInfoList()){ + DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); + FaceInfo faceInfo = detectionInfo.getFaceInfo(); + float[] features = null; + Image subImage = null; + //人脸对齐 + if(config.isAlign()){ + //人脸对齐 + double[][] pointsArray = FaceUtils.facePoints(faceInfo.getKeyPoints()); + djlImagePreprocessor.enableCrop(rectangle).enableAffine(pointsArray, 96, 112); + subImage = djlImagePreprocessor.process(); + }else{ + //裁剪 + djlImagePreprocessor.enableCrop(rectangle); + if(config.isCropFace()){ + subImage = djlImagePreprocessor.process(); + } + } + features = featureExtraction(subImage); + if (subImage != null && subImage.getWrappedImage() instanceof Mat) { + ((Mat)subImage.getWrappedImage()).release(); + } + faceInfo.setFeature(features); + } + } + return detectedResult; + } + + @Override + public R extractTopFaceFeature(Image image) { + R detectedResult = config.getDetectModel().detect(image); + if(!detectedResult.isSuccess()){ + return R.fail(detectedResult.getCode(), detectedResult.getMessage()); + } + if(Objects.isNull(detectedResult.getData()) || Objects.isNull(detectedResult.getData().getDetectionInfoList()) || detectedResult.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + float[] features = null; + try (NDManager manager = model.getNDManager().newSubManager()) { + DJLImageFacePreprocessor djlImagePreprocessor = new DJLImageFacePreprocessor(image, manager); + //只取第一个人脸 + DetectionInfo detectionInfo = detectedResult.getData().getDetectionInfoList().get(0); + DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); + FaceInfo faceInfo = detectionInfo.getFaceInfo(); + Image subImage = null; + //人脸对齐 + if(config.isAlign()){ + //人脸对齐 + double[][] pointsArray = FaceUtils.facePoints(faceInfo.getKeyPoints()); + djlImagePreprocessor.enableCrop(rectangle).enableAffine(pointsArray, 96, 112); + subImage = djlImagePreprocessor.process(); + }else{ + //裁剪 + djlImagePreprocessor.enableCrop(rectangle); + if(config.isCropFace()){ + subImage = djlImagePreprocessor.process(); + } + } + features = featureExtraction(subImage); + if (subImage != null && subImage.getWrappedImage() instanceof Mat) { + ((Mat)subImage.getWrappedImage()).release(); + } + } + return Objects.isNull(features) ? R.fail(R.Status.Unknown) : R.ok(features); + } + + + @Override + public Image drawSearchResult(Image image, FaceSearchParams params, String displayField) { + R detectionResponse = search(image, params); + Image drawImage = ImageUtils.copy(image); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(drawImage); + BufferedImageUtils.drawFaceSearchResult(bufferedImage, detectionResponse.getData(), displayField); + return SmartImageFactory.getInstance().fromBufferedImage(bufferedImage); + } @Override public void close() { + if (fromFactory) { + FaceRecModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -717,4 +889,15 @@ public class CommonFaceRecModel implements FaceRecModel{ public GenericObjectPool> getPool() { return predictorPool; } + + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java index 2d3b9ef..37b5e80 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecModel.java @@ -38,22 +38,38 @@ public interface FaceRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + + + /** + * 特征比较 + * @param image1 图1 + * @param image2 图2 + * @return + */ + default R featureComparison(Image image1, Image image2){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 特征比较 * @param imagePath1 图1路径 * @param imagePath2 图2路径 * @return */ + @Deprecated default R featureComparison(String imagePath1, String imagePath2){ throw new UnsupportedOperationException("默认不支持该功能"); } + + /** * 特征比较 * @param sourceImage1 图1BufferedImage * @param sourceImag2 图2BufferedImage * @return */ + @Deprecated default R featureComparison(BufferedImage sourceImage1, BufferedImage sourceImag2){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -65,10 +81,22 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData2 * @return */ + @Deprecated default R featureComparison(byte[] imageData1, byte[] imageData2){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 注册人脸 + * 提取分数最高人脸进行注册 + * @param faceRegisterInfo 注册人脸信息 + * @param image + * @return + */ + default R register(FaceRegisterInfo faceRegisterInfo, Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 注册人脸 * 提取分数最高人脸进行注册 @@ -76,6 +104,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R register(FaceRegisterInfo faceRegisterInfo, String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -87,6 +116,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param inputStream * @return */ + @Deprecated default R register(FaceRegisterInfo faceRegisterInfo, InputStream inputStream){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -98,6 +128,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default R register(FaceRegisterInfo faceRegisterInfo, BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -110,6 +141,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default R register(FaceRegisterInfo faceRegisterInfo, byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -125,6 +157,17 @@ public interface FaceRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 更新或注册人脸 + * 自动提取分数最高人脸进行更新 + * @param faceRegisterInfo 注册人脸信息 + * @param image + * @return + */ + default void upsertFace(FaceRegisterInfo faceRegisterInfo, Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 更新或注册人脸 * 自动提取分数最高人脸进行更新 @@ -132,6 +175,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imagePath * @return */ + @Deprecated default void upsertFace(FaceRegisterInfo faceRegisterInfo, String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -144,6 +188,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default void upsertFace(FaceRegisterInfo faceRegisterInfo, BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -167,10 +212,21 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default void upsertFace(FaceRegisterInfo faceRegisterInfo, byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 查询人脸(查询图片中所有人脸) + * @param image + * @param params 人脸查询参数 + * @return + */ + default R search(Image image, FaceSearchParams params){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 查询人脸(查询图片中所有人脸) @@ -178,6 +234,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param params 人脸查询参数 * @return */ + @Deprecated default R search(String imagePath, FaceSearchParams params){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -189,6 +246,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default R search(BufferedImage sourceImage, FaceSearchParams params){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -199,6 +257,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default R search(byte[] imageData, FaceSearchParams params){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -213,6 +272,18 @@ public interface FaceRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 查询人脸 + * 从图像中提取分数最高的人脸特征,并在人脸库中进行 1:N 查询 + * 适用于单人脸场景 + * @param image + * @param params 人脸查询参数 + * @return + */ + default R> searchByTopFace(Image image, FaceSearchParams params){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 查询人脸 * 从图像中提取分数最高的人脸特征,并在人脸库中进行 1:N 查询 @@ -221,6 +292,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param params 人脸查询参数 * @return */ + @Deprecated default R> searchByTopFace(String imagePath, FaceSearchParams params){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -233,6 +305,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default R> searchByTopFace(BufferedImage sourceImage, FaceSearchParams params){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -244,6 +317,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData * @return */ + @Deprecated default R> searchByTopFace(byte[] imageData, FaceSearchParams params){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -286,12 +360,24 @@ public interface FaceRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + + /** + * 特征提取(所有人脸) + * 适用于多人脸场景 + * @param image + * @return + */ + default R extractFeatures(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 特征提取(所有人脸) * 适用于多人脸场景 * @param imagePath 图片路径 * @return */ + @Deprecated default R extractFeatures(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -302,6 +388,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData 图片字节流 * @return */ + @Deprecated default R extractFeatures(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -312,10 +399,21 @@ public interface FaceRecModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default R extractFeatures(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 特征提取(提取分数最高人脸特征) + * 适用于单人脸场景 + * @param image + * @return + */ + default R extractTopFaceFeature(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 特征提取(提取分数最高人脸特征) @@ -323,6 +421,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default R extractTopFaceFeature(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -333,6 +432,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R extractTopFaceFeature(String imagePath){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -343,6 +443,7 @@ public interface FaceRecModel extends AutoCloseable{ * @param imageData 图片字节流 * @return */ + @Deprecated default R extractTopFaceFeature(byte[] imageData){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -370,9 +471,23 @@ public interface FaceRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 绘制人脸搜索结果 + * @param image + * @param params + * @param displayField + */ + default Image drawSearchResult(Image image, FaceSearchParams params, String displayField){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + default GenericObjectPool> getPool() { throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecPreprocessConfig.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecPreprocessConfig.java new file mode 100644 index 0000000..4c61c39 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecPreprocessConfig.java @@ -0,0 +1,114 @@ +package cn.smartjavaai.face.model.facerec; + +import ai.djl.modality.cv.Image; +import lombok.Data; + +/** + * 人脸识别参数 + * @author dwj + * @date 2025/9/18 + */ +@Data +public class FaceRecPreprocessConfig { + + private int inputWidth; + private int inputHeight; + + /** + * 像素类型 + */ + private Image.Flag imageFlag; + + /** + * 是否使用管道 + */ + private boolean usePipeline; + + /** + * 是否归一化 + */ + private boolean normalize; + + /** + * 归一化 mean + */ + private float[] mean; + + /** + * 归一化 std + */ + private float[] std; + + /** + * 输出索引(默认取第 0 个) + */ + private int outputIndex; + + private FaceRecPreprocessConfig(Builder builder) { + this.inputWidth = builder.inputWidth; + this.inputHeight = builder.inputHeight; + this.imageFlag = builder.imageFlag; + this.usePipeline = builder.usePipeline; + this.normalize = builder.normalize; + this.mean = builder.mean; + this.std = builder.std; + this.outputIndex = builder.outputIndex; + } + + // ========= Builder ========= + public static class Builder { + private int inputWidth = 112; // 默认值 + private int inputHeight = 112; // 默认值 + private Image.Flag imageFlag = Image.Flag.COLOR; // 默认彩色 + private boolean usePipeline = true; + private boolean normalize = true; + + private float[] mean = new float[]{0.5F, 0.5F, 0.5F}; + private float[] std = new float[]{0.5F, 0.5F, 0.5F}; + private int outputIndex; + + + public Builder inputSize(int width, int height) { + this.inputWidth = width; + this.inputHeight = height; + return this; + } + + public Builder imageFlag(Image.Flag flag) { + this.imageFlag = flag; + return this; + } + + public Builder usePipeline(boolean usePipeline) { + this.usePipeline = usePipeline; + return this; + } + + public Builder normalize(boolean normalize) { + this.normalize = normalize; + return this; + } + + public Builder mean(float... mean) { + this.mean = mean; + return this; + } + + public Builder std(float... std) { + this.std = std; + return this; + } + + public Builder outputIndex(int outputIndex) { + this.outputIndex = outputIndex; + return this; + } + + public FaceRecPreprocessConfig build() { + return new FaceRecPreprocessConfig(this); + } + } + + + +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecTranslatorBuilder.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecTranslatorBuilder.java new file mode 100644 index 0000000..a7e590f --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/FaceRecTranslatorBuilder.java @@ -0,0 +1,7 @@ +package cn.smartjavaai.face.model.facerec; + +/** + * @author dwj + */ +public class FaceRecTranslatorBuilder { +} diff --git a/face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java index 0a195fe..f29db53 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/SeetaFace6FaceRecModel.java @@ -1,11 +1,14 @@ package cn.smartjavaai.face.model.facerec; import ai.djl.engine.Engine; +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionInfo; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.entity.face.FaceInfo; import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.face.config.FaceRecConfig; @@ -16,7 +19,10 @@ 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.factory.FaceDetModelFactory; +import cn.smartjavaai.face.factory.FaceRecModelFactory; import cn.smartjavaai.face.utils.FaceUtils; +import cn.smartjavaai.face.utils.Seetaface6Utils; import cn.smartjavaai.face.vector.config.MilvusConfig; import cn.smartjavaai.face.vector.config.SQLiteConfig; import cn.smartjavaai.face.vector.core.VectorDBClient; @@ -29,6 +35,7 @@ import com.seeta.sdk.*; import cn.smartjavaai.face.seetaface.NativeLoader; import io.milvus.param.MetricType; import lombok.extern.slf4j.Slf4j; +import org.apache.commons.collections.CollectionUtils; import org.apache.commons.lang3.StringUtils; import javax.imageio.ImageIO; @@ -213,7 +220,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public R featureComparison(BufferedImage image1, BufferedImage image2) { - if(!ImageUtils.isImageValid(image1) || !ImageUtils.isImageValid(image2)){ + if(!BufferedImageUtils.isImageValid(image1) || !BufferedImageUtils.isImageValid(image2)){ return R.fail(R.Status.INVALID_IMAGE); } R feature1 = extractTopFaceFeature(image1); @@ -263,7 +270,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public R register(FaceRegisterInfo faceRegisterInfo, BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } //提取特征向量 @@ -376,7 +383,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public R search(BufferedImage image, FaceSearchParams params) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } //提取所有人脸特征 @@ -385,13 +392,16 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ return detectionResponse; } //设置默认值 - float threshold = Objects.isNull(params.getThreshold()) ? FaceDetectConstant.SEETAFACE_DEFAULT_SIMILARITY_THRESHOLD : params.getThreshold(); + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); for (DetectionInfo detectionInfo : detectionResponse.getData().getDetectionInfoList()){ if(Objects.nonNull(detectionInfo.getFaceInfo()) && Objects.nonNull(detectionInfo.getFaceInfo().getFeature())){ List searchResults = vectorDBClient.search(detectionInfo.getFaceInfo().getFeature(), searchParams); + if (CollectionUtils.isEmpty(searchResults)){ + return R.fail(1000, "未找到匹配结果"); + } detectionInfo.getFaceInfo().setFaceSearchResults(searchResults); } } @@ -425,7 +435,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ throw new FaceException("人脸查询参数为空"); } //设置默认值 - float threshold = Objects.isNull(params.getThreshold()) ? FaceDetectConstant.SEETAFACE_DEFAULT_SIMILARITY_THRESHOLD : params.getThreshold(); + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); @@ -450,7 +460,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public R> searchByTopFace(BufferedImage sourceImage, FaceSearchParams params) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ return R.fail(R.Status.INVALID_IMAGE); } //提取分数最高人脸特征 @@ -459,7 +469,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ return R.fail(featureResponse.getCode(), featureResponse.getMessage()); } //设置默认值 - float threshold = Objects.isNull(params.getThreshold()) ? FaceDetectConstant.SEETAFACE_DEFAULT_SIMILARITY_THRESHOLD : params.getThreshold(); + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); @@ -586,7 +596,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public R extractFeatures(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } FaceDetector faceDetector = null; @@ -594,7 +604,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ FaceRecognizer faceRecognizer = null; try { SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); + imageData.data = BufferedImageUtils.getMatrixBGR(image); faceRecognizer = faceRecognizerPool.borrowObject(); //默认使用Seetaface6检测模型 if(Objects.isNull(config.getDetectModel())){ @@ -637,7 +647,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ if(Objects.isNull(faceInfo) || Objects.isNull(faceInfo.getKeyPoints())){ return R.fail(R.Status.Unknown.getCode(), "未检测到人脸关键点"); } - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(faceInfo.getKeyPoints()); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(faceInfo.getKeyPoints()); //CropFaceV2 + ExtractCroppedFace 已包含裁剪+人脸对齐 boolean isSuccess = faceRecognizer.Extract(imageData, pointFS, features); if(!isSuccess){ @@ -679,12 +689,12 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ @Override public R extractTopFaceFeature(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } float[] features = null; SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); + imageData.data = BufferedImageUtils.getMatrixBGR(image); FaceDetector faceDetector = null; FaceLandmarker faceLandmarker = null; FaceRecognizer faceRecognizer = null; @@ -712,7 +722,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ return R.fail(R.Status.NO_FACE_DETECTED); } DetectionInfo detectionInfo = detectResponse.getData().getDetectionInfoList().get(0); - pointFS = FaceUtils.convertToSeetaPointF(detectionInfo.getFaceInfo().getKeyPoints()); + pointFS = Seetaface6Utils.convertToSeetaPointF(detectionInfo.getFaceInfo().getKeyPoints()); } //提取特征 features = new float[faceRecognizer.GetExtractFeatureSize()]; @@ -761,7 +771,7 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ try { faceRecognizer = faceRecognizerPool.borrowObject(); SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); + imageData.data = BufferedImageUtils.getMatrixBGR(image); //提取特征 float[] features = new float[faceRecognizer.GetExtractFeatureSize()]; faceRecognizer.ExtractCroppedFace(imageData, features); @@ -904,8 +914,264 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ return R.ok(vectorDBClient.listFaces(pageNum, pageSize)); } + + @Override + public R featureComparison(Image image1, Image image2) { + R feature1 = extractTopFaceFeature(image1); + if(!feature1.isSuccess()){ + return R.fail(feature1.getCode(), feature1.getMessage()); + } + R feature2 = extractTopFaceFeature(image2); + if(!feature2.isSuccess()){ + return R.fail(feature2.getCode(), feature2.getMessage()); + } + return R.ok(calculSimilar(feature1.getData(), feature2.getData())); + } + + @Override + public R register(FaceRegisterInfo faceRegisterInfo, Image image) { + //提取特征向量 + R featureResponse = extractTopFaceFeature(image); + if(!featureResponse.isSuccess()){ + return R.fail(featureResponse.getCode(), featureResponse.getMessage()); + } + return register(faceRegisterInfo, featureResponse.getData()); + } + + @Override + public void upsertFace(FaceRegisterInfo faceRegisterInfo, Image image) { + if(vectorDBClient == null){ + throw new VectorDBException("向量数据库未初始化成功"); + } + if(Objects.isNull(faceRegisterInfo)){ + throw new FaceException("注册信息为空"); + } + if(StringUtils.isBlank(faceRegisterInfo.getId())){ + throw new FaceException("注册信息中ID为空"); + } + //提取最大人脸特征 + R featureResponse = extractTopFaceFeature(image); + if(!featureResponse.isSuccess()){ + throw new FaceException(featureResponse.getMessage()); + } + upsertFace(faceRegisterInfo, featureResponse.getData()); + } + + @Override + public R search(Image image, FaceSearchParams params) { + //提取所有人脸特征 + R detectionResponse = extractFeatures(image); + if(!detectionResponse.isSuccess()){ + return detectionResponse; + } + //设置默认值 + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); + int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); + boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); + FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); + for (DetectionInfo detectionInfo : detectionResponse.getData().getDetectionInfoList()){ + if(Objects.nonNull(detectionInfo.getFaceInfo()) && Objects.nonNull(detectionInfo.getFaceInfo().getFeature())){ + List searchResults = vectorDBClient.search(detectionInfo.getFaceInfo().getFeature(), searchParams); + if (CollectionUtils.isEmpty(searchResults)){ + return R.fail(1000, "未找到匹配结果"); + } + detectionInfo.getFaceInfo().setFaceSearchResults(searchResults); + } + } + return detectionResponse; + } + + @Override + public R> searchByTopFace(Image image, FaceSearchParams params) { + //提取分数最高人脸特征 + R featureResponse = extractTopFaceFeature(image); + if(!featureResponse.isSuccess()){ + return R.fail(featureResponse.getCode(), featureResponse.getMessage()); + } + //设置默认值 + float threshold = Objects.isNull(params.getThreshold()) ? config.getModelEnum().getThreshold() : params.getThreshold(); + int topK = Objects.isNull(params.getTopK()) ? 1 : params.getTopK(); + boolean normalize = Objects.isNull(params.getNormalizeSimilarity()) ? NORMALIZE_SIMILARITY : params.getNormalizeSimilarity(); + FaceSearchParams searchParams = new FaceSearchParams(topK, threshold, normalize); + List searchResults = vectorDBClient.search(featureResponse.getData(), searchParams); + return R.ok(searchResults); + } + + @Override + public R extractFeatures(Image image) { + FaceDetector faceDetector = null; + FaceLandmarker faceLandmarker = null; + FaceRecognizer faceRecognizer = null; + try { + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + faceRecognizer = faceRecognizerPool.borrowObject(); + //默认使用Seetaface6检测模型 + if(Objects.isNull(config.getDetectModel())){ + List featureList = new ArrayList(); + List seetaPointFSList = new ArrayList(); + faceDetector = faceDetectorPool.borrowObject(); + faceLandmarker = faceLandmarkerPool.borrowObject(); + //检测人脸 + SeetaRect[] seetaResult = faceDetector.Detect(imageData); + if(Objects.isNull(seetaResult) || seetaResult.length == 0){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + for(SeetaRect seetaRect : seetaResult){ + //提取人脸的5点人脸标识 + SeetaPointF[] pointFS = new SeetaPointF[faceLandmarker.number()]; + faceLandmarker.mark(imageData, seetaRect, pointFS); + //提取特征 + float[] features = new float[faceRecognizer.GetExtractFeatureSize()]; + //CropFaceV2 + ExtractCroppedFace 已包含裁剪+人脸对齐 + boolean isSuccess = faceRecognizer.Extract(imageData, pointFS, features); + if(!isSuccess){ + return R.fail(R.Status.Unknown.getCode(), "人脸特征提取失败"); + } + featureList.add(features); + seetaPointFSList.add(pointFS); + } + return R.ok(FaceUtils.featuresConvertToResponse(seetaResult, seetaPointFSList, featureList)); + }else{ + R detectResponse = config.getDetectModel().detect(image); + if(!detectResponse.isSuccess()){ + return detectResponse; + } + if(Objects.isNull(detectResponse.getData()) || Objects.isNull(detectResponse.getData().getDetectionInfoList()) || detectResponse.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + for(DetectionInfo detectionInfo : detectResponse.getData().getDetectionInfoList()){ + //提取特征 + float[] features = new float[faceRecognizer.GetExtractFeatureSize()]; + FaceInfo faceInfo = detectionInfo.getFaceInfo(); + if(Objects.isNull(faceInfo) || Objects.isNull(faceInfo.getKeyPoints())){ + return R.fail(R.Status.Unknown.getCode(), "未检测到人脸关键点"); + } + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(faceInfo.getKeyPoints()); + //CropFaceV2 + ExtractCroppedFace 已包含裁剪+人脸对齐 + boolean isSuccess = faceRecognizer.Extract(imageData, pointFS, features); + if(!isSuccess){ + return R.fail(R.Status.Unknown.getCode(), "人脸特征提取失败"); + } + faceInfo.setFeature(features); + } + return detectResponse; + } + } catch (FaceException e) { + throw e; + } catch (Exception e) { + throw new FaceException("人脸特征提取异常", e); + }finally { + if (faceDetector != null) { + try { + faceDetectorPool.returnObject(faceDetector); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + if (faceLandmarker != null) { + try { + faceLandmarkerPool.returnObject(faceLandmarker); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + if (faceRecognizer != null) { + try { + faceRecognizerPool.returnObject(faceRecognizer); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R extractTopFaceFeature(Image image) { + float[] features = null; + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + FaceDetector faceDetector = null; + FaceLandmarker faceLandmarker = null; + FaceRecognizer faceRecognizer = null; + try { + faceRecognizer = faceRecognizerPool.borrowObject(); + //提取人脸的5点人脸标识 + SeetaPointF[] pointFS = null; + //默认使用Seetaface6检测模型 + if(Objects.isNull(config.getDetectModel())){ + faceDetector = faceDetectorPool.borrowObject(); + faceLandmarker = faceLandmarkerPool.borrowObject(); + //检测人脸 + SeetaRect[] seetaResult = faceDetector.Detect(imageData); + if(Objects.isNull(seetaResult) || seetaResult.length == 0){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + pointFS = new SeetaPointF[faceLandmarker.number()]; + faceLandmarker.mark(imageData, seetaResult[0], pointFS); + }else{ + R detectResponse = config.getDetectModel().detect(image); + if(!detectResponse.isSuccess()){ + return R.fail(detectResponse.getCode(), detectResponse.getMessage()); + } + if(Objects.isNull(detectResponse.getData()) || Objects.isNull(detectResponse.getData().getDetectionInfoList()) || detectResponse.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + DetectionInfo detectionInfo = detectResponse.getData().getDetectionInfoList().get(0); + pointFS = Seetaface6Utils.convertToSeetaPointF(detectionInfo.getFaceInfo().getKeyPoints()); + } + //提取特征 + features = new float[faceRecognizer.GetExtractFeatureSize()]; + //CropFaceV2 + ExtractCroppedFace 已包含裁剪+人脸对齐 + boolean isSuccess = faceRecognizer.Extract(imageData, pointFS, features); + if(!isSuccess){ + return R.fail(R.Status.Unknown.getCode(), "人脸特征提取失败"); + } + return R.ok(features); + } catch (FaceException e) { + throw e; + } catch (Exception e) { + throw new FaceException("目标检测错误", e); + }finally { + if (faceDetector != null) { + try { + faceDetectorPool.returnObject(faceDetector); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + if (faceLandmarker != null) { + try { + faceLandmarkerPool.returnObject(faceLandmarker); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + if (faceRecognizer != null) { + try { + faceRecognizerPool.returnObject(faceRecognizer); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + + @Override + public Image drawSearchResult(Image image, FaceSearchParams params, String displayField) { + R detectionResponse = search(image, params); + Image drawImage = ImageUtils.copy(image); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(drawImage); + BufferedImageUtils.drawFaceSearchResult(bufferedImage, detectionResponse.getData(), displayField); + return SmartImageFactory.getInstance().fromBufferedImage(bufferedImage); + } + @Override public void close() throws Exception { + if (fromFactory) { + FaceRecModelFactory.removeFromCache(config.getModelEnum()); + } if(Objects.nonNull(faceDetectorPool)){ faceDetectorPool.close(); } @@ -923,6 +1189,9 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ } } + + + @Override public boolean isLoadFaceCompleted() { return isLoadCompleted; @@ -944,4 +1213,15 @@ public class SeetaFace6FaceRecModel implements FaceRecModel{ public FaceDatabasePool getFaceDatabasePool() { return faceDatabasePool; } + + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/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 index de6b708..3461417 100644 --- a/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 @@ -4,18 +4,20 @@ import ai.djl.Device; import ai.djl.modality.cv.Image; 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.FaceRecConfig; import cn.smartjavaai.face.constant.FaceNetConstant; import cn.smartjavaai.face.enums.FaceRecModelEnum; -import cn.smartjavaai.face.model.facerec.translator.FaceFeatureTranslator; -import cn.smartjavaai.face.model.facerec.translator.FaceNetRecTranslator; +import cn.smartjavaai.face.model.facerec.FaceRecPreprocessConfig; +import cn.smartjavaai.face.model.facerec.translator.*; import org.apache.commons.lang3.StringUtils; import java.nio.file.Paths; import java.util.Arrays; import java.util.List; import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; import java.util.stream.Collectors; /** @@ -29,73 +31,40 @@ public class FaceRecCriteriaFactory { if(!Objects.isNull(config.getDevice())){ device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId()); } - Criteria criteria = null; - if(config.getModelEnum() == FaceRecModelEnum.FACENET_MODEL){ - criteria = - Criteria.builder() - .setTypes(Image.class, float[].class) - .optModelName("face_feature") // specify model file prefix - .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : - FaceNetConstant.MODEL_URL) - .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) - .optTranslator(new FaceNetRecTranslator()) - .optDevice(device) - .optEngine("PyTorch") // Use PyTorch engine - .optProgress(new ProgressBar()) - .build(); - }else if (config.getModelEnum() == FaceRecModelEnum.INSIGHT_FACE_MOBILE_FACENET_MODEL){ - if(StringUtils.isBlank(config.getModelPath())){ - throw new RuntimeException("请指定模型路径"); - } - List mean = Arrays.asList(0.5f,0.5f,0.5f,0.5f,0.5f,0.5f); - String normalize = mean.stream().map(Object::toString).collect(Collectors.joining(",")); - criteria = Criteria.builder() - .setTypes(Image.class, float[].class) - .optModelPath(Paths.get(config.getModelPath())) -// .optTranslatorFactory(new ImageFeatureExtractorFactory()) -// .optArgument("normalize", normalize) -// .optArgument("resize", "112,112") - .optTranslator(new FaceFeatureTranslator()) - .optEngine("PyTorch") // Use PyTorch engine - .optProgress(new ProgressBar()) - .optDevice(device) - .build(); - }else if (config.getModelEnum() == FaceRecModelEnum.INSIGHT_FACE_IRSE50_MODEL){ - if(StringUtils.isBlank(config.getModelPath())){ - throw new RuntimeException("请指定模型路径"); - } - List mean = Arrays.asList(0.5f,0.5f,0.5f,0.5f,0.5f,0.5f); - String normalize = mean.stream().map(Object::toString).collect(Collectors.joining(",")); - criteria = Criteria.builder() - .setTypes(Image.class, float[].class) - .optModelPath(Paths.get(config.getModelPath())) -// .optTranslatorFactory(new ImageFeatureExtractorFactory()) -// .optArgument("normalize", normalize) -// .optArgument("resize", "112,112") - .optTranslator(new FaceFeatureTranslator()) - .optEngine("PyTorch") // Use PyTorch engine - .optDevice(device) - .optProgress(new ProgressBar()) - .build(); - }else if (config.getModelEnum() == FaceRecModelEnum.ELASTIC_FACE_MODEL){ - if(StringUtils.isBlank(config.getModelPath())){ - throw new RuntimeException("请指定模型路径"); - } - List mean = Arrays.asList(0.5f,0.5f,0.5f,0.5f,0.5f,0.5f); - String normalize = mean.stream().map(Object::toString).collect(Collectors.joining(",")); - criteria = Criteria.builder() - .setTypes(Image.class, float[].class) - .optModelPath(Paths.get(config.getModelPath())) -// .optTranslatorFactory(new ImageFeatureExtractorFactory()) -// .optArgument("normalize", normalize) -// .optArgument("resize", "112,112") - .optTranslator(new FaceFeatureTranslator()) - .optEngine("PyTorch") // Use PyTorch engine - .optDevice(device) - .optProgress(new ProgressBar()) - .build(); - } + Translator translator = getFaceRecTranslator(config); + Criteria criteria = + Criteria.builder() + .setTypes(Image.class, float[].class) + .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : + FaceNetConstant.MODEL_URL) + .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) + .optTranslator(translator) + .optDevice(device) + .optEngine(config.getModelEnum().getEngine()) + .optProgress(new ProgressBar()) + .build(); return criteria; } + /** + * 获取人脸识别模型Translator + * @param config + * @return + */ + public static Translator getFaceRecTranslator(FaceRecConfig config) { + FaceRecPreprocessConfig preprocessConfig = new FaceRecPreprocessConfig.Builder() + .inputSize(config.getModelEnum().getInputWidth(), config.getModelEnum().getInputHeight()) + .build(); + switch (config.getModelEnum()) { + case VGG_FACE: + preprocessConfig = new FaceRecPreprocessConfig.Builder() + .inputSize(config.getModelEnum().getInputWidth(), config.getModelEnum().getInputHeight()) + .usePipeline(false) + .normalize(false) + .build(); + break; + } + return new CommonFaceRecTranslator(preprocessConfig); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/CommonFaceRecTranslator.java b/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/CommonFaceRecTranslator.java new file mode 100644 index 0000000..f7631c2 --- /dev/null +++ b/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/CommonFaceRecTranslator.java @@ -0,0 +1,77 @@ +package cn.smartjavaai.face.model.facerec.translator; + +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.transform.Normalize; +import ai.djl.modality.cv.transform.Resize; +import ai.djl.modality.cv.transform.ToTensor; +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.Batchifier; +import ai.djl.translate.Pipeline; +import ai.djl.translate.Translator; +import ai.djl.translate.TranslatorContext; +import cn.smartjavaai.face.model.facerec.FaceRecPreprocessConfig; + +/** + * 通用人脸识别模型转换器 + * @author dwj + */ +public class CommonFaceRecTranslator implements Translator { + + private FaceRecPreprocessConfig preprocessConfig; + + public CommonFaceRecTranslator(FaceRecPreprocessConfig preprocessConfig) { + this.preprocessConfig = preprocessConfig; + } + + /** + * {@inheritDoc} + */ + @Override + public NDList processInput(TranslatorContext ctx, Image input) { + NDArray array = input.toNDArray(ctx.getNDManager(), preprocessConfig.getImageFlag()); + NDList ndList = null; + if(preprocessConfig.isUsePipeline()){ + Pipeline pipeline = new Pipeline(); + if(input.getWidth() != preprocessConfig.getInputWidth() || input.getHeight() != preprocessConfig.getInputHeight()){ + pipeline.add(new Resize(preprocessConfig.getInputWidth(), preprocessConfig.getInputHeight())); + } + pipeline.add(new ToTensor()); + if(preprocessConfig.isNormalize()){ + pipeline.add(new Normalize( + preprocessConfig.getMean(), + preprocessConfig.getStd())); + } + ndList = pipeline.transform(new NDList(array)); + }else{ + if(input.getWidth() != preprocessConfig.getInputWidth() || input.getHeight() != preprocessConfig.getInputHeight()){ + array = NDImageUtils.resize(array, preprocessConfig.getInputWidth(), preprocessConfig.getInputHeight()); + } + array = array.toType(DataType.FLOAT32, false); + if (preprocessConfig.isNormalize()){ + array = array.sub(preprocessConfig.getMean()[0]).div(preprocessConfig.getStd()[0]); + } + array = array.transpose(2, 0, 1); + return new NDList(array); + } + return ndList; + } + + /** + * {@inheritDoc} + */ + @Override + public float[] processOutput(TranslatorContext ctx, NDList list) { + NDArray embedding = list.get(preprocessConfig.getOutputIndex()); + embedding = embedding.div(embedding.norm()); // L2归一化 + return embedding.toFloatArray(); + } + + @Override + public Batchifier getBatchifier() { + return Batchifier.STACK; + } + +} diff --git a/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 deleted file mode 100644 index 6b623af..0000000 --- a/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceFeatureTranslator.java +++ /dev/null @@ -1,59 +0,0 @@ -package cn.smartjavaai.face.model.facerec.translator; - -import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.transform.Normalize; -import ai.djl.modality.cv.transform.Resize; -import ai.djl.modality.cv.transform.ToTensor; -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.translate.Batchifier; -import ai.djl.translate.Pipeline; -import ai.djl.translate.Translator; -import ai.djl.translate.TranslatorContext; - -/** - * facenet人脸特征提取Translator - * @author dwj - * @date 2025/3/31 - */ -public final class FaceFeatureTranslator implements Translator { - - - - public FaceFeatureTranslator() { - } - - /** - * {@inheritDoc} - */ - @Override - public NDList processInput(TranslatorContext ctx, Image input) { - NDArray array = input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR); - Pipeline pipeline = new Pipeline(); - if(input.getWidth() != 112 || input.getHeight() != 112){ - pipeline.add(new Resize(112,112)); - } - pipeline - .add(new ToTensor()) - .add(new Normalize( - new float[]{0.5F, 0.5F, 0.5F}, - new float[]{0.5F, 0.5F, 0.5F})); - - return pipeline.transform(new NDList(array)); - } - - /** - * {@inheritDoc} - */ - @Override - public float[] processOutput(TranslatorContext ctx, NDList list) { - NDArray embedding = list.singletonOrThrow(); - embedding = embedding.div(embedding.norm()); // L2归一化 - return embedding.toFloatArray(); - } - - @Override - public Batchifier getBatchifier() { - return Batchifier.STACK; - } -} diff --git a/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 deleted file mode 100644 index 5c21867..0000000 --- a/face/src/main/java/cn/smartjavaai/face/model/facerec/translator/FaceNetRecTranslator.java +++ /dev/null @@ -1,57 +0,0 @@ -package cn.smartjavaai.face.model.facerec.translator; - -import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.transform.Normalize; -import ai.djl.modality.cv.transform.Resize; -import ai.djl.modality.cv.transform.ToTensor; -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.translate.Batchifier; -import ai.djl.translate.Pipeline; -import ai.djl.translate.Translator; -import ai.djl.translate.TranslatorContext; - -/** - * facenet人脸特征提取Translator - * @author dwj - * @date 2025/3/31 - */ -public final class FaceNetRecTranslator implements Translator { - - - - public FaceNetRecTranslator() { - } - - /** - * {@inheritDoc} - */ - @Override - public NDList processInput(TranslatorContext ctx, Image input) { - NDArray array = input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR); - Pipeline pipeline = new Pipeline(); - pipeline - //.add(new Resize(112)) - .add(new ToTensor()) - .add(new Normalize( - new float[]{0.5F, 0.5F, 0.5F}, - new float[]{0.5F, 0.5F, 0.5F})); - - return pipeline.transform(new NDList(array)); - } - - /** - * {@inheritDoc} - */ - @Override - public float[] processOutput(TranslatorContext ctx, NDList list) { - NDArray embedding = list.singletonOrThrow(); - embedding = embedding.div(embedding.norm()); // L2归一化 - return embedding.toFloatArray(); - } - - @Override - public Batchifier getBatchifier() { - return Batchifier.STACK; - } -} diff --git a/face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java index f1fbd0a..9a73be3 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/liveness/CommonLivenessModel.java @@ -11,31 +11,35 @@ import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ZooModel; import ai.djl.training.util.ProgressBar; import ai.djl.util.JsonUtils; -import cn.smartjavaai.common.entity.DetectionInfo; -import cn.smartjavaai.common.entity.DetectionRectangle; -import cn.smartjavaai.common.entity.DetectionResponse; -import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.cv.SmartImageFactory; +import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.face.FaceInfo; import cn.smartjavaai.common.entity.face.LivenessResult; import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.face.LivenessStatus; import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.preprocess.BufferedImagePreprocessor; +import cn.smartjavaai.common.preprocess.DJLImagePreprocessor; import cn.smartjavaai.common.utils.*; import cn.smartjavaai.face.config.LivenessConfig; import cn.smartjavaai.face.constant.MiniVisionConstant; +import cn.smartjavaai.face.entity.FaceQualityResult; import cn.smartjavaai.face.enums.LivenessModelEnum; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.FaceDetModelFactory; +import cn.smartjavaai.face.factory.LivenessModelFactory; import cn.smartjavaai.face.model.liveness.criterial.LivenessCriteriaFactory; import cn.smartjavaai.face.model.liveness.translator.MiniVisionTranslator; import com.seeta.sdk.FaceAntiSpoofing; import lombok.extern.slf4j.Slf4j; +import nu.pattern.OpenCV; import org.apache.commons.lang3.StringUtils; import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.impl.GenericObjectPool; import org.bytedeco.javacv.FFmpegFrameGrabber; import org.bytedeco.javacv.Frame; import org.bytedeco.javacv.Java2DFrameUtils; +import org.bytedeco.javacv.OpenCVFrameConverter; import org.opencv.core.Mat; import javax.imageio.ImageIO; @@ -56,11 +60,19 @@ import java.util.*; @Slf4j public class CommonLivenessModel implements LivenessDetModel{ + static { + //视频功能需要 + OpenCV.loadLocally(); + } + protected GenericObjectPool> predictorPool; protected LivenessConfig config; protected ZooModel model; + + private OpenCVFrameConverter.ToOrgOpenCvCoreMat converterToMat = null; + @Override public void loadModel(LivenessConfig config) { if(Objects.isNull(config)){ @@ -69,6 +81,9 @@ public class CommonLivenessModel implements LivenessDetModel{ if(StringUtils.isBlank(config.getModelPath())){ throw new FaceException("modelPath不能为空"); } + if(Objects.isNull(config.getDetectModel())){ + throw new FaceException("未指定检测模型"); + } this.config = config; //设置真人阈值 Float realityThreshold = Objects.isNull(config.getRealityThreshold()) ? MiniVisionConstant.REALITY_THRESHOLD : config.getRealityThreshold(); @@ -92,254 +107,6 @@ public class CommonLivenessModel implements LivenessDetModel{ } - @Override - public R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle) { - if(!ImageUtils.isImageValid(image)){ - throw new FaceException("图像无效"); - } - Predictor predictor = null; - Image djlImage = null; - try { - predictor = predictorPool.borrowObject(); - //预处理图片 - BufferedImage processedImage = image; - if(config.getModelEnum() == LivenessModelEnum.IIC_FL_MODEL){ - processedImage = new BufferedImagePreprocessor(image, faceDetectionRectangle) - .setExtendRatio(96f / 112f) - .enableSquarePadding(true) - .enableScaling(true) - .setTargetSize(128) - .enableCenterCrop(true) - .setCenterCropSize(112) - .process(); - } - djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(processedImage)); - Float result = predictor.predict(djlImage); - if(result >= config.getRealityThreshold()){ - return R.ok(new LivenessResult(LivenessStatus.LIVE, result)); - }else{ - float nonLiveScore = BigDecimal.ONE.subtract(new BigDecimal(result)).floatValue(); - return R.ok(new LivenessResult(LivenessStatus.NON_LIVE, nonLiveScore)); - } - } catch (Exception e) { - throw new FaceException("活体检测错误", e); - }finally { - if (predictor != null) { - try { - predictorPool.returnObject(predictor); //归还 - } catch (Exception e) { - log.warn("归还Predictor失败", e); - try { - predictor.close(); // 归还失败才销毁 - } catch (Exception ex) { - log.error("关闭Predictor失败", ex); - } - } - } - if (djlImage != null){ - ((Mat)djlImage.getWrappedImage()).release(); - } - - } - } - - @Override - public R detect(String imagePath, DetectionRectangle faceDetectionRectangle) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detect(image, faceDetectionRectangle); - } - - @Override - public R detect(byte[] imageData, DetectionRectangle faceDetectionRectangle) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionRectangle); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @Override - public R detectBase64(String base64Image, DetectionRectangle faceDetectionRectangle) { - if(StringUtils.isBlank(base64Image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detect(imageData, faceDetectionRectangle); - } - - @Override - public R> detect(String imagePath, DetectionResponse faceDetectionResponse) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detect(image, faceDetectionResponse); - } - - @Override - public R> detect(byte[] imageData, DetectionResponse faceDetectionResponse) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionResponse); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @Override - public R> detect(BufferedImage image, DetectionResponse faceDetectionResponse) { - if(!ImageUtils.isImageValid(image)){ - R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ - R.fail(R.Status.NO_FACE_DETECTED); - } - List livenessStatusList = new ArrayList(); - for(DetectionInfo detectionInfo : faceDetectionResponse.getDetectionInfoList()){ - R result = detect(image, detectionInfo.getDetectionRectangle()); - if(!result.isSuccess()){ - return R.fail(result.getCode(), result.getMessage()); - } - livenessStatusList.add(result.getData()); - } - return R.ok(livenessStatusList); - } - - @Override - public R> detectBase64(String base64Image, DetectionResponse faceDetectionResponse) { - if(StringUtils.isBlank(base64Image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detect(imageData, faceDetectionResponse); - } - - - @Override - public R detectTopFace(BufferedImage image) { - if(Objects.isNull(config.getDetectModel())){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "未指定检测模型"); - } - R faceDetectionResponse = config.getDetectModel().detect(image); - if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ - return R.fail(R.Status.NO_FACE_DETECTED); - } - return detect(image, faceDetectionResponse.getData().getDetectionInfoList().get(0).getDetectionRectangle()); - } - - @Override - public R detectTopFace(String imagePath) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detectTopFace(image); - } - - @Override - public R detectTopFace(byte[] imageData) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detectTopFace(ImageIO.read(new ByteArrayInputStream(imageData))); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @Override - public R detectTopFaceBase64(String base64Image) { - if(StringUtils.isBlank(base64Image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - byte[] imageData = Base64ImageUtils.base64ToImage(base64Image); - return detectTopFace(imageData); - } - - @Override - public R detect(String imagePath) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detect(image); - } - - @Override - public R detect(BufferedImage image) { - if(Objects.isNull(config.getDetectModel())){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "未指定检测模型"); - } - R faceDetectionResponse = config.getDetectModel().detect(image); - if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ - return R.fail(R.Status.NO_FACE_DETECTED); - } - for(DetectionInfo detectionInfo : faceDetectionResponse.getData().getDetectionInfoList()){ - R result = detect(image, detectionInfo.getDetectionRectangle()); - if(!result.isSuccess()){ - return R.fail(result.getCode(), result.getMessage()); - } - if(Objects.isNull(detectionInfo.getFaceInfo())){ - detectionInfo.setFaceInfo(new FaceInfo()); - } - detectionInfo.getFaceInfo().setLivenessStatus(result.getData()); - } - return faceDetectionResponse; - } - - @Override - public R detect(byte[] imageData) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData))); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @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 detectVideo(InputStream videoInputStream) { @@ -376,20 +143,18 @@ public class CommonLivenessModel implements LivenessDetModel{ // 获取当前帧 Frame frame = grabber.grabImage(); if (frame != null) { - BufferedImage bufferedImage = Java2DFrameUtils.toBufferedImage(frame); - R livenessStatus = detectTopFace(bufferedImage); - if(!livenessStatus.isSuccess()){ - log.debug("第" + frameIndex + "帧处理失败:" + livenessStatus.getMessage()); + if(converterToMat == null){ + converterToMat = new OpenCVFrameConverter.ToOrgOpenCvCoreMat(); + } + Mat mat = converterToMat.convert(frame); + R livenessScore = detectTopFace(SmartImageFactory.getInstance().fromMat(mat)); + mat.release(); + if(!livenessScore.isSuccess()){ + log.debug("第" + frameIndex + "帧处理失败:" + livenessScore.getMessage()); continue; }else{ - log.debug("第" + frameIndex + "帧活体检测结果:" + JsonUtils.toJson(livenessStatus)); - float liveScore = 0; - if(livenessStatus.getData().getStatus() == LivenessStatus.LIVE){ - liveScore = livenessStatus.getData().getScore(); - }else{ - liveScore = BigDecimal.ONE.subtract(BigDecimal.valueOf(livenessStatus.getData().getScore())).floatValue(); - } - scoreWindow.add(liveScore); + log.debug("第" + frameIndex + "帧活体检测结果:" + livenessScore); + scoreWindow.add(livenessScore.getData().getScore()); } // 如果累计检测帧数 >= 配置值,开始判断 if (scoreWindow.size() >= config.getFrameCount()) { @@ -398,14 +163,9 @@ public class CommonLivenessModel implements LivenessDetModel{ .average() .orElse(0.0); log.debug("滑动窗口平均得分: {}", avgScore); - if (avgScore >= config.getRealityThreshold()) { - grabber.stop(); - return R.ok(new LivenessResult(LivenessStatus.LIVE, avgScore)); - } else { - grabber.stop(); - float nonLiveScore = BigDecimal.ONE.subtract(BigDecimal.valueOf(avgScore)).floatValue(); - return R.ok(new LivenessResult(LivenessStatus.NON_LIVE, nonLiveScore)); - } + grabber.stop(); + LivenessStatus livenessStatus = avgScore > config.getRealityThreshold() ? LivenessStatus.LIVE : LivenessStatus.NON_LIVE; + return R.ok(new LivenessResult(livenessStatus, avgScore)); } } } @@ -420,6 +180,97 @@ public class CommonLivenessModel implements LivenessDetModel{ } + @Override + public R detect(Image image) { + if(Objects.isNull(config.getDetectModel())){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "未指定检测模型"); + } + R faceDetectionResponse = config.getDetectModel().detect(image); + if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + for(DetectionInfo detectionInfo : faceDetectionResponse.getData().getDetectionInfoList()){ + R result = detect(image, detectionInfo.getDetectionRectangle()); + if(!result.isSuccess()){ + return R.fail(result.getCode(), result.getMessage()); + } + if(Objects.isNull(detectionInfo.getFaceInfo())){ + detectionInfo.setFaceInfo(new FaceInfo()); + } + detectionInfo.getFaceInfo().setLivenessStatus(result.getData()); + } + return faceDetectionResponse; + } + + @Override + public R> detect(Image image, DetectionResponse faceDetectionResponse) { + if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ + R.fail(R.Status.NO_FACE_DETECTED); + } + List livenessStatusList = new ArrayList(); + for(DetectionInfo detectionInfo : faceDetectionResponse.getDetectionInfoList()){ + R result = detect(image, detectionInfo.getDetectionRectangle()); + if(!result.isSuccess()){ + return R.fail(result.getCode(), result.getMessage()); + } + livenessStatusList.add(result.getData()); + } + return R.ok(livenessStatusList); + } + + @Override + public R detect(Image image, DetectionRectangle faceDetectionRectangle) { + Predictor predictor = null; + //预处理图片 + Image processedImage = null; + try { + predictor = predictorPool.borrowObject(); + if(config.getModelEnum() == LivenessModelEnum.IIC_FL_MODEL){ + processedImage = new DJLImagePreprocessor(image, faceDetectionRectangle) + .setExtendRatio(96f / 112f) + .enableSquarePadding(true) + .enableScaling(true) + .setTargetSize(128) + .enableCenterCrop(true) + .setCenterCropSize(112) + .process(); + } + Float result = null; + if(processedImage != null){ + result = predictor.predict(processedImage); + ImageUtils.releaseOpenCVMat(processedImage); + }else{ + result = predictor.predict(image); + } + LivenessStatus status = result >= config.getRealityThreshold() ? LivenessStatus.LIVE : LivenessStatus.NON_LIVE; + return R.ok(new LivenessResult(status, result)); + } catch (Exception e) { + throw new FaceException("活体检测错误", e); + }finally { + if (predictor != null) { + try { + predictorPool.returnObject(predictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + @Override + public R detectTopFace(Image image) { + R faceDetectionResponse = config.getDetectModel().detect(image); + if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ + return R.fail(R.Status.NO_FACE_DETECTED); + } + return detect(image, faceDetectionResponse.getData().getDetectionInfoList().get(0).getDetectionRectangle()); + } + @Override public GenericObjectPool> getPool() { return predictorPool; @@ -427,6 +278,9 @@ public class CommonLivenessModel implements LivenessDetModel{ @Override public void close() throws Exception { + if (fromFactory) { + LivenessModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -441,6 +295,15 @@ public class CommonLivenessModel implements LivenessDetModel{ } catch (Exception e) { log.warn("关闭 model 失败", e); } + } + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; } } diff --git a/face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java index 5256c4e..58755b8 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/liveness/LivenessDetModel.java @@ -26,254 +26,6 @@ public interface LivenessDetModel extends AutoCloseable{ */ void loadModel(LivenessConfig config); // 加载模型 - - /** - * 活体检测(多人脸) - * @param imagePath 图片路径 - * @return - */ - default R detect(String imagePath){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(多人脸) - * @param image BufferedImage - * @return - */ - default R detect(BufferedImage image){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(多人脸) - * @param imageData 图片字节流 - * @return - */ - default R detect(byte[] imageData){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(多人脸) - * @param base64Image - * @return - */ - default R detectBase64(String base64Image){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(多人脸) - * @param imagePath 图片路径 - * @param faceDetectionResponse 人脸检测结果 - * @return - */ - default R> detect(String imagePath, DetectionResponse faceDetectionResponse){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(多人脸) - * @param imageData 图片数据 - * @param faceDetectionResponse 人脸检测结果 - * @return - */ - default R> detect(byte[] imageData,DetectionResponse faceDetectionResponse){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(多人脸) - * @param image BufferedImage - * @param faceDetectionResponse 人脸检测结果 - * @return - */ - default R> detect(BufferedImage image,DetectionResponse faceDetectionResponse){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(多人脸) - * @param base64Image - * @param faceDetectionResponse 人脸检测结果 - * @return - */ - default R> detectBase64(String base64Image,DetectionResponse faceDetectionResponse){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(单人脸) - * @param imagePath 图片路径 - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detect(String imagePath, DetectionRectangle faceDetectionRectangle, List keyPoints){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(单人脸) - * @param imageData - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detect(byte[] imageData, DetectionRectangle faceDetectionRectangle, List keyPoints){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - - - /** - * 活体检测(单人脸) - * @param image BufferedImage - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(单人脸) - * @param base64Image - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detectBase64(String base64Image, DetectionRectangle faceDetectionRectangle, List keyPoints){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(单人脸) - * @param imagePath 图片路径 - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detect(String imagePath, DetectionRectangle faceDetectionRectangle){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(单人脸) - * @param imageData - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detect(byte[] imageData, DetectionRectangle faceDetectionRectangle){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - - - /** - * 活体检测(单人脸) - * @param image BufferedImage - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(单人脸) - * @param base64Image - * @param faceDetectionRectangle 人脸检测结果-人脸框 - * @return - */ - default R detectBase64(String base64Image, DetectionRectangle faceDetectionRectangle){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(分数最高人脸) - * @param image - * @return - */ - default R detectTopFace(BufferedImage image){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(分数最高人脸) - * @param imagePath - * @return - */ - default R detectTopFace(String imagePath){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - /** - * 活体检测(分数最高人脸) - * @param imageData - * @return - */ - default R detectTopFace(byte[] imageData){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - /** - * 活体检测(分数最高人脸) - * @param base64Image - * @return - */ - default R detectTopFaceBase64(String base64Image){ - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - - /** - * 视频活体检测(逐帧检测) - * @param frameImage - * @param faceDetectionRectangle - * @return - */ -// default R detectVideoByFrame(BufferedImage frameImage, DetectionRectangle faceDetectionRectangle, List keyPoints){ -// throw new UnsupportedOperationException("默认不支持该功能"); -// } - - /** - * 视频活体检测(逐帧检测) - * @param frameData - * @param faceDetectionRectangle - * @return - */ -// default R detectVideoByFrame(byte[] frameData, DetectionRectangle faceDetectionRectangle, List keyPoints){ -// throw new UnsupportedOperationException("默认不支持该功能"); -// } - - /** - * 视频活体检测(逐帧检测) - * @param frameImageData - * @return - */ -// default R detectVideoByFrame(byte[] frameImageData){ -// throw new UnsupportedOperationException("默认不支持该功能"); -// } - - /** - * 视频活体检测(逐帧检测) - * @param frameImageData - * @return - */ -// default R detectVideoByFrame(BufferedImage frameImageData){ -// throw new UnsupportedOperationException("默认不支持该功能"); -// } - /** * 视频活体检测 * @param videoInputStream @@ -293,6 +45,53 @@ public interface LivenessDetModel extends AutoCloseable{ } + /** + * 活体检测(多人脸) + * @param image + * @return + */ + default R detect(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 活体检测(多人脸) + * @param image + * @param faceDetectionResponse 人脸检测结果 + * @return + */ + default R> detect(Image image, DetectionResponse faceDetectionResponse){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 活体检测(单人脸) + * @param image + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return + */ + default R detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 活体检测(单人脸) + * @param image + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return + */ + default R detect(Image image, DetectionRectangle faceDetectionRectangle){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 活体检测(分数最高人脸) + * @param image + * @return + */ + default R detectTopFace(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } default GenericObjectPool> getPool() { @@ -300,5 +99,10 @@ public interface LivenessDetModel extends AutoCloseable{ } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + } diff --git a/face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java index 2646f8c..247a6c9 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/liveness/MiniVisionLivenessModel.java @@ -12,20 +12,18 @@ import ai.djl.repository.zoo.ZooModel; import ai.djl.training.util.ProgressBar; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.face.FaceInfo; -import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.entity.face.LivenessResult; +import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.face.LivenessStatus; import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.preprocess.BufferedImagePreprocessor; -import cn.smartjavaai.common.utils.ArrayUtils; -import cn.smartjavaai.common.utils.Base64ImageUtils; -import cn.smartjavaai.common.utils.FileUtils; -import cn.smartjavaai.common.utils.ImageUtils; +import cn.smartjavaai.common.preprocess.DJLImagePreprocessor; +import cn.smartjavaai.common.utils.*; import cn.smartjavaai.face.config.LivenessConfig; import cn.smartjavaai.face.constant.MiniVisionConstant; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.LivenessModelFactory; import cn.smartjavaai.face.model.liveness.translator.MiniVisionTranslator; -import cn.smartjavaai.common.utils.OpenCVUtils; import com.seeta.sdk.*; import lombok.extern.slf4j.Slf4j; import org.apache.commons.lang3.StringUtils; @@ -151,10 +149,7 @@ public class MiniVisionLivenessModel extends CommonLivenessModel{ @Override - public R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle) { - if(!ImageUtils.isImageValid(image)){ - throw new FaceException("图像无效"); - } + public R detect(Image image, DetectionRectangle faceDetectionRectangle) { Predictor predictor = null; Predictor sePredictor = null; try { @@ -162,30 +157,27 @@ public class MiniVisionLivenessModel extends CommonLivenessModel{ float[] seResult = null; if(Objects.nonNull(predictorPool)){ //预处理图片 - BufferedImage processedImage = new BufferedImagePreprocessor(image, faceDetectionRectangle) + Image processedImage = new DJLImagePreprocessor(image, faceDetectionRectangle) .setExtendRatio(2.7f) .enableSquarePadding(true) .enableScaling(true) .setTargetSize(80) .process(); predictor = predictorPool.borrowObject(); - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(processedImage)); - result = predictor.predict(djlImage); - ((Mat)djlImage.getWrappedImage()).release(); - + result = predictor.predict(processedImage); + ImageUtils.releaseOpenCVMat(processedImage); } if(Objects.nonNull(sePredictorPool)){ //预处理图片 - BufferedImage processedImage = new BufferedImagePreprocessor(image, faceDetectionRectangle) + Image processedImage = new DJLImagePreprocessor(image, faceDetectionRectangle) .setExtendRatio(4) .enableSquarePadding(true) .enableScaling(true) .setTargetSize(80) .process(); - Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(processedImage)); sePredictor = sePredictorPool.borrowObject(); - seResult = sePredictor.predict(djlImage); - ((Mat)djlImage.getWrappedImage()).release(); + seResult = sePredictor.predict(processedImage); + ImageUtils.releaseOpenCVMat(processedImage); } if(Objects.isNull(result) && Objects.isNull(seResult)){ throw new FaceException("活体检测错误"); @@ -195,15 +187,12 @@ public class MiniVisionLivenessModel extends CommonLivenessModel{ BigDecimal score = Objects.isNull(result) ? BigDecimal.ZERO : BigDecimal.valueOf(result[maxIndex]); BigDecimal seScore = Objects.isNull(seResult) ? BigDecimal.ZERO : BigDecimal.valueOf(seResult[maxIndex]); BigDecimal avgSocre = score.add(seScore).divide(BigDecimal.valueOf(2), 2, RoundingMode.HALF_UP); + //活体 if(maxIndex == 1){ - if(avgSocre.floatValue() >= config.getRealityThreshold()){ - return R.ok(new LivenessResult(LivenessStatus.LIVE, avgSocre.floatValue())); - }else{ - float nonLiveScore = BigDecimal.ONE.subtract(avgSocre).floatValue(); - return R.ok(new LivenessResult(LivenessStatus.NON_LIVE, nonLiveScore)); - } - }else{ - return R.ok(new LivenessResult(LivenessStatus.NON_LIVE, avgSocre.floatValue())); + LivenessStatus livenessStatus = avgSocre.floatValue() > config.getRealityThreshold() ? LivenessStatus.LIVE : LivenessStatus.NON_LIVE; + return R.ok(new LivenessResult(livenessStatus, avgSocre.floatValue())); + }else{//非活体 + return R.ok(new LivenessResult(LivenessStatus.NON_LIVE, BigDecimal.ONE.subtract(avgSocre).floatValue())); } } catch (Exception e) { throw new FaceException("活体检测错误", e); @@ -245,6 +234,9 @@ public class MiniVisionLivenessModel extends CommonLivenessModel{ @Override public void close() throws Exception { + if (fromFactory) { + LivenessModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -284,4 +276,14 @@ public class MiniVisionLivenessModel extends CommonLivenessModel{ MINIFASNET_V1_SE, FUSION // 融合模型 } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java b/face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java index 119f22c..ec1ae5e 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/liveness/Seetaface6LivenessModel.java @@ -1,21 +1,27 @@ package cn.smartjavaai.face.model.liveness; import ai.djl.engine.Engine; +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.face.FaceInfo; import cn.smartjavaai.common.entity.face.LivenessResult; import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.face.config.LivenessConfig; import cn.smartjavaai.common.enums.face.LivenessStatus; import cn.smartjavaai.face.constant.LivenessConstant; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.LivenessModelFactory; import cn.smartjavaai.face.seetaface.NativeLoader; import cn.smartjavaai.face.utils.FaceUtils; +import cn.smartjavaai.face.utils.Seetaface6Utils; import com.seeta.pool.*; import com.seeta.sdk.*; import lombok.extern.slf4j.Slf4j; +import nu.pattern.OpenCV; import org.apache.commons.lang3.StringUtils; import org.bytedeco.javacv.FFmpegFrameGrabber; import org.bytedeco.javacv.Frame; @@ -44,6 +50,11 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ private LivenessConfig config; + static { + //视频功能需要 + OpenCV.loadLocally(); + } + @Override public void loadModel(LivenessConfig config) { if(StringUtils.isBlank(config.getModelPath())){ @@ -129,10 +140,7 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ } - private R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints, boolean isImage) { - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } + private R detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints, boolean isImage) { if(Objects.isNull(faceDetectionRectangle)){ return R.fail(R.Status.NO_FACE_DETECTED); } @@ -145,8 +153,8 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ faceAntiSpoofing = faceAntiSpoofingPool.borrowObject(); SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(faceDetectionRectangle); - SeetaPointF[] landmarks = FaceUtils.convertToSeetaPointF(keyPoints); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(faceDetectionRectangle); + SeetaPointF[] landmarks = Seetaface6Utils.convertToSeetaPointF(keyPoints); //检测图片 if(isImage){ status = faceAntiSpoofing.Predict(imageData, seetaRect, landmarks); @@ -154,7 +162,7 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ //检测视频 status = faceAntiSpoofing.PredictVideo(imageData, seetaRect, landmarks); } - return R.ok(new LivenessResult(FaceUtils.convertToLivenessStatus(status))); + return R.ok(new LivenessResult(Seetaface6Utils.convertToLivenessStatus(status))); } catch (Exception e) { throw new FaceException("活体检测错误", e); } finally { @@ -168,38 +176,61 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ } } - @Override - public R detect(String imagePath) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + + private R detectVideo(FFmpegFrameGrabber grabber) { + FaceAntiSpoofing faceAntiSpoofing = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); + faceAntiSpoofing = faceAntiSpoofingPool.borrowObject(); + //重置视频 + faceAntiSpoofing.ResetVideo(); + grabber.start(); + // 获取视频总帧数 + int totalFrames = grabber.getLengthInFrames(); + int videoFrameCountConfig = faceAntiSpoofing.GetVideoFrameCount(); + log.debug("视频总帧数:{},检测帧数:{}", totalFrames, videoFrameCountConfig); + if(totalFrames < videoFrameCountConfig){ + return R.fail(1001, "视频帧数低于检测帧数"); + } + // 逐帧处理视频 + for (int frameIndex = 0; frameIndex < totalFrames; frameIndex++) { + if(frameIndex >= config.getMaxVideoDetectFrames()){ + return R.fail(10002, "超出最大检测帧数:" + config.getMaxVideoDetectFrames()); + } + // 获取当前帧 + Frame frame = grabber.grabImage(); + if (frame != null) { + BufferedImage bufferedImage = Java2DFrameUtils.toBufferedImage(frame); + Image image = SmartImageFactory.getInstance().fromBufferedImage(bufferedImage); + R livenessStatus = detectTopFace(image, false); + if(!livenessStatus.isSuccess()){ + log.debug("第" + frameIndex + "帧处理失败:" + livenessStatus.getMessage()); + continue; + } + //满足检测帧数之后停止检测 + if(livenessStatus.getData().getStatus() != LivenessStatus.DETECTING){ + return livenessStatus; + } + } + } + grabber.stop(); + } catch (FFmpegFrameGrabber.Exception e) { + throw new FaceException(e); + } catch (Exception e) { + throw new FaceException(e); + } finally { + if (faceAntiSpoofing != null) { + try { + faceAntiSpoofingPool.returnObject(faceAntiSpoofing); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } } - return detect(image); + return R.fail(1000, "有效帧数量不足,无法完成活体检测"); } @Override - public R detect(byte[] imageData) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData))); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @Override - public R detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } + public R detect(Image image) { FaceAntiSpoofing faceAntiSpoofing = null; FaceLandmarker faceLandmarker = null; FaceDetector detectPredictor = null; @@ -224,9 +255,9 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ seetaPointFSList.add(landmarks); //检测图片 FaceAntiSpoofing.Status status = faceAntiSpoofing.Predict(imageData, seetaRect, landmarks); - livenessStatusList.add(FaceUtils.convertToLivenessStatus(status)); + livenessStatusList.add(Seetaface6Utils.convertToLivenessStatus(status)); } - return R.ok(FaceUtils.convertToDetectionResponse(seetaResult, seetaPointFSList, livenessStatusList)); + return R.ok(Seetaface6Utils.convertToDetectionResponse(seetaResult, seetaPointFSList, livenessStatusList)); }else{ R faceDetectionResponse = config.getDetectModel().detect(image); if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ @@ -272,26 +303,7 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ } @Override - public R> detect(String imagePath, DetectionResponse faceDetectionResponse) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detect(image, faceDetectionResponse); - } - - - @Override - public R> detect(BufferedImage image, DetectionResponse faceDetectionResponse) { - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } + public R> detect(Image image, DetectionResponse faceDetectionResponse) { if(Objects.isNull(faceDetectionResponse) || Objects.isNull(faceDetectionResponse.getDetectionInfoList()) || faceDetectionResponse.getDetectionInfoList().isEmpty()){ return R.fail(R.Status.NO_FACE_DETECTED); } @@ -310,73 +322,17 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ return R.ok(livenessStatusList); } - - @Override - public R> detect(byte[] imageData, DetectionResponse faceDetectionResponse) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionResponse); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - @Override - public R detect(String imagePath, DetectionRectangle faceDetectionRectangle, List keyPoints) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detect(image, faceDetectionRectangle, keyPoints); - } - - @Override - public R detect(BufferedImage image, DetectionRectangle faceDetectionRectangle, List keyPoints) { + public R detect(Image image, DetectionRectangle faceDetectionRectangle, List keyPoints) { return detect(image, faceDetectionRectangle, keyPoints, true); } @Override - public R detect(byte[] imageData, DetectionRectangle faceDetectionRectangle, List keyPoints) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionRectangle, keyPoints); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } + public R detectTopFace(Image image) { + return detectTopFace(image, true); } - - @Override - public R detectTopFace(String imagePath) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return detectTopFace(image); - } - - - private R detectTopFace(BufferedImage image, boolean isImage) { - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } + private R detectTopFace(Image image, boolean isImage) { FaceAntiSpoofing faceAntiSpoofing = null; FaceLandmarker faceLandmarker = null; FaceDetector detectPredictor = null; @@ -402,14 +358,15 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ }else{ status = faceAntiSpoofing.PredictVideo(imageData, seetaResult[0], landmarks); } - return R.ok(new LivenessResult(FaceUtils.convertToLivenessStatus(status))); + return R.ok(new LivenessResult(Seetaface6Utils.convertToLivenessStatus(status))); }else{ R faceDetectionResponse = config.getDetectModel().detect(image); if(Objects.isNull(faceDetectionResponse.getData()) || Objects.isNull(faceDetectionResponse.getData().getDetectionInfoList()) || faceDetectionResponse.getData().getDetectionInfoList().isEmpty()){ return R.fail(R.Status.NO_FACE_DETECTED); } DetectionInfo detectionInfo = faceDetectionResponse.getData().getDetectionInfoList().get(0); - return detect(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getKeyPoints(), isImage); + R detectionResponseR = detect(image, detectionInfo.getDetectionRectangle(), detectionInfo.getFaceInfo().getKeyPoints(), isImage); + return detectionResponseR; } } catch (Exception e) { throw new FaceException("活体检测错误", e); @@ -438,58 +395,6 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ } } - - @Override - public R detectTopFace(BufferedImage image) { - return detectTopFace(image, true); - } - - @Override - public R detectTopFace(byte[] imageData) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detectTopFace(ImageIO.read(new ByteArrayInputStream(imageData))); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - - public R detectVideoByFrame(BufferedImage frameImage, DetectionRectangle faceDetectionRectangle, List keyPoints) { - if(!ImageUtils.isImageValid(frameImage)){ - return R.fail(R.Status.INVALID_IMAGE); - } - return detect(frameImage,faceDetectionRectangle, keyPoints,false); - } - - public R detectVideoByFrame(byte[] frameData, DetectionRectangle faceDetectionRectangle, List keyPoints) { - if(Objects.isNull(frameData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detect(ImageIO.read(new ByteArrayInputStream(frameData)), faceDetectionRectangle, keyPoints, false); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - public R detectVideoByFrame(byte[] frameImageData) { - if(Objects.isNull(frameImageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return detectVideoByFrame(ImageIO.read(new ByteArrayInputStream(frameImageData))); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - public R detectVideoByFrame(BufferedImage frameImageData) { - return detectTopFace(frameImageData, false); - } - @Override public R detectVideo(InputStream videoInputStream) { if(Objects.isNull(videoInputStream)){ @@ -506,57 +411,6 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ return detectVideo(new FFmpegFrameGrabber(videoPath)); } - private R detectVideo(FFmpegFrameGrabber grabber) { - FaceAntiSpoofing faceAntiSpoofing = null; - try { - faceAntiSpoofing = faceAntiSpoofingPool.borrowObject(); - //重置视频 - faceAntiSpoofing.ResetVideo(); - grabber.start(); - // 获取视频总帧数 - int totalFrames = grabber.getLengthInFrames(); - int videoFrameCountConfig = faceAntiSpoofing.GetVideoFrameCount(); - log.debug("视频总帧数:{},检测帧数:{}", totalFrames, videoFrameCountConfig); - if(totalFrames < videoFrameCountConfig){ - return R.fail(1001, "视频帧数低于检测帧数"); - } - // 逐帧处理视频 - for (int frameIndex = 0; frameIndex < totalFrames; frameIndex++) { - if(frameIndex >= config.getMaxVideoDetectFrames()){ - return R.fail(10002, "超出最大检测帧数:" + config.getMaxVideoDetectFrames()); - } - // 获取当前帧 - Frame frame = grabber.grabImage(); - if (frame != null) { - BufferedImage bufferedImage = Java2DFrameUtils.toBufferedImage(frame); - R livenessStatus = detectVideoByFrame(bufferedImage); - if(!livenessStatus.isSuccess()){ - log.debug("第" + frameIndex + "帧处理失败:" + livenessStatus.getMessage()); - continue; - } - //满足检测帧数之后停止检测 - if(livenessStatus.getData().getStatus() != LivenessStatus.DETECTING){ - return livenessStatus; - } - } - } - grabber.stop(); - } catch (FFmpegFrameGrabber.Exception e) { - throw new FaceException(e); - } catch (Exception e) { - throw new FaceException(e); - } finally { - if (faceAntiSpoofing != null) { - try { - faceAntiSpoofingPool.returnObject(faceAntiSpoofing); - } catch (Exception e) { - log.warn("归还Predictor失败", e); - } - } - } - return R.fail(1000, "有效帧数量不足,无法完成活体检测"); - } - public FaceDetectorPool getFaceDetectorPool() { return faceDetectorPool; } @@ -571,6 +425,9 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ @Override public void close() throws Exception { + if (fromFactory) { + LivenessModelFactory.removeFromCache(config.getModelEnum()); + } try { if (faceDetectorPool != null) { faceDetectorPool.close(); @@ -593,4 +450,14 @@ public class Seetaface6LivenessModel implements LivenessDetModel{ log.warn("关闭 predictorPool 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/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 index c3003c7..1c22225 100644 --- a/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 @@ -12,7 +12,6 @@ import cn.smartjavaai.face.config.LivenessConfig; import cn.smartjavaai.face.constant.FaceNetConstant; import cn.smartjavaai.face.enums.FaceRecModelEnum; import cn.smartjavaai.face.enums.LivenessModelEnum; -import cn.smartjavaai.face.model.facerec.translator.FaceNetRecTranslator; import cn.smartjavaai.face.model.liveness.translator.IicFrTranslator; import org.apache.commons.lang3.StringUtils; diff --git a/face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java b/face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java index 301cc38..3bf5b5e 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/quality/FaceQualityModel.java @@ -1,5 +1,6 @@ package cn.smartjavaai.face.model.quality; +import ai.djl.modality.cv.Image; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.Point; import cn.smartjavaai.common.entity.R; @@ -31,6 +32,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateBrightness(BufferedImage image, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -42,6 +44,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateBrightness(String imagePath, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -53,6 +56,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateBrightness(byte[] imageData, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -65,6 +69,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateClarity(BufferedImage image, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -76,6 +81,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateClarity(String imagePath, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -87,6 +93,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateClarity(byte[] imageData, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -98,6 +105,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateCompleteness(BufferedImage image, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -109,6 +117,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateCompleteness(String imagePath, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -120,6 +129,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateCompleteness(byte[] imageData, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -131,6 +141,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluatePose(BufferedImage image, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -142,6 +153,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluatePose(String imagePath, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -153,6 +165,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluatePose(byte[] imageData, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -164,6 +177,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateResolution(BufferedImage image, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -175,6 +189,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateResolution(String imagePath, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -186,6 +201,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateResolution(byte[] imageData, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -198,6 +214,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateAll(String imagePath, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -209,6 +226,7 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateAll(BufferedImage image, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -221,9 +239,80 @@ public interface FaceQualityModel extends AutoCloseable{ * @param keyPoints * @return */ + @Deprecated default R evaluateAll(byte[] imageData, DetectionRectangle rectangle, List keyPoints){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 亮度评估 + * @param image + * @param rectangle + * @param keyPoints + * @return + */ + default R evaluateBrightness(Image image, DetectionRectangle rectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 清晰度评估 + * @param image + * @param rectangle + * @param keyPoints + * @return + */ + default R evaluateClarity(Image image, DetectionRectangle rectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 完整度评估 + * @param image + * @param rectangle + * @param keyPoints + * @return + */ + default R evaluateCompleteness(Image image, DetectionRectangle rectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 人脸姿态评估 + * @param image + * @param rectangle + * @param keyPoints + * @return + */ + default R evaluatePose(Image image, DetectionRectangle rectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 人脸分辨率评估 + * @param image + * @param rectangle + * @param keyPoints + * @return + */ + default R evaluateResolution(Image image, DetectionRectangle rectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 评估所有 + * @param image + * @param rectangle + * @param keyPoints + * @return + */ + default R evaluateAll(Image image, DetectionRectangle rectangle, List keyPoints){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java b/face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java index 976b20b..5fdeebe 100644 --- a/face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java +++ b/face/src/main/java/cn/smartjavaai/face/model/quality/Seetaface6QualityModel.java @@ -1,8 +1,12 @@ package cn.smartjavaai.face.model.quality; import ai.djl.engine.Engine; +import ai.djl.modality.cv.Image; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; +import cn.smartjavaai.common.entity.face.ExpressionResult; import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.PoolUtils; @@ -11,9 +15,12 @@ import cn.smartjavaai.face.entity.FaceQualityResult; import cn.smartjavaai.face.entity.FaceQualitySummary; import cn.smartjavaai.face.enums.QualityGrade; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.factory.FaceQualityModelFactory; +import cn.smartjavaai.face.factory.LivenessModelFactory; import cn.smartjavaai.face.seetaface.ClarityDLResult; import cn.smartjavaai.face.seetaface.NativeLoader; import cn.smartjavaai.face.utils.FaceUtils; +import cn.smartjavaai.face.utils.Seetaface6Utils; import com.seeta.pool.*; import com.seeta.sdk.*; import lombok.extern.slf4j.Slf4j; @@ -98,42 +105,10 @@ public class Seetaface6QualityModel implements FaceQualityModel { @Override public R evaluateBrightness(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { - //参数检查 - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(rectangle)){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); - } - if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); - } - QualityOfBrightness qualityOfBrightness = null; - try { - if(Objects.isNull(this.qualityOfBrightnessPool)){ - this.qualityOfBrightnessPool = new QualityOfBrightnessPool(new SeetaConfSetting()); - qualityOfBrightnessPool.setMaxTotal(predictorPoolSize); - } - qualityOfBrightness = qualityOfBrightnessPool.borrowObject(); - SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); - float[] scores = new float[1]; - QualityOfBrightness.QualityLevel level = qualityOfBrightness.check(imageData, seetaRect, pointFS, scores); - FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); - return R.ok(result); - } catch (Exception e) { - throw new FaceException("亮度评估错误", e); - } finally { - if (qualityOfBrightness != null) { - try { - qualityOfBrightnessPool.returnObject(qualityOfBrightness); - } catch (Exception e) { - log.warn("归还Predictor失败", e); - } - } - } + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = evaluateBrightness(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -141,14 +116,16 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = evaluateBrightness(image, rectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return evaluateBrightness(image, rectangle, keyPoints); } @Override @@ -156,51 +133,23 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return evaluateBrightness(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = evaluateBrightness(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override public R evaluateClarity(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { - //参数检查 - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(rectangle)){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); - } - if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); - } - QualityOfClarity qualityOfClarity = null; - try { - if(Objects.isNull(this.qualityOfClarityPool)){ - this.qualityOfClarityPool = new QualityOfClarityPool(new SeetaConfSetting()); - qualityOfClarityPool.setMaxTotal(predictorPoolSize); - } - qualityOfClarity = qualityOfClarityPool.borrowObject(); - SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); - float[] scores = new float[1]; - QualityOfClarity.QualityLevel level = qualityOfClarity.check(imageData, seetaRect, pointFS, scores); - FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); - return R.ok(result); - } catch (Exception e) { - throw new FaceException("清晰度评估错误", e); - } finally { - if (qualityOfClarity != null) { - try { - qualityOfClarityPool.returnObject(qualityOfClarity); - } catch (Exception e) { - log.warn("归还Predictor失败", e); - } - } - } + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = evaluateClarity(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -208,14 +157,16 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = evaluateClarity(image, rectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return evaluateClarity(image, rectangle, keyPoints); } @Override @@ -223,51 +174,23 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return evaluateClarity(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = evaluateClarity(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override public R evaluateCompleteness(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { - //参数检查 - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(rectangle)){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); - } - if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); - } - QualityOfIntegrity qualityOfIntegrity = null; - try { - if(Objects.isNull(this.qualityOfIntegrityPool)){ - this.qualityOfIntegrityPool = new QualityOfIntegrityPool(new SeetaConfSetting()); - qualityOfIntegrityPool.setMaxTotal(predictorPoolSize); - } - qualityOfIntegrity = qualityOfIntegrityPool.borrowObject(); - SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); - float[] scores = new float[1]; - QualityOfIntegrity.QualityLevel level = qualityOfIntegrity.check(imageData, seetaRect, pointFS, scores); - FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); - return R.ok(result); - } catch (Exception e) { - throw new FaceException("完整度评估错误", e); - } finally { - if (qualityOfIntegrity != null) { - try { - qualityOfIntegrityPool.returnObject(qualityOfIntegrity); - } catch (Exception e) { - log.warn("归还Predictor失败", e); - } - } - } + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = evaluateCompleteness(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -275,14 +198,16 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = evaluateCompleteness(image, rectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return evaluateCompleteness(image, rectangle, keyPoints); } @Override @@ -290,51 +215,23 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return evaluateCompleteness(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = evaluateCompleteness(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override public R evaluatePose(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { - //参数检查 - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(rectangle)){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); - } - if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); - } - QualityOfPose qualityOfPose = null; - try { - if(Objects.isNull(this.qualityOfPosePool)){ - this.qualityOfPosePool = new QualityOfPosePool(new SeetaConfSetting()); - qualityOfPosePool.setMaxTotal(predictorPoolSize); - } - qualityOfPose = qualityOfPosePool.borrowObject(); - SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); - float[] scores = new float[1]; - QualityOfPose.QualityLevel level = qualityOfPose.check(imageData, seetaRect, pointFS, scores); - FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); - return R.ok(result); - } catch (Exception e) { - throw new FaceException("姿态评估错误", e); - } finally { - if (qualityOfPose != null) { - try { - qualityOfPosePool.returnObject(qualityOfPose); - } catch (Exception e) { - log.warn("归还Predictor失败", e); - } - } - } + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = evaluatePose(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -342,14 +239,16 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = evaluatePose(image, rectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return evaluatePose(image, rectangle, keyPoints); } @Override @@ -357,51 +256,23 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return evaluatePose(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = evaluatePose(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override public R evaluateResolution(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { - //参数检查 - if(!ImageUtils.isImageValid(image)){ - return R.fail(R.Status.INVALID_IMAGE); - } - if(Objects.isNull(rectangle)){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); - } - if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ - return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); - } - QualityOfResolution qualityOfResolution = null; - try { - if(Objects.isNull(this.qualityOfResolutionPool)){ - this.qualityOfResolutionPool = new QualityOfResolutionPool(new SeetaConfSetting()); - qualityOfResolutionPool.setMaxTotal(predictorPoolSize); - } - qualityOfResolution = qualityOfResolutionPool.borrowObject(); - SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); - float[] scores = new float[1]; - QualityOfResolution.QualityLevel level = qualityOfResolution.check(imageData, seetaRect, pointFS, scores); - FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); - return R.ok(result); - } catch (Exception e) { - throw new FaceException("姿态评估错误", e); - } finally { - if (qualityOfResolution != null) { - try { - qualityOfResolutionPool.returnObject(qualityOfResolution); - } catch (Exception e) { - log.warn("归还Predictor失败", e); - } - } - } + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = evaluateResolution(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } @Override @@ -409,14 +280,16 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = evaluateResolution(image, rectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return evaluateResolution(image, rectangle, keyPoints); } @Override @@ -424,17 +297,21 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(Objects.isNull(imageData)){ return R.fail(R.Status.INVALID_IMAGE); } + Image imageDjl = null; try { - return evaluateResolution(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new FaceException("错误的图像", e); + throw new RuntimeException(e); } + R detectionResponseR = evaluateResolution(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; } - public R evaluateClarityWithDL(BufferedImage image, List keyPoints) { + public R evaluateClarityWithDL(Image image, List keyPoints) { //参数检查 - if(!ImageUtils.isImageValid(image)){ + if(Objects.isNull(image)){ return R.fail(R.Status.INVALID_IMAGE); } if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ @@ -456,7 +333,7 @@ public class Seetaface6QualityModel implements FaceQualityModel { qualityOfLBN = qualityOfLBNPool.borrowObject(); SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); imageData.data = ImageUtils.getMatrixBGR(image); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); int[] light = new int[1]; int[] blur = new int[1]; int[] noise = new int[1]; @@ -476,34 +353,9 @@ public class Seetaface6QualityModel implements FaceQualityModel { } - public R evaluateClarityWithDL(String imagePath, List keyPoints) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return evaluateClarityWithDL(image, keyPoints); - } - - public R evaluateClarityWithDL(byte[] imageData, List keyPoints) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return evaluateClarityWithDL(ImageIO.read(new ByteArrayInputStream(imageData)), keyPoints); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - - public R evaluatePoseWithDL(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { + public R evaluatePoseWithDL(Image image, DetectionRectangle rectangle, List keyPoints) { //参数检查 - if(!ImageUtils.isImageValid(image)){ + if(Objects.isNull(image)){ return R.fail(R.Status.INVALID_IMAGE); } if(Objects.isNull(rectangle)){ @@ -528,8 +380,8 @@ public class Seetaface6QualityModel implements FaceQualityModel { qualityOfPoseEx = qualityOfPoseExPool.borrowObject(); SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); imageData.data = ImageUtils.getMatrixBGR(image); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); float[] scores = new float[1]; QualityOfPoseEx.QualityLevel level = qualityOfPoseEx.check(imageData, seetaRect, pointFS, scores); FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); @@ -547,30 +399,6 @@ public class Seetaface6QualityModel implements FaceQualityModel { } } - public R evaluatePoseWithDL(String imagePath, DetectionRectangle rectangle, List keyPoints) { - if(!FileUtils.isFileExists(imagePath)){ - return R.fail(R.Status.FILE_NOT_FOUND); - } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; - try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); - } catch (IOException e) { - throw new FaceException("无效图片路径", e); - } - return evaluatePoseWithDL(image, rectangle, keyPoints); - } - - public R evaluatePoseWithDL(byte[] imageData, DetectionRectangle rectangle, List keyPoints) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return evaluatePoseWithDL(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } /** * 获取清晰度模型配置(深度学习) @@ -613,20 +441,246 @@ public class Seetaface6QualityModel implements FaceQualityModel { if(!FileUtils.isFileExists(imagePath)){ return R.fail(R.Status.FILE_NOT_FOUND); } - // 将图片路径转换为 BufferedImage - BufferedImage image = null; + Image image = null; try { - image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + image = SmartImageFactory.getInstance().fromFile(imagePath); + R detectionResponseR = evaluateAll(image, rectangle, keyPoints); + return detectionResponseR; } catch (IOException e) { throw new FaceException("无效图片路径", e); + }finally { + ImageUtils.releaseOpenCVMat(image); } - return evaluateAll(image, rectangle, keyPoints); } @Override public R evaluateAll(BufferedImage image, DetectionRectangle rectangle, List keyPoints) { + Image imageDjl = SmartImageFactory.getInstance().fromBufferedImage(image); + R detectionResponseR = evaluateAll(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; + } + + @Override + public R evaluateAll(byte[] imageData, DetectionRectangle rectangle, List keyPoints) { + if(Objects.isNull(imageData)){ + return R.fail(R.Status.INVALID_IMAGE); + } + Image imageDjl = null; + try { + imageDjl = SmartImageFactory.getInstance().fromBytes(imageData); + } catch (IOException e) { + throw new RuntimeException(e); + } + R detectionResponseR = evaluateAll(imageDjl, rectangle, keyPoints); + ImageUtils.releaseOpenCVMat(imageDjl); + return detectionResponseR; + } + + @Override + public R evaluateBrightness(Image image, DetectionRectangle rectangle, List keyPoints) { //参数检查 - if(!ImageUtils.isImageValid(image)){ + if(Objects.isNull(image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + if(Objects.isNull(rectangle)){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); + } + if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); + } + QualityOfBrightness qualityOfBrightness = null; + try { + if(Objects.isNull(this.qualityOfBrightnessPool)){ + this.qualityOfBrightnessPool = new QualityOfBrightnessPool(new SeetaConfSetting()); + qualityOfBrightnessPool.setMaxTotal(predictorPoolSize); + } + qualityOfBrightness = qualityOfBrightnessPool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); + float[] scores = new float[1]; + QualityOfBrightness.QualityLevel level = qualityOfBrightness.check(imageData, seetaRect, pointFS, scores); + FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); + return R.ok(result); + } catch (Exception e) { + throw new FaceException("亮度评估错误", e); + } finally { + if (qualityOfBrightness != null) { + try { + qualityOfBrightnessPool.returnObject(qualityOfBrightness); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R evaluateClarity(Image image, DetectionRectangle rectangle, List keyPoints) { + //参数检查 + if(Objects.isNull(image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + if(Objects.isNull(rectangle)){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); + } + if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); + } + QualityOfClarity qualityOfClarity = null; + try { + if(Objects.isNull(this.qualityOfClarityPool)){ + this.qualityOfClarityPool = new QualityOfClarityPool(new SeetaConfSetting()); + qualityOfClarityPool.setMaxTotal(predictorPoolSize); + } + qualityOfClarity = qualityOfClarityPool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); + float[] scores = new float[1]; + QualityOfClarity.QualityLevel level = qualityOfClarity.check(imageData, seetaRect, pointFS, scores); + FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); + return R.ok(result); + } catch (Exception e) { + throw new FaceException("清晰度评估错误", e); + } finally { + if (qualityOfClarity != null) { + try { + qualityOfClarityPool.returnObject(qualityOfClarity); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R evaluateCompleteness(Image image, DetectionRectangle rectangle, List keyPoints) { + //参数检查 + if(Objects.isNull(image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + if(Objects.isNull(rectangle)){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); + } + if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); + } + QualityOfIntegrity qualityOfIntegrity = null; + try { + if(Objects.isNull(this.qualityOfIntegrityPool)){ + this.qualityOfIntegrityPool = new QualityOfIntegrityPool(new SeetaConfSetting()); + qualityOfIntegrityPool.setMaxTotal(predictorPoolSize); + } + qualityOfIntegrity = qualityOfIntegrityPool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); + float[] scores = new float[1]; + QualityOfIntegrity.QualityLevel level = qualityOfIntegrity.check(imageData, seetaRect, pointFS, scores); + FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); + return R.ok(result); + } catch (Exception e) { + throw new FaceException("完整度评估错误", e); + } finally { + if (qualityOfIntegrity != null) { + try { + qualityOfIntegrityPool.returnObject(qualityOfIntegrity); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R evaluatePose(Image image, DetectionRectangle rectangle, List keyPoints) { + //参数检查 + if(Objects.isNull(image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + if(Objects.isNull(rectangle)){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); + } + if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); + } + QualityOfPose qualityOfPose = null; + try { + if(Objects.isNull(this.qualityOfPosePool)){ + this.qualityOfPosePool = new QualityOfPosePool(new SeetaConfSetting()); + qualityOfPosePool.setMaxTotal(predictorPoolSize); + } + qualityOfPose = qualityOfPosePool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); + float[] scores = new float[1]; + QualityOfPose.QualityLevel level = qualityOfPose.check(imageData, seetaRect, pointFS, scores); + FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); + return R.ok(result); + } catch (Exception e) { + throw new FaceException("姿态评估错误", e); + } finally { + if (qualityOfPose != null) { + try { + qualityOfPosePool.returnObject(qualityOfPose); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R evaluateResolution(Image image, DetectionRectangle rectangle, List keyPoints) { + //参数检查 + if(Objects.isNull(image)){ + return R.fail(R.Status.INVALID_IMAGE); + } + if(Objects.isNull(rectangle)){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "rectangle为空"); + } + if(Objects.isNull(keyPoints) || keyPoints.isEmpty()){ + return R.fail(R.Status.PARAM_ERROR.getCode(), "keyPoints为空"); + } + QualityOfResolution qualityOfResolution = null; + try { + if(Objects.isNull(this.qualityOfResolutionPool)){ + this.qualityOfResolutionPool = new QualityOfResolutionPool(new SeetaConfSetting()); + qualityOfResolutionPool.setMaxTotal(predictorPoolSize); + } + qualityOfResolution = qualityOfResolutionPool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = ImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); + float[] scores = new float[1]; + QualityOfResolution.QualityLevel level = qualityOfResolution.check(imageData, seetaRect, pointFS, scores); + FaceQualityResult result = new FaceQualityResult(scores[0], QualityGrade.valueOf(level.name())); + return R.ok(result); + } catch (Exception e) { + throw new FaceException("姿态评估错误", e); + } finally { + if (qualityOfResolution != null) { + try { + qualityOfResolutionPool.returnObject(qualityOfResolution); + } catch (Exception e) { + log.warn("归还Predictor失败", e); + } + } + } + } + + @Override + public R evaluateAll(Image image, DetectionRectangle rectangle, List keyPoints) { + //参数检查 + if(Objects.isNull(image)){ return R.fail(R.Status.INVALID_IMAGE); } if(Objects.isNull(rectangle)){ @@ -669,8 +723,8 @@ public class Seetaface6QualityModel implements FaceQualityModel { qualityOfResolution = qualityOfResolutionPool.borrowObject(); SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); imageData.data = ImageUtils.getMatrixBGR(image); - SeetaRect seetaRect = FaceUtils.convertToSeetaRect(rectangle); - SeetaPointF[] pointFS = FaceUtils.convertToSeetaPointF(keyPoints); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(rectangle); + SeetaPointF[] pointFS = Seetaface6Utils.convertToSeetaPointF(keyPoints); float[] scoresBrightness = new float[1]; QualityOfBrightness.QualityLevel level = qualityOfBrightness.check(imageData, seetaRect, pointFS, scoresBrightness); summary.setBrightness(new FaceQualityResult(scoresBrightness[0], QualityGrade.valueOf(level.name()))); @@ -698,18 +752,6 @@ public class Seetaface6QualityModel implements FaceQualityModel { } } - @Override - public R evaluateAll(byte[] imageData, DetectionRectangle rectangle, List keyPoints) { - if(Objects.isNull(imageData)){ - return R.fail(R.Status.INVALID_IMAGE); - } - try { - return evaluateAll(ImageIO.read(new ByteArrayInputStream(imageData)), rectangle, keyPoints); - } catch (IOException e) { - throw new FaceException("错误的图像", e); - } - } - public QualityOfBrightnessPool getQualityOfBrightnessPool() { return qualityOfBrightnessPool; } @@ -740,6 +782,9 @@ public class Seetaface6QualityModel implements FaceQualityModel { @Override public void close() throws Exception { + if (fromFactory) { + FaceQualityModelFactory.removeFromCache(config.getModelEnum()); + } if(Objects.nonNull(qualityOfBrightnessPool)){ qualityOfBrightnessPool.close(); } @@ -762,4 +807,14 @@ public class Seetaface6QualityModel implements FaceQualityModel { qualityOfResolutionPool.close(); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java b/face/src/main/java/cn/smartjavaai/face/preprocess/DJLImageFacePreprocessor.java similarity index 76% rename from face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java rename to face/src/main/java/cn/smartjavaai/face/preprocess/DJLImageFacePreprocessor.java index 9060cc5..16cb294 100644 --- a/face/src/main/java/cn/smartjavaai/face/preprocess/DJLImagePreprocessor.java +++ b/face/src/main/java/cn/smartjavaai/face/preprocess/DJLImageFacePreprocessor.java @@ -3,13 +3,14 @@ package cn.smartjavaai.face.preprocess; import ai.djl.modality.cv.Image; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDManager; -import ai.djl.opencv.OpenCVImageFactory; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.face.utils.FaceAlignUtils; import cn.smartjavaai.face.utils.FaceUtils; import org.opencv.core.Mat; +import java.awt.image.BufferedImage; import java.util.Objects; /** @@ -17,7 +18,7 @@ import java.util.Objects; * @author dwj * @date 2025/6/27 */ -public class DJLImagePreprocessor { +public class DJLImageFacePreprocessor { private Image image; @@ -33,20 +34,20 @@ public class DJLImagePreprocessor { private int affineTargetWidth; private int affineTargetHeight; - public DJLImagePreprocessor(Image image, NDManager manager) { + public DJLImageFacePreprocessor(Image image, NDManager manager) { this.image = image; this.manager = manager; } // 启用裁剪 - public DJLImagePreprocessor enableCrop(DetectionRectangle rect) { + public DJLImageFacePreprocessor enableCrop(DetectionRectangle rect) { this.enableCrop = true; this.cropRect = rect; return this; } // 启用仿射变换 - public DJLImagePreprocessor enableAffine(double[][] keyPoints, int targetWidth, int targetHeight) { + public DJLImageFacePreprocessor enableAffine(double[][] keyPoints, int targetWidth, int targetHeight) { if(Objects.isNull(keyPoints)){ throw new IllegalArgumentException("keyPoints must be not null"); } @@ -83,8 +84,11 @@ public class DJLImagePreprocessor { } // 5点仿射变换 Mat affine_matrix = OpenCVUtils.toOpenCVMat(manager, srcPoints, dstPoints); - Mat mat = FaceAlignUtils.warpAffine((Mat) image.getWrappedImage(), affine_matrix, width, height); - Image alignedImg = OpenCVImageFactory.getInstance().fromImage(mat); + Object imageObj = image.getWrappedImage(); + Mat src = image.getWrappedImage() instanceof Mat ? (Mat) imageObj : OpenCVUtils.image2Mat((BufferedImage) imageObj); + Mat mat = FaceAlignUtils.warpAffine(src, affine_matrix, width, height); + Image alignedImg = SmartImageFactory.getInstance().fromMat(mat); + affine_matrix.release(); return alignedImg; } diff --git a/face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java b/face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java index ed22020..9c7175f 100644 --- a/face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java +++ b/face/src/main/java/cn/smartjavaai/face/utils/FaceAlignUtils.java @@ -4,6 +4,7 @@ import ai.djl.modality.cv.Image; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDManager; import ai.djl.opencv.OpenCVImageFactory; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.OpenCVUtils; import com.seeta.sdk.SeetaImageData; @@ -56,7 +57,7 @@ public class FaceAlignUtils { public static SeetaImageData faceAlign(BufferedImage sourceImage, SeetaPointF[] pointFS) { NDManager manager = NDManager.newBaseManager(); //获取子图中人脸关键点坐标 - double[][] pointsArray = FaceUtils.facePoints(pointFS); + double[][] pointsArray = Seetaface6Utils.facePoints(pointFS); NDArray srcPoints = manager.create(pointsArray); NDArray dstPoints = FaceUtils.faceTemplate512x512(manager); // 5点仿射变换 @@ -64,7 +65,7 @@ public class FaceAlignUtils { Mat mat = FaceAlignUtils.warpAffine(OpenCVUtils.image2Mat(sourceImage), affine_matrix); BufferedImage alignImage = OpenCVUtils.mat2Image(mat); SeetaImageData imageData = new SeetaImageData(alignImage.getWidth(), alignImage.getHeight(), 3); - imageData.data = ImageUtils.getMatrixBGR(alignImage); + imageData.data = BufferedImageUtils.getMatrixBGR(alignImage); return imageData; } } diff --git a/face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java b/face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java index 8e8d2c4..67ccd67 100644 --- a/face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java +++ b/face/src/main/java/cn/smartjavaai/face/utils/FaceUtils.java @@ -3,8 +3,10 @@ package cn.smartjavaai.face.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.Landmark; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.types.DataType; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.Point; import cn.smartjavaai.common.entity.face.FaceAttribute; @@ -13,9 +15,12 @@ import cn.smartjavaai.common.entity.face.HeadPose; import cn.smartjavaai.common.entity.face.LivenessResult; import cn.smartjavaai.common.enums.face.EyeStatus; import cn.smartjavaai.common.enums.face.GenderType; +import cn.smartjavaai.common.utils.BufferedImageUtils; +import cn.smartjavaai.common.utils.Graphics2DUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.enums.face.LivenessStatus; import cn.smartjavaai.face.exception.FaceException; +import cn.smartjavaai.face.model.facedect.mtcnn.MtcnnBatchResult; import com.seeta.sdk.*; import javax.imageio.ImageIO; @@ -158,101 +163,7 @@ public class FaceUtils { return new DetectionResponse(detectionInfoList); } - /** - * 绘制人脸框 - * @param sourceImage - * @param detectionResponse - * @param savePath - * @throws IOException - */ - public static void drawBoundingBoxes(BufferedImage sourceImage, DetectionResponse detectionResponse, String savePath) throws IOException { - if(!ImageUtils.isImageValid(sourceImage)){ - throw new FaceException("图像无效"); - } - if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getDetectionInfoList()) || detectionResponse.getDetectionInfoList().isEmpty()){ - throw new FaceException("无目标数据"); - } - Graphics2D graphics = sourceImage.createGraphics(); - graphics.setColor(Color.RED);// 边框颜色 - graphics.setStroke(new BasicStroke(2)); // 线宽2像素 - graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, - RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 - int stroke = 2; - for(DetectionInfo detectionInfo : detectionResponse.getDetectionInfoList()){ - DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); - graphics.setColor(Color.RED);// 边框颜色 - graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); - String className = "face"; - if (detectionInfo.getScore() > 0){ - int percent = (int) Math.round(detectionInfo.getScore() * 100); - className = "face " + percent + "%"; - } - drawText(graphics, className , rectangle.getX(), rectangle.getY(), stroke, 4); - //绘制人脸关键点 - if(detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getKeyPoints() != null && - !detectionInfo.getFaceInfo().getKeyPoints().isEmpty()){ - drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); - } - } - graphics.dispose(); - ImageIO.write(sourceImage, "png", new File(savePath)); - } - /** - * 绘制人脸框 - * @param sourceImage - * @param detectionResponse - * @throws IOException - */ - public static BufferedImage drawBoundingBoxes(BufferedImage sourceImage, DetectionResponse detectionResponse) throws IOException { - if(!ImageUtils.isImageValid(sourceImage)){ - throw new FaceException("图像无效"); - } - if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getDetectionInfoList()) || detectionResponse.getDetectionInfoList().isEmpty()){ - throw new FaceException("无目标数据"); - } - Graphics2D graphics = sourceImage.createGraphics(); - graphics.setColor(Color.RED);// 边框颜色 - graphics.setStroke(new BasicStroke(2)); // 线宽2像素 - graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, - RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 - int stroke = 2; - for(DetectionInfo detectionInfo : detectionResponse.getDetectionInfoList()){ - DetectionRectangle rectangle = detectionInfo.getDetectionRectangle(); - graphics.setColor(Color.RED);// 边框颜色 - graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); - drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4); - //绘制人脸关键点 - if(detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getKeyPoints() != null && - !detectionInfo.getFaceInfo().getKeyPoints().isEmpty()){ - drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); - } - } - graphics.dispose(); - return sourceImage; - } - - /** - * 绘制文字 - * @param g - * @param text - * @param x - * @param y - * @param stroke - * @param padding - */ - private static void drawText(Graphics2D g, String text, int x, int y, int stroke, int padding) { - FontMetrics metrics = g.getFontMetrics(); - x += stroke / 2; - y += stroke / 2; - int width = metrics.stringWidth(text) + padding * 2 - stroke / 2; - int height = metrics.getHeight() + metrics.getDescent(); - int ascent = metrics.getAscent(); - java.awt.Rectangle background = new java.awt.Rectangle(x, y, width, height); - g.fill(background); - g.setPaint(Color.WHITE); - g.drawString(text, x + padding, y + ascent); - } /** * 修正检测框 @@ -306,28 +217,7 @@ public class FaceUtils { return pointsArray; } - /** - * 子图中人脸关键点坐标 - Coordinates of key points in the image - * - * @param pointFS - * @return - */ - public static double[][] facePoints(SeetaPointF[] pointFS) { - // 图中关键点坐标 - Coordinates of key points in the image - // 1. left_eye_x , left_eye_y - // 2. right_eye_x , right_eye_y - // 3. nose_x , nose_y - // 4. left_mouth_x , left_mouth_y - // 5. right_mouth_x , right_mouth_y - double[][] pointsArray = new double[5][2]; // 保存人脸关键点 - Save facial key points - int i = 0; - for (SeetaPointF point : pointFS) { - pointsArray[i][0] = point.getX(); - pointsArray[i][1] = point.getY(); - i++; - } - return pointsArray; - } + /** * 512x512的目标点 - Target point of 512x512 @@ -386,180 +276,6 @@ public class FaceUtils { return points; } - /** - * bgr转图片 - * @return 图片 - */ - public static BufferedImage toBufferedImage(SeetaImageData seetaImageData) { - int type = BufferedImage.TYPE_3BYTE_BGR; - BufferedImage image = new BufferedImage(seetaImageData.width, seetaImageData.height, type); - image.getRaster().setDataElements(0, 0, seetaImageData.width, seetaImageData.height, seetaImageData.data); - return image; - } - - /** - * 绘制人脸关键点 - * @param g - * @param keyPoints - */ - private static void drawLandmarks(Graphics2D g, List keyPoints) { - g.setColor(new Color(246, 96, 0)); - BasicStroke bStroke = new BasicStroke(4.0F, 0, 0); - g.setStroke(bStroke); - for (Point point : keyPoints){ - g.drawRect((int)point.getX(), (int)point.getY(), 2, 2); - } - } - - - /** - * 将DetectionRectangle转换为SeetaRect - * @param detectionRectangle - * @return - */ - public static SeetaRect convertToSeetaRect(DetectionRectangle detectionRectangle){ - SeetaRect seetaRect = new SeetaRect(); - seetaRect.x = detectionRectangle.getX(); - seetaRect.y = detectionRectangle.getY(); - seetaRect.width = detectionRectangle.getWidth(); - seetaRect.height = detectionRectangle.getHeight(); - return seetaRect; - } - - - /** - * 将PointList转换为SeetaPointF[] - * @param pointList - * @return - */ - public static SeetaPointF[] convertToSeetaPointF(List pointList){ - return pointList.stream() - .map(p -> { - SeetaPointF sp = new SeetaPointF(); - sp.x = p.getX(); - sp.y = p.getY(); - return sp; - }) - .toArray(SeetaPointF[]::new); - } - - /** - * 将SeetaAntiSpoofing.Status转换为LivenessStatus - * @param status - * @return - */ - public static LivenessStatus convertToLivenessStatus(FaceAntiSpoofing.Status status){ - if(status == null){ - return LivenessStatus.UNKNOWN; - } - switch (status) { - case REAL: - return LivenessStatus.LIVE; - case SPOOF: - return LivenessStatus.NON_LIVE; - case FUZZY: - return LivenessStatus.UNKNOWN; - case DETECTING: - return LivenessStatus.DETECTING; - default: - return LivenessStatus.UNKNOWN; // 默认返回未知 - } - } - - /** - * 转为GenderType - * @param gender - * @return - */ - public static GenderType convertToGenderType(GenderPredictor.GENDER gender){ - if(gender == null){ - return GenderType.UNKNOWN; - } - switch (gender) { - case MALE: - return GenderType.MALE; - case FEMALE: - return GenderType.FEMALE; - default: - return GenderType.UNKNOWN; // 默认返回未知 - } - } - - /** - * 转为EyeStatus - * @param eyeState - * @return - */ - public static EyeStatus convertToEyeStatus(EyeStateDetector.EYE_STATE eyeState){ - if(eyeState == null){ - return EyeStatus.UNKNOWN; - } - switch (eyeState) { - case EYE_OPEN: - return EyeStatus.OPEN; - case EYE_CLOSE: - return EyeStatus.CLOSED; - case EYE_RANDOM: - return EyeStatus.NON_EYE_REGION; - default: - return EyeStatus.UNKNOWN; // 默认返回未知 - } - } - - public static DetectionResponse convertToFaceAttributeResponse(SeetaRect[] seetaResult, List seetaPointFSList, List faceAttributeList){ - if(Objects.isNull(seetaResult) || seetaResult.length == 0){ - return null; - } - List detectionInfoList = new ArrayList(); - for(int i = 0; i < seetaResult.length; i++){ - SeetaRect rect = seetaResult[i]; - DetectionRectangle rectangle = new DetectionRectangle(rect.x, rect.y, rect.width, rect.height); - FaceInfo faceInfo = new FaceInfo(); - if(seetaPointFSList != null && seetaPointFSList.size() > 0){ - SeetaPointF[] seetaPointFS = seetaPointFSList.get(i); - List keyPoints = Arrays.stream(seetaPointFS) - .map(p -> new Point(p.x, p.y)) - .collect(Collectors.toList()); - faceInfo.setKeyPoints(keyPoints); - } - if(faceAttributeList != null && faceAttributeList.size() > 0){ - faceInfo.setFaceAttribute(faceAttributeList.get(i)); - } - detectionInfoList.add(new DetectionInfo(rectangle, 0, faceInfo)); - } - return new DetectionResponse(detectionInfoList); - } - - /** - * 转换为FaceDetectedResult - * @param seetaResult - * @return - */ - public static DetectionResponse convertToDetectionResponse(SeetaRect[] seetaResult, List seetaPointFSList, List livenessStatusList){ - if(Objects.isNull(seetaResult) || seetaResult.length == 0){ - return null; - } - DetectionResponse detectionResponse = new DetectionResponse(); - List detectionInfoList = new ArrayList(); - for(int i = 0; i < seetaResult.length; i++){ - SeetaRect rect = seetaResult[i]; - SeetaPointF[] seetaPointFS = seetaPointFSList.get(i); - //过滤置信度 - /*if(config.getConfidenceThreshold() > 0){ - continue; - }*/ - DetectionRectangle rectangle = new DetectionRectangle(rect.x, rect.y, rect.width, rect.height); - List keyPoints = Arrays.stream(seetaPointFS) - .map(p -> new Point(p.x, p.y)) - .collect(Collectors.toList()); - FaceInfo faceInfo = new FaceInfo(keyPoints); - faceInfo.setLivenessStatus(new LivenessResult(livenessStatusList.get(i))); - DetectionInfo detectionInfo = new DetectionInfo(rectangle, 0, faceInfo); - detectionInfoList.add(detectionInfo); - } - detectionResponse.setDetectionInfoList(detectionInfoList); - return detectionResponse; - } /** * 绘制人脸属性 @@ -569,7 +285,7 @@ public class FaceUtils { * @throws IOException */ public static void drawBoxesWithFaceAttribute(BufferedImage sourceImage, DetectionResponse detectionResponse, String savePath) throws IOException { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ throw new FaceException("图像无效"); } if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getDetectionInfoList()) || detectionResponse.getDetectionInfoList().isEmpty()){ @@ -589,7 +305,7 @@ public class FaceUtils { //绘制人脸关键点 if(detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getKeyPoints() != null && !detectionInfo.getFaceInfo().getKeyPoints().isEmpty()){ - drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); + Graphics2DUtils.drawLandmarks(graphics, detectionInfo.getFaceInfo().getKeyPoints()); } // 判断人脸框是否足够大 if (rectangle.getHeight() > 60 && detectionInfo.getFaceInfo() != null && detectionInfo.getFaceInfo().getFaceAttribute() != null) { @@ -638,7 +354,7 @@ public class FaceUtils { lines.add("姿态: P=" + pitch + " Y=" + yaw + " R=" + roll); } if (!lines.isEmpty()) { - drawMultilineTextWithBackground(graphics, lines, rectangle.getX(), rectangle.getY()); // 适当偏移 + Graphics2DUtils.drawMultilineTextWithBackground(graphics, lines, rectangle.getX(), rectangle.getY()); // 适当偏移 } } @@ -647,27 +363,7 @@ public class FaceUtils { ImageIO.write(sourceImage, "png", new File(savePath)); } - private static void drawMultilineTextWithBackground(Graphics2D g, List lines, int x, int y) { - Font font = new Font("SansSerif", Font.PLAIN, 14); - g.setFont(font); - FontMetrics fm = g.getFontMetrics(); - int lineHeight = fm.getHeight(); - int maxWidth = lines.stream().mapToInt(fm::stringWidth).max().orElse(0); - int padding = 4; - int boxWidth = maxWidth + padding * 2; - int boxHeight = lineHeight * lines.size() + padding * 2; - - // 背景矩形 - g.setColor(new Color(0, 0, 0, 128)); - g.fillRoundRect(x, y, boxWidth, boxHeight, 8, 8); - - // 绘制每一行文字 - g.setColor(Color.WHITE); - for (int i = 0; i < lines.size(); i++) { - g.drawString(lines.get(i), x + padding, y + padding + (i + 1) * lineHeight - 4); - } - } /** * 将 Milvus 查询返回的得分转换为 0~1 范围的相似度 @@ -691,6 +387,81 @@ public class FaceUtils { } } + /** + * 裁剪人脸 + * @param image + * @param rectangle + * @return + */ + public static Image cropFace(Image image, DetectionRectangle rectangle){ + return image.getSubImage(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); + } + + /** + * 绘制人脸 + * @param image + * @param rectangle + * @return + */ +// public static Image drawFaceName(Image image, DetectionResponse detectionResponse){ +// +// } + + + + /** + * 将 Mtcnn 批量结果转换为 DJL 的 DetectedObjects + * @param mtcnnBatchResult + * @param imageWidth + * @param imageHeight + * @return + */ + public static DetectedObjects toDetectedObjects(MtcnnBatchResult mtcnnBatchResult, int imageWidth, int imageHeight) { + List classNames = new ArrayList<>(); + List probs = new ArrayList<>(); + List boxes = new ArrayList<>(); + + NDArray boxesND = mtcnnBatchResult.boxes.get(0); + NDArray probsND = mtcnnBatchResult.probs.get(0); + NDArray pointsND = mtcnnBatchResult.points.get(0); + if(pointsND != null){ + pointsND = pointsND.toType(DataType.FLOAT64, false); + } + if(boxesND == null || probsND == null || pointsND == null){ + return new DetectedObjects(classNames, probs, boxes); + } + long numBoxes = boxesND.getShape().get(0); + for (int i = 0; i < numBoxes; i++) { + NDArray box = boxesND.get(i); // [x1, y1, x2, y2] + NDArray prob = probsND.get(i); + NDArray pointND = pointsND.get(i); // shape [5,2] + + float x1 = box.getFloat(0); + float y1 = box.getFloat(1); + float x2 = box.getFloat(2); + float y2 = box.getFloat(3); + + // 转换为 DJL 的 Rectangle,需要归一化到 [0,1] + double x = x1 / imageWidth; + double y = y1 / imageHeight; + double w = (x2 - x1) / imageWidth; + double h = (y2 - y1) / imageHeight; + + List keyPoints = new ArrayList<>(); + double[] flatPoints = pointND.toDoubleArray(); // 一维长度 10 + for (int p = 0; p < 5; p++) { + keyPoints.add(new ai.djl.modality.cv.output.Point(flatPoints[p * 2], flatPoints[p * 2 + 1])); + } + Landmark landmark = + new Landmark(x, y, w, h, keyPoints); +// BoundingBox rect = new ai.djl.modality.cv.output.Rectangle(x, y, w, h); + classNames.add("Face"); // 默认类别是人脸 + probs.add((double) prob.getFloat()); + boxes.add(landmark); + } + return new DetectedObjects(classNames, probs, boxes); + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java b/face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java index a4c4c3b..0a6ffa1 100644 --- a/face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java +++ b/face/src/main/java/cn/smartjavaai/face/utils/Seetaface6Utils.java @@ -1,6 +1,24 @@ package cn.smartjavaai.face.utils; +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.FaceAttribute; +import cn.smartjavaai.common.entity.face.FaceInfo; +import cn.smartjavaai.common.entity.face.LivenessResult; +import cn.smartjavaai.common.enums.face.EyeStatus; +import cn.smartjavaai.common.enums.face.GenderType; +import cn.smartjavaai.common.enums.face.LivenessStatus; import cn.smartjavaai.face.enums.QualityGrade; +import com.seeta.sdk.*; + +import java.awt.image.BufferedImage; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Objects; +import java.util.stream.Collectors; /** * Seetaface6工具类 @@ -10,4 +28,188 @@ import cn.smartjavaai.face.enums.QualityGrade; public class Seetaface6Utils { + /** + * 子图中人脸关键点坐标 - Coordinates of key points in the image + * + * @param pointFS + * @return + */ + public static double[][] facePoints(SeetaPointF[] pointFS) { + // 图中关键点坐标 - Coordinates of key points in the image + // 1. left_eye_x , left_eye_y + // 2. right_eye_x , right_eye_y + // 3. nose_x , nose_y + // 4. left_mouth_x , left_mouth_y + // 5. right_mouth_x , right_mouth_y + double[][] pointsArray = new double[5][2]; // 保存人脸关键点 - Save facial key points + int i = 0; + for (SeetaPointF point : pointFS) { + pointsArray[i][0] = point.getX(); + pointsArray[i][1] = point.getY(); + i++; + } + return pointsArray; + } + + /** + * bgr转图片 + * @return 图片 + */ + public static BufferedImage toBufferedImage(SeetaImageData seetaImageData) { + int type = BufferedImage.TYPE_3BYTE_BGR; + BufferedImage image = new BufferedImage(seetaImageData.width, seetaImageData.height, type); + image.getRaster().setDataElements(0, 0, seetaImageData.width, seetaImageData.height, seetaImageData.data); + return image; + } + + + /** + * 将DetectionRectangle转换为SeetaRect + * @param detectionRectangle + * @return + */ + public static SeetaRect convertToSeetaRect(DetectionRectangle detectionRectangle){ + SeetaRect seetaRect = new SeetaRect(); + seetaRect.x = detectionRectangle.getX(); + seetaRect.y = detectionRectangle.getY(); + seetaRect.width = detectionRectangle.getWidth(); + seetaRect.height = detectionRectangle.getHeight(); + return seetaRect; + } + + /** + * 将PointList转换为SeetaPointF[] + * @param pointList + * @return + */ + public static SeetaPointF[] convertToSeetaPointF(List pointList){ + return pointList.stream() + .map(p -> { + SeetaPointF sp = new SeetaPointF(); + sp.x = p.getX(); + sp.y = p.getY(); + return sp; + }) + .toArray(SeetaPointF[]::new); + } + + /** + * 将SeetaAntiSpoofing.Status转换为LivenessStatus + * @param status + * @return + */ + public static LivenessStatus convertToLivenessStatus(FaceAntiSpoofing.Status status){ + if(status == null){ + return LivenessStatus.UNKNOWN; + } + switch (status) { + case REAL: + return LivenessStatus.LIVE; + case SPOOF: + return LivenessStatus.NON_LIVE; + case FUZZY: + return LivenessStatus.UNKNOWN; + case DETECTING: + return LivenessStatus.DETECTING; + default: + return LivenessStatus.UNKNOWN; // 默认返回未知 + } + } + + + /** + * 转为GenderType + * @param gender + * @return + */ + public static GenderType convertToGenderType(GenderPredictor.GENDER gender){ + if(gender == null){ + return GenderType.UNKNOWN; + } + switch (gender) { + case MALE: + return GenderType.MALE; + case FEMALE: + return GenderType.FEMALE; + default: + return GenderType.UNKNOWN; // 默认返回未知 + } + } + + /** + * 转为EyeStatus + * @param eyeState + * @return + */ + public static EyeStatus convertToEyeStatus(EyeStateDetector.EYE_STATE eyeState){ + if(eyeState == null){ + return EyeStatus.UNKNOWN; + } + switch (eyeState) { + case EYE_OPEN: + return EyeStatus.OPEN; + case EYE_CLOSE: + return EyeStatus.CLOSED; + case EYE_RANDOM: + return EyeStatus.NON_EYE_REGION; + default: + return EyeStatus.UNKNOWN; // 默认返回未知 + } + } + + public static DetectionResponse convertToFaceAttributeResponse(SeetaRect[] seetaResult, List seetaPointFSList, List faceAttributeList){ + if(Objects.isNull(seetaResult) || seetaResult.length == 0){ + return null; + } + List detectionInfoList = new ArrayList(); + for(int i = 0; i < seetaResult.length; i++){ + SeetaRect rect = seetaResult[i]; + DetectionRectangle rectangle = new DetectionRectangle(rect.x, rect.y, rect.width, rect.height); + FaceInfo faceInfo = new FaceInfo(); + if(seetaPointFSList != null && seetaPointFSList.size() > 0){ + SeetaPointF[] seetaPointFS = seetaPointFSList.get(i); + List keyPoints = Arrays.stream(seetaPointFS) + .map(p -> new Point(p.x, p.y)) + .collect(Collectors.toList()); + faceInfo.setKeyPoints(keyPoints); + } + if(faceAttributeList != null && faceAttributeList.size() > 0){ + faceInfo.setFaceAttribute(faceAttributeList.get(i)); + } + detectionInfoList.add(new DetectionInfo(rectangle, 0, faceInfo)); + } + return new DetectionResponse(detectionInfoList); + } + + /** + * 转换为FaceDetectedResult + * @param seetaResult + * @return + */ + public static DetectionResponse convertToDetectionResponse(SeetaRect[] seetaResult, List seetaPointFSList, List livenessStatusList){ + if(Objects.isNull(seetaResult) || seetaResult.length == 0){ + return null; + } + DetectionResponse detectionResponse = new DetectionResponse(); + List detectionInfoList = new ArrayList(); + for(int i = 0; i < seetaResult.length; i++){ + SeetaRect rect = seetaResult[i]; + SeetaPointF[] seetaPointFS = seetaPointFSList.get(i); + //过滤置信度 + /*if(config.getConfidenceThreshold() > 0){ + continue; + }*/ + DetectionRectangle rectangle = new DetectionRectangle(rect.x, rect.y, rect.width, rect.height); + List keyPoints = Arrays.stream(seetaPointFS) + .map(p -> new Point(p.x, p.y)) + .collect(Collectors.toList()); + FaceInfo faceInfo = new FaceInfo(keyPoints); + faceInfo.setLivenessStatus(new LivenessResult(livenessStatusList.get(i))); + DetectionInfo detectionInfo = new DetectionInfo(rectangle, 0, faceInfo); + detectionInfoList.add(detectionInfo); + } + detectionResponse.setDetectionInfoList(detectionInfoList); + return detectionResponse; + } + } diff --git a/face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java b/face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java index 11e2437..79f3d47 100644 --- a/face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java +++ b/face/src/main/java/cn/smartjavaai/face/vector/core/SQLiteClient.java @@ -237,7 +237,7 @@ public class SQLiteClient implements VectorDBClient { private void loadAllFeaturesToMemory() { try { int pageSize = 1000; - int page = 0; + int page = 1; while (true) { List batch = faceDao.findFace(page, pageSize); if (CollectionUtils.isEmpty(batch)) { diff --git a/face/src/test/java/Test.java b/face/src/test/java/Test.java index 5f4c97c..c357d7e 100644 --- a/face/src/test/java/Test.java +++ b/face/src/test/java/Test.java @@ -1,9 +1,12 @@ import ai.djl.Application; +import ai.djl.modality.cv.Image; +import ai.djl.modality.cv.ImageFactory; 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.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.face.config.FaceDetConfig; @@ -20,6 +23,8 @@ import cn.smartjavaai.face.utils.SimilarityUtil; import lombok.extern.slf4j.Slf4j; import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Paths; import java.util.List; import java.util.Map; @@ -85,9 +90,18 @@ public class Test { // } - FaceDetModel faceDetModel = getFaceDetModel(); - R result = faceDetModel.detectAndDraw("/Users/wenjie/Downloads/facetest/00974.png", "/Users/wenjie/Downloads/xx333.png"); - log.info("result:{}", result.isSuccess() + " msg:" + result.getMessage()); +// FaceDetModel faceDetModel = getFaceDetModel(); +// R result = faceDetModel.detectAndDraw("/Users/wenjie/Downloads/facetest/00974.png", "/Users/wenjie/Downloads/xx333.png"); +// log.info("result:{}", result.isSuccess() + " msg:" + result.getMessage()); + + Image ime = ImageFactory.getInstance().fromFile(Paths.get("/Users/wenjie/Downloads/facetest/00974.png")); +// SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); + Image image = SmartImageFactory.getInstance().fromFile(Paths.get("/Users/wenjie/Downloads/facetest/00974.png")); +// image.save(Files.newOutputStream(Paths.get("/Users/wenjie/Downloads/xx333.png")), "png"); +// +// + Image ime2 = ImageFactory.getInstance().fromFile(Paths.get("/Users/wenjie/Downloads/facetest/00974.png")); + image.getSubImage(0, 0, 100, 100); } diff --git a/ocr/pom.xml b/ocr/pom.xml index f2b3263..3662c59 100644 --- a/ocr/pom.xml +++ b/ocr/pom.xml @@ -6,7 +6,7 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 ocr @@ -42,7 +42,7 @@ - 1.0.24 + 1.0.25 ocr SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java b/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java index e265c30..e9397c2 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/entity/OcrInfo.java @@ -1,5 +1,6 @@ package cn.smartjavaai.ocr.entity; +import ai.djl.modality.cv.Image; import lombok.Data; import java.util.ArrayList; @@ -20,7 +21,7 @@ public class OcrInfo { private String fullText; - private String base64Img; + private transient Image drawnImage; public OcrInfo(List> lineList, String fullText) { diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java index 7f79255..3429999 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/factory/OcrModelFactory.java @@ -158,6 +158,7 @@ public class OcrModelFactory { throw new OcrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -179,6 +180,7 @@ public class OcrModelFactory { throw new OcrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -199,6 +201,7 @@ public class OcrModelFactory { throw new OcrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -220,4 +223,61 @@ public class OcrModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + commonDetModelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + commonDetModelMap.clear(); + + commonRecModelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + commonRecModelMap.clear(); + + directionModelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + directionModelMap.clear(); + } + + /** + * 移除缓存的检测模型 + * @param modelEnum + */ + public static void removeDetModelFromCache(CommonDetModelEnum modelEnum) { + commonDetModelMap.remove(modelEnum); + } + + /** + * 移除缓存的识别模型 + * @param modelEnum + */ + public static void removeRecModelFromCache(CommonRecModelEnum modelEnum) { + commonRecModelMap.remove(modelEnum); + } + + /** + * 移除缓存的方向分类模型 + * @param modelEnum + */ + public static void removeDirectionModelFromCache(DirectionModelEnum modelEnum) { + directionModelMap.remove(modelEnum); + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java index 8ffe729..796e763 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/factory/PlateModelFactory.java @@ -4,9 +4,7 @@ import cn.smartjavaai.common.config.Config; import cn.smartjavaai.ocr.config.PlateDetModelConfig; import cn.smartjavaai.ocr.config.PlateRecModelConfig; import cn.smartjavaai.ocr.config.TableStructureConfig; -import cn.smartjavaai.ocr.enums.PlateDetModelEnum; -import cn.smartjavaai.ocr.enums.PlateRecModelEnum; -import cn.smartjavaai.ocr.enums.TableStructureModelEnum; +import cn.smartjavaai.ocr.enums.*; import cn.smartjavaai.ocr.exception.OcrException; import cn.smartjavaai.ocr.model.plate.CRNNPlateRecModel; import cn.smartjavaai.ocr.model.plate.PlateDetModel; @@ -133,6 +131,7 @@ public class PlateModelFactory { throw new OcrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -153,6 +152,7 @@ public class PlateModelFactory { throw new OcrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -165,4 +165,44 @@ public class PlateModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + detModelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + detModelMap.clear(); + + recModelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + recModelMap.clear(); + + } + + /** + * 移除缓存的检测模型 + * @param modelEnum + */ + public static void removeDetModelFromCache(PlateDetModelEnum modelEnum) { + detModelMap.remove(modelEnum); + } + + /** + * 移除缓存的识别模型 + * @param modelEnum + */ + public static void removeRecModelFromCache(PlateRecModelEnum modelEnum) { + recModelMap.remove(modelEnum); + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java b/ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java index 9723173..619d602 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/factory/TableRecModelFactory.java @@ -103,6 +103,7 @@ public class TableRecModelFactory { throw new OcrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -116,4 +117,26 @@ public class TableRecModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + tableStructureModelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + tableStructureModelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(TableStructureModelEnum modelEnum) { + tableStructureModelMap.remove(modelEnum); + } + } diff --git a/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 index 214775e..272b30c 100644 --- a/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 @@ -27,6 +27,8 @@ public interface OcrCommonDetModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + + @Deprecated default List detect(String imagePath) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -37,6 +39,7 @@ public interface OcrCommonDetModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default List detect(BufferedImage image) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -47,6 +50,7 @@ public interface OcrCommonDetModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default List detect(byte[] imageData) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -74,15 +78,26 @@ public interface OcrCommonDetModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default BufferedImage detectAndDraw(BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 检测并绘制结果 + * @param sourceImage + * @return + */ + default Image detectAndDraw(Image sourceImage){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 文本检测(批量) * @param imageList BufferedImage * @return */ + @Deprecated default List> batchDetect(List imageList) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -100,5 +115,9 @@ public interface OcrCommonDetModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java index 369883f..5a2e828 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/detect/OcrCommonDetModelImpl.java @@ -11,13 +11,15 @@ import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelZoo; import ai.djl.repository.zoo.ZooModel; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; -import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.ocr.config.OcrDetModelConfig; import cn.smartjavaai.ocr.entity.OcrBox; import cn.smartjavaai.ocr.exception.OcrException; +import cn.smartjavaai.ocr.factory.OcrModelFactory; import cn.smartjavaai.ocr.model.common.detect.criteria.OcrCommonDetCriterialFactory; import cn.smartjavaai.ocr.utils.OcrUtils; import lombok.extern.slf4j.Slf4j; @@ -81,12 +83,12 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); } catch (IOException e) { throw new OcrException("无效的图片", e); } List ocrBoxList = detect(img); - ((Mat)img.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(img); return ocrBoxList; } @@ -103,16 +105,15 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ throw new OcrException("图像文件不存在"); } try { - Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); List boxList = detect(img); if(Objects.isNull(boxList) || boxList.isEmpty()){ throw new OcrException("未检测到文字"); } - OcrUtils.drawRect((Mat)img.getWrappedImage(), boxList); - Path output = Paths.get(outputPath); - log.debug("Saving to {}", output.toAbsolutePath().toString()); - img.save(Files.newOutputStream(output), "png"); - ((Mat) img.getWrappedImage()).release(); + OcrUtils.drawOcrDetResult(img, boxList, 12); + //4通道保存jpg会有问题 + ImageUtils.save(img, Paths.get(outputPath), "png"); + ImageUtils.releaseOpenCVMat(img); } catch (IOException e) { throw new OcrException(e); } @@ -121,12 +122,12 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ @Override public List detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new OcrException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(image); List ocrBoxList = detect(img); - ((Mat)img.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(img); return ocrBoxList; } @@ -145,29 +146,30 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ @Override public BufferedImage detectAndDraw(BufferedImage sourceImage) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ throw new OcrException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); List ocrBoxList = detect(img); if(Objects.isNull(ocrBoxList) || ocrBoxList.isEmpty()){ throw new OcrException("未检测到文字"); } - OcrUtils.drawRect((Mat)img.getWrappedImage(), ocrBoxList); - try { - ByteArrayOutputStream outputStream = new ByteArrayOutputStream(); - // 调用 save 方法将 Image 写入字节流 - img.save(outputStream, "png"); - // 将字节流转换为 BufferedImage - byte[] imageBytes = outputStream.toByteArray(); - return ImageIO.read(new ByteArrayInputStream(imageBytes)); - } catch (IOException e) { - throw new OcrException("导出图片失败", e); - } finally { - if (img != null){ - ((Mat) img.getWrappedImage()).release(); - } + OcrUtils.drawOcrDetResult(img, ocrBoxList, 12); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(img); + ImageUtils.releaseOpenCVMat(img); + return bufferedImage; + + } + + @Override + public Image detectAndDraw(Image sourceImage) { + List ocrBoxList = detect(sourceImage); + if(Objects.isNull(ocrBoxList) || ocrBoxList.isEmpty()){ + throw new OcrException("未检测到文字"); } + Image img = ImageUtils.copy(sourceImage); + OcrUtils.drawOcrDetResult(img, ocrBoxList, 12); + return img; } @Override @@ -175,13 +177,13 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ List djlImageList = new ArrayList<>(imageList.size()); try { for (BufferedImage bufferedImage : imageList) { - djlImageList.add(ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(bufferedImage))); + djlImageList.add(SmartImageFactory.getInstance().fromBufferedImage(bufferedImage)); } return batchDetectDJLImage(djlImageList); } catch (Exception e) { throw new OcrException(e); } finally { - djlImageList.forEach(image -> ((Mat)image.getWrappedImage()).release()); + djlImageList.forEach(image -> ImageUtils.releaseOpenCVMat(image)); } } @@ -215,6 +217,10 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ } } + + + + @Override public GenericObjectPool> getPool() { return detPredictorPool; @@ -222,6 +228,9 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ @Override public void close() throws Exception { + if (fromFactory) { + OcrModelFactory.removeDetModelFromCache(config.getModelEnum()); + } try { if (detPredictorPool != null) { detPredictorPool.close(); @@ -238,5 +247,15 @@ public class OcrCommonDetModelImpl implements OcrCommonDetModel{ } } + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } + } diff --git a/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 index 23476d0..93be75d 100644 --- a/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 @@ -12,7 +12,8 @@ import ai.djl.ndarray.types.Shape; import ai.djl.translate.Batchifier; import ai.djl.translate.Translator; import ai.djl.translate.TranslatorContext; -import cn.smartjavaai.ocr.opencv.OcrNDArrayUtils; +import cn.smartjavaai.common.utils.DJLCommonUtils; +import cn.smartjavaai.common.utils.OpenCVUtils; import org.opencv.core.*; import org.opencv.imgproc.Imgproc; @@ -98,10 +99,10 @@ public class PPOCRDetTranslator implements Translator { if (this.use_dilation) { Mat mask = new Mat(); //convert from NDArray to Mat - Mat srcMat = OcrNDArrayUtils.uint8NDArrayToMat(segmentation); + Mat srcMat = DJLCommonUtils.uint8NDArrayToMat(segmentation); // size 越小,腐蚀的单位越小,图片越接近原图 // Mat dilation_kernel = Imgproc.getStructuringElement(Imgproc.MORPH_RECT, new Size(2, 2)); - Mat dilation_kernel = OcrNDArrayUtils.uint8ArrayToMat(new byte[][]{{1, 1}, {1, 1}}); + Mat dilation_kernel = OpenCVUtils.uint8ArrayToMat(new byte[][]{{1, 1}, {1, 1}}); /** * 膨胀说明: 图像的一部分区域与指定的核进行卷积, 求核的最`大`值并赋值给指定区域。 膨胀可以理解为图像中`高亮区域`的'领域扩大'。 * 意思是高亮部分会侵蚀不是高亮的部分,使高亮部分越来越多。 @@ -115,7 +116,7 @@ public class PPOCRDetTranslator implements Translator { srcMat.release(); dilation_kernel.release(); } else { - Mat srcMat = OcrNDArrayUtils.uint8NDArrayToMat(segmentation); + Mat srcMat = DJLCommonUtils.uint8NDArrayToMat(segmentation); //destination Matrix Scalar scalar = new Scalar(255); Core.multiply(srcMat, scalar, newMask); @@ -462,20 +463,20 @@ public class PPOCRDetTranslator implements Translator { box.set(new NDIndex(":, 1"), box.get(":, 1").sub(ymin)); //mask - convert from NDArray to Mat - Mat maskMat = OcrNDArrayUtils.uint8NDArrayToMat(mask); + Mat maskMat = DJLCommonUtils.uint8NDArrayToMat(mask); //mask - convert from NDArray to Mat - 4 rows, 2 cols - Mat boxMat = OcrNDArrayUtils.floatNDArrayToMat(box, CvType.CV_32S); + Mat boxMat = DJLCommonUtils.floatNDArrayToMat(box, CvType.CV_32S); // boxMat.reshape(1, new int[]{1, 4, 2}); List pts = new ArrayList<>(); - MatOfPoint matOfPoint = OcrNDArrayUtils.matToMatOfPoint(boxMat); // new MatOfPoint(boxMat); + MatOfPoint matOfPoint = OpenCVUtils.matToMatOfPoint(boxMat); // new MatOfPoint(boxMat); pts.add(matOfPoint); Imgproc.fillPoly(maskMat, pts, new Scalar(1)); NDArray subBitMap = bitmap.get(ymin + ":" + (ymax + 1) + "," + xmin + ":" + (xmax + 1)); - Mat bitMapMat = OcrNDArrayUtils.floatNDArrayToMat(subBitMap); + Mat bitMapMat = DJLCommonUtils.floatNDArrayToMat(subBitMap); Scalar score = Core.mean(bitMapMat, maskMat); float scoreValue = (float) score.val[0]; diff --git a/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 index 7e96f1c..5f348cf 100644 --- a/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 @@ -40,6 +40,7 @@ public interface OcrDirectionModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default List detect(String imagePath) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -50,6 +51,7 @@ public interface OcrDirectionModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default List detect(BufferedImage image) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -60,6 +62,7 @@ public interface OcrDirectionModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default List detect(byte[] imageData) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -103,12 +106,26 @@ public interface OcrDirectionModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default BufferedImage detectAndDraw(BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 检测并绘制结果 + * @param sourceImage + * @return + */ + default Image detectAndDraw(Image sourceImage){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + default GenericObjectPool> getPool() { throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java index 3567460..2b70eb1 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/direction/PPOCRMobileV2ClsModel.java @@ -1,44 +1,35 @@ package cn.smartjavaai.ocr.model.common.direction; -import ai.djl.Device; import ai.djl.MalformedModelException; import ai.djl.engine.Engine; import ai.djl.inference.Predictor; import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.ImageFactory; import ai.djl.ndarray.NDManager; import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelZoo; import ai.djl.repository.zoo.ZooModel; -import ai.djl.training.util.ProgressBar; -import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; -import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.ocr.config.DirectionModelConfig; import cn.smartjavaai.ocr.entity.*; import cn.smartjavaai.ocr.enums.AngleEnum; import cn.smartjavaai.ocr.exception.OcrException; +import cn.smartjavaai.ocr.factory.OcrModelFactory; import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel; import cn.smartjavaai.ocr.model.common.direction.criteria.DirectionCriteriaFactory; -import cn.smartjavaai.ocr.model.common.direction.translator.PpWordRotateTranslator; import cn.smartjavaai.ocr.utils.OcrUtils; import lombok.extern.slf4j.Slf4j; import org.apache.commons.collections.CollectionUtils; import org.apache.commons.lang3.StringUtils; -import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.impl.GenericObjectPool; import org.opencv.core.Mat; -import javax.imageio.ImageIO; import java.awt.image.BufferedImage; -import java.io.ByteArrayInputStream; -import java.io.ByteArrayOutputStream; import java.io.IOException; -import java.nio.file.Files; -import java.nio.file.Path; import java.nio.file.Paths; import java.util.ArrayList; import java.util.Collections; @@ -63,6 +54,8 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { private OcrCommonDetModel textDetModel; + public static final int FONT_SIZE = 45; + @Override public void loadModel(DirectionModelConfig config){ @@ -100,14 +93,12 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); return detect(img); } catch (IOException e) { throw new OcrException("无效的图片", e); }finally { - if(img != null){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -122,7 +113,7 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { if(Objects.isNull(boxeList) || boxeList.isEmpty()){ throw new OcrException("未检测到文本"); } - Mat srcMat = (Mat) image.getWrappedImage(); + Mat srcMat = ImageUtils.toMat(image); return detect(boxeList, srcMat); } @@ -171,7 +162,7 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { // } @Override - public List detect(List boxList,Mat srcMat){ + public List detect(List boxList, Mat srcMat){ if(Objects.isNull(boxList) || boxList.isEmpty()){ throw new OcrException("boxList为空"); } @@ -189,32 +180,30 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); List itemList = detect(img); if(Objects.isNull(itemList) || itemList.isEmpty()){ throw new OcrException("未检测到文字"); } - OcrUtils.drawRectWithText((Mat) img.getWrappedImage(), itemList); - Path output = Paths.get(outputPath); - log.debug("Saving to {}", output.toAbsolutePath().toString()); - img.save(Files.newOutputStream(output), "png"); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(img); + OcrUtils.drawOcrResult(bufferedImage, itemList, FONT_SIZE); + log.debug("Saving to {}", outputPath); + BufferedImageUtils.saveImage(bufferedImage, outputPath); } catch (IOException e) { throw new OcrException(e); } finally { - if (img != null){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @Override public List detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new OcrException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(image); List ocrItemList = detect(img); - ((Mat)img.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(img); return ocrItemList; } @@ -223,39 +212,41 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { if(Objects.isNull(imageData)){ throw new OcrException("图像无效"); } + Image img = null; try { - BufferedImage image = ImageIO.read(new ByteArrayInputStream(imageData)); - return detect(image); + img = SmartImageFactory.getInstance().fromBytes(imageData); } catch (IOException e) { - throw new OcrException("错误的图像", e); + throw new RuntimeException(e); } + List ocrItemList = detect(img); + ImageUtils.releaseOpenCVMat(img); + return ocrItemList; } @Override public BufferedImage detectAndDraw(BufferedImage sourceImage) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ throw new OcrException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); List ocrItemList = detect(img); if(Objects.isNull(ocrItemList) || ocrItemList.isEmpty()){ throw new OcrException("未检测到文字"); } - OcrUtils.drawRectWithText((Mat) img.getWrappedImage(), ocrItemList); - try { - ByteArrayOutputStream outputStream = new ByteArrayOutputStream(); - // 调用 save 方法将 Image 写入字节流 - img.save(outputStream, "png"); - // 将字节流转换为 BufferedImage - byte[] imageBytes = outputStream.toByteArray(); - return ImageIO.read(new ByteArrayInputStream(imageBytes)); - } catch (IOException e) { - throw new OcrException("导出图片失败", e); - } finally { - if (img != null){ - ((Mat) img.getWrappedImage()).release(); - } + BufferedImage drawImage = BufferedImageUtils.copyBufferedImage(sourceImage); + OcrUtils.drawOcrResult(drawImage, ocrItemList, FONT_SIZE); + return drawImage; + } + + @Override + public Image detectAndDraw(Image sourceImage) { + List ocrItemList = detect(sourceImage); + if(Objects.isNull(ocrItemList) || ocrItemList.isEmpty()){ + throw new OcrException("未检测到文字"); } + Image drawImage = ImageUtils.copy(sourceImage); + OcrUtils.drawOcrResult(drawImage, ocrItemList, FONT_SIZE); + return drawImage; } @Override @@ -291,7 +282,7 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { //高宽比 > 1.5 纵向 if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) { //旋转图片90度 - subImg = OcrUtils.rotateImg(manager, subImg); + subImg = ImageUtils.rotateImg(manager, subImg); isRotatedList.add(true); imageList.add(subImg); }else{ @@ -303,6 +294,8 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { } List> result = new ArrayList<>(); List directionInfos = batchDetect(imageList); + //释放 + imageList.forEach(image -> ImageUtils.releaseOpenCVMat(image)); if(CollectionUtils.isEmpty(directionInfos)){ throw new OcrException("方向检测失败"); } @@ -378,6 +371,9 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { @Override public void close() throws Exception { + if (fromFactory) { + OcrModelFactory.removeDirectionModelFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -393,4 +389,14 @@ public class PPOCRMobileV2ClsModel implements OcrDirectionModel { log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java index adef8c6..98e90ed 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModel.java @@ -47,6 +47,7 @@ public interface OcrCommonRecModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default OcrInfo recognize(String imagePath, OcrRecOptions options) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -66,6 +67,7 @@ public interface OcrCommonRecModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default OcrInfo recognize(BufferedImage image, OcrRecOptions options) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -76,6 +78,7 @@ public interface OcrCommonRecModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default OcrInfo recognize(byte[] imageData, OcrRecOptions options) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -90,11 +93,22 @@ public interface OcrCommonRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 识别并绘制结果 + * @param image + * @return + */ + @Deprecated + default OcrInfo recognizeAndDraw(Image image, int fontSize, OcrRecOptions options){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + /** * 识别并绘制结果 * @param sourceImage * @return */ + @Deprecated default BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize, OcrRecOptions options){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -103,6 +117,7 @@ public interface OcrCommonRecModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default String recognizeAndDrawToBase64(byte[] imageData, int fontSize, OcrRecOptions options){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -112,10 +127,12 @@ public interface OcrCommonRecModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default OcrInfo recognizeAndDraw(byte[] imageData, int fontSize, OcrRecOptions options){ throw new UnsupportedOperationException("默认不支持该功能"); } + @Deprecated default List batchRecognize(List imageList, OcrRecOptions options) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -128,4 +145,8 @@ public interface OcrCommonRecModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java index e3e704e..7d9eeb4 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/common/recognize/OcrCommonRecModelImpl.java @@ -4,7 +4,6 @@ import ai.djl.MalformedModelException; import ai.djl.engine.Engine; import ai.djl.inference.Predictor; import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.ImageFactory; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDManager; import ai.djl.repository.zoo.Criteria; @@ -12,14 +11,16 @@ import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelZoo; import ai.djl.repository.zoo.ZooModel; import cn.hutool.core.img.ImgUtil; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; -import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.ocr.config.OcrRecModelConfig; import cn.smartjavaai.ocr.config.OcrRecOptions; import cn.smartjavaai.ocr.entity.*; import cn.smartjavaai.ocr.exception.OcrException; +import cn.smartjavaai.ocr.factory.OcrModelFactory; import cn.smartjavaai.ocr.model.common.detect.OcrCommonDetModel; import cn.smartjavaai.ocr.model.common.direction.OcrDirectionModel; import cn.smartjavaai.ocr.model.common.recognize.criteria.OcrCommonRecCriterialFactory; @@ -33,7 +34,6 @@ import org.opencv.core.Mat; import javax.imageio.ImageIO; import java.awt.image.BufferedImage; import java.io.ByteArrayInputStream; -import java.io.ByteArrayOutputStream; import java.io.IOException; import java.nio.file.Paths; import java.util.*; @@ -95,14 +95,12 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); return recognize(img, options); } catch (IOException e) { throw new OcrException("无效的图片", e); } finally { - if (img != null) { - ((Mat) img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -138,7 +136,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { //高宽比 > 1.5 if (subImg.getHeight() * 1.0 / subImg.getWidth() > 1.5) { //旋转图片90度 - subImg = OcrUtils.rotateImg(manager, subImg); + subImg = ImageUtils.rotateImg(manager, subImg); //ImageUtils.saveImage(subImg, i + "rotate.png", "build/output"); } imageList.add(subImg); @@ -236,30 +234,31 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { if (!FileUtils.isFileExists(imagePath)) { throw new OcrException("图像文件不存在"); } + Image img = null; try { - Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); OcrInfo ocrInfo = recognize(img, options); if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) { throw new OcrException("未检测到文字"); } - Mat wrappedImage = (Mat) img.getWrappedImage(); - BufferedImage bufferedImage = OpenCVUtils.mat2Image(wrappedImage); - OcrUtils.drawRectWithText(bufferedImage, ocrInfo, fontSize); - ImageUtils.saveImage(bufferedImage, outputPath); - wrappedImage.release(); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(img); + OcrUtils.drawOcrResult(bufferedImage, ocrInfo, fontSize); + BufferedImageUtils.saveImage(bufferedImage, outputPath); } catch (IOException e) { throw new OcrException(e); + }finally { + ImageUtils.releaseOpenCVMat(img); } } @Override public OcrInfo recognize(BufferedImage image, OcrRecOptions options) { - if (!ImageUtils.isImageValid(image)) { + if (!BufferedImageUtils.isImageValid(image)) { throw new OcrException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(image); OcrInfo ocrInfo = recognize(img, options); - ((Mat) img.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(img); return ocrInfo; } @@ -278,15 +277,15 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { @Override public BufferedImage recognizeAndDraw(BufferedImage sourceImage, int fontSize, OcrRecOptions options) { - if (!ImageUtils.isImageValid(sourceImage)) { + if (!BufferedImageUtils.isImageValid(sourceImage)) { throw new OcrException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); OcrInfo ocrInfo = recognize(img, options); if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) { throw new OcrException("未检测到文字"); } - OcrUtils.drawRectWithText(sourceImage, ocrInfo, fontSize); + OcrUtils.drawOcrResult(sourceImage, ocrInfo, fontSize); return sourceImage; } @@ -301,7 +300,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { } try { BufferedImage sourceImage = ImageIO.read(new ByteArrayInputStream(imageData)); - OcrUtils.drawRectWithText(sourceImage, ocrInfo, fontSize); + OcrUtils.drawOcrResult(sourceImage, ocrInfo, fontSize); return ImgUtil.toBase64(sourceImage, "png"); } catch (IOException e) { throw new OcrException("导出图片失败", e); @@ -313,18 +312,21 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { if (Objects.isNull(imageData)) { throw new OcrException("图像无效"); } - OcrInfo ocrInfo = recognize(imageData, options); + Image img = null; + try { + img = SmartImageFactory.getInstance().fromBytes(imageData); + } catch (IOException e) { + throw new RuntimeException(e); + } + OcrInfo ocrInfo = recognize(img, options); if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) { throw new OcrException("未检测到文字"); } - try { - BufferedImage sourceImage = ImageIO.read(new ByteArrayInputStream(imageData)); - OcrUtils.drawRectWithText(sourceImage, ocrInfo, fontSize); - ocrInfo.setBase64Img(ImgUtil.toBase64(sourceImage, "png")); - return ocrInfo; - } catch (IOException e) { - throw new OcrException("导出图片失败", e); - } + //opencv中文乱码,使用BufferedImage + BufferedImage sourceImage = ImageUtils.toBufferedImage(img); + OcrUtils.drawOcrResult(sourceImage, ocrInfo, fontSize); + ocrInfo.setDrawnImage(SmartImageFactory.getInstance().fromBufferedImage(sourceImage)); + return ocrInfo; } @Override @@ -332,13 +334,13 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { List djlImageList = new ArrayList<>(imageList.size()); try { for (BufferedImage bufferedImage : imageList) { - djlImageList.add(ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(bufferedImage))); + djlImageList.add(SmartImageFactory.getInstance().fromBufferedImage(bufferedImage)); } return batchRecognizeDJLImage(djlImageList, options); } catch (Exception e) { throw new OcrException(e); } finally { - djlImageList.forEach(image -> ((Mat) image.getWrappedImage()).release()); + djlImageList.forEach(image -> ImageUtils.releaseOpenCVMat(image)); } } @@ -370,7 +372,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { throw new OcrException("请配置方向模型"); } List matList = imageList.stream() - .map(image -> (Mat) image.getWrappedImage()) + .map(image -> ImageUtils.toMat(image)) .collect(Collectors.toList()); List> ocrItemList = directionModel.batchDetect(boxeList, matList); if (CollectionUtils.isEmpty(ocrItemList) || ocrItemList.size() != imageList.size()) { @@ -378,7 +380,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { } allImageAlignList = new ArrayList(); for (int i = 0; i < ocrItemList.size(); i++) { - Mat srcMat = (Mat) imageList.get(i).getWrappedImage(); + Mat srcMat = ImageUtils.toMat(imageList.get(i)); List imageAlignList = batchAlignWithDirection(ocrItemList.get(i), srcMat, manager); // for(int j = 0; j < imageAlignList.size(); j++){ // ImageUtils.saveImage(imageAlignList.get(j),"dir-"+i+"-"+j+".png","/Users/xxx/Downloads/testing33"); @@ -387,10 +389,10 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { } } else { for (int i = 0; i < boxeList.size(); i++) { - Mat srcMat = (Mat) imageList.get(i).getWrappedImage(); + Mat srcMat = ImageUtils.toMat(imageList.get(i)); List imageAlignList = batchAlign(boxeList.get(i), srcMat, manager); // for(int j = 0; j < imageAlignList.size(); j++){ -// ImageUtils.saveImage(imageAlignList.get(j),i+"-"+j+".png","/Users/xxx/Downloads/testing33"); +// ImageUtils.saveImage(imageAlignList.get(j),i+"-"+j+".png","/Users/wenjie/Downloads/testing33"); // } allImageAlignList.addAll(imageAlignList); } @@ -435,7 +437,7 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { try { predictor = recPredictorPool.borrowObject(); List textList = predictor.batchPredict(imageAlignList); - imageAlignList.forEach(subImg -> ((Mat) subImg.getWrappedImage()).release()); + imageAlignList.forEach(subImg -> ImageUtils.releaseOpenCVMat(subImg)); return textList; } catch (Exception e) { throw new OcrException("OCR检测错误", e); @@ -455,6 +457,19 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { } } + + @Override + public OcrInfo recognizeAndDraw(Image image, int fontSize, OcrRecOptions options) { + OcrInfo ocrInfo = recognize(image, options); + if (Objects.isNull(ocrInfo) || Objects.isNull(ocrInfo.getLineList()) || ocrInfo.getLineList().isEmpty()) { + throw new OcrException("未检测到文字"); + } + BufferedImage sourceImage = ImageUtils.toBufferedImage(image); + OcrUtils.drawOcrResult(sourceImage, ocrInfo, fontSize); + ocrInfo.setDrawnImage(SmartImageFactory.getInstance().fromBufferedImage(sourceImage)); + return ocrInfo; + } + @Override public void setTextDetModel(OcrCommonDetModel detModel) { this.textDetModel = detModel; @@ -482,6 +497,9 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { @Override public void close() throws Exception { + if (fromFactory) { + OcrModelFactory.removeRecModelFromCache(config.getRecModelEnum()); + } try { if (recPredictorPool != null) { recPredictorPool.close(); @@ -497,4 +515,14 @@ public class OcrCommonRecModelImpl implements OcrCommonRecModel { log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java index ba95975..a97cada 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/CRNNPlateRecModel.java @@ -4,7 +4,6 @@ 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.repository.zoo.Criteria; @@ -13,19 +12,18 @@ import ai.djl.repository.zoo.ModelZoo; import ai.djl.repository.zoo.ZooModel; import cn.hutool.core.lang.UUID; import cn.hutool.core.lang.generator.UUIDGenerator; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.R; 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.common.utils.*; import cn.smartjavaai.ocr.config.PlateDetModelConfig; import cn.smartjavaai.ocr.config.PlateRecModelConfig; import cn.smartjavaai.ocr.entity.PlateInfo; import cn.smartjavaai.ocr.entity.PlateResult; import cn.smartjavaai.ocr.enums.PlateType; import cn.smartjavaai.ocr.exception.OcrException; +import cn.smartjavaai.ocr.factory.PlateModelFactory; import cn.smartjavaai.ocr.model.plate.criteria.PlateDetCriterialFactory; import cn.smartjavaai.ocr.model.plate.criteria.PlateRecCriterialFactory; import cn.smartjavaai.ocr.utils.OcrUtils; @@ -97,13 +95,13 @@ public class CRNNPlateRecModel implements PlateRecModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); R> plateResult = recognize(img); return plateResult; } catch (IOException e) { throw new OcrException("无效的图片", e); } finally { - ((Mat)img.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(img); } } @@ -118,12 +116,12 @@ public class CRNNPlateRecModel implements PlateRecModel{ @Override public R> recognize(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + Image img = SmartImageFactory.getInstance().fromBufferedImage(image); R> plateResult = recognize(img); - ((Mat)img.getWrappedImage()).release(); + ImageUtils.releaseOpenCVMat(img); return plateResult; } @@ -140,7 +138,7 @@ public class CRNNPlateRecModel implements PlateRecModel{ if(Objects.isNull(config.getPlateDetModel())){ return R.fail(R.Status.PARAM_ERROR.getCode(), "未指定车牌检测模型"); } - DetectedObjects detectedObjects = config.getPlateDetModel().detect(image); + DetectedObjects detectedObjects = config.getPlateDetModel().detectCore(image); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ return R.fail(R.Status.NO_OBJECT_DETECTED); } @@ -151,23 +149,26 @@ public class CRNNPlateRecModel implements PlateRecModel{ for (PlateInfo plateInfo : plateInfoList){ DetectionRectangle detectionRectangle = plateInfo.getDetectionRectangle(); // Image subImage = image.getSubImage(detectionRectangle.getX(), detectionRectangle.getY(), detectionRectangle.getWidth(), detectionRectangle.getHeight()); + Mat imageMat = ImageUtils.toMat(image); //透视变换 - Image subImage = OcrUtils.transformAndCrop((Mat)image.getWrappedImage(), plateInfo.getBox()); + Mat subMat = OcrUtils.transformAndCropToMat(imageMat, plateInfo.getBox()); //双层车牌 if(plateInfo.getPlateType() == PlateType.DOUBLE){ - Mat mergeImage = getSplitMerge((Mat)subImage.getWrappedImage()); - subImage = ImageFactory.getInstance().fromImage(mergeImage); + subMat = getSplitMerge(subMat); } + Image subImage = SmartImageFactory.getInstance().fromMat(subMat); PlateResult plateResult = predictor.predict(subImage); if(Objects.nonNull(plateResult)){ plateInfo.setPlateNumber(plateResult.getPlateNo()); plateInfo.setPlateColor(plateResult.getPlateColor()); } + ImageUtils.releaseOpenCVMat(subImage); } return R.ok(plateInfoList); } catch (Exception e) { throw new OcrException("车牌识别错误", e); }finally { + if (predictor != null) { try { recPredictorPool.returnObject(predictor); //归还 @@ -247,14 +248,12 @@ public class CRNNPlateRecModel implements PlateRecModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromInputStream(inputStream); + img = SmartImageFactory.getInstance().fromInputStream(inputStream); return recognize(img); } catch (IOException e) { throw new OcrException("无效图片输入流", e); } finally { - if (img != null){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -265,7 +264,7 @@ public class CRNNPlateRecModel implements PlateRecModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); R> plateResult = recognize(img); if(!plateResult.isSuccess()){ return R.fail(plateResult.getCode(), plateResult.getMessage()); @@ -273,22 +272,20 @@ public class CRNNPlateRecModel implements PlateRecModel{ if(CollectionUtils.isEmpty(plateResult.getData())){ return R.fail(R.Status.NO_OBJECT_DETECTED); } - BufferedImage bufferedImage = OpenCVUtils.mat2Image((Mat)img.getWrappedImage()); + BufferedImage bufferedImage = ImageUtils.toBufferedImage(img); OcrUtils.drawPlateInfo(bufferedImage, plateResult.getData()); ImageIO.write(bufferedImage, "png", new File(outputPath)); return R.ok(); } catch (IOException e) { throw new OcrException(e); } finally { - if (img != null){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @Override public R recognizeAndDraw(BufferedImage sourceImage) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ return R.fail(R.Status.INVALID_IMAGE); } try { @@ -306,6 +303,25 @@ public class CRNNPlateRecModel implements PlateRecModel{ } } + @Override + public R recognizeAndDraw(Image image) { + try { + R> plateResult = recognize(image); + if(!plateResult.isSuccess()){ + return R.fail(plateResult.getCode(), plateResult.getMessage()); + } + if(CollectionUtils.isEmpty(plateResult.getData())){ + return R.fail(R.Status.NO_OBJECT_DETECTED); + } + //opencv中文乱码,使用BufferedImage + BufferedImage sourceImage = ImageUtils.toBufferedImage(image); + OcrUtils.drawPlateInfo(sourceImage, plateResult.getData()); + Image drawImage = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); + return R.ok(drawImage); + } catch (Exception e) { + throw new OcrException("导出图片失败", e); + } + } @Override public GenericObjectPool> getPool() { @@ -314,6 +330,9 @@ public class CRNNPlateRecModel implements PlateRecModel{ @Override public void close() throws Exception { + if (fromFactory) { + PlateModelFactory.removeRecModelFromCache(config.getModelEnum()); + } try { if (recPredictorPool != null) { recPredictorPool.close(); @@ -329,4 +348,14 @@ public class CRNNPlateRecModel implements PlateRecModel{ log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java index 2f31059..0adc57f 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateDetModel.java @@ -32,6 +32,7 @@ public interface PlateDetModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R> detect(String imagePath) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -42,6 +43,7 @@ public interface PlateDetModel extends AutoCloseable{ * @param inputStream * @return */ + @Deprecated default R> detect(InputStream inputStream) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -52,6 +54,7 @@ public interface PlateDetModel extends AutoCloseable{ * @param base64Image * @return */ + @Deprecated default R> detectBase64(String base64Image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -62,6 +65,7 @@ public interface PlateDetModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default R> detect(BufferedImage image) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -72,6 +76,7 @@ public interface PlateDetModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default R> detect(byte[] imageData) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -81,7 +86,17 @@ public interface PlateDetModel extends AutoCloseable{ * @param image DJL Image * @return */ - default DetectedObjects detect(Image image){ + default DetectedObjects detectCore(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + + /** + * 车牌检测 + * @param image DJL Image + * @return + */ + default R> detect(Image image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -100,16 +115,28 @@ public interface PlateDetModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default R detectAndDraw(BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } + /** + * 检测并绘制结果 + * @param image + * @return + */ + default Image detectAndDraw(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + default GenericObjectPool> getPool(){ throw new UnsupportedOperationException("默认不支持该功能"); } - + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java index 68b27d0..98e8a6a 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/PlateRecModel.java @@ -31,6 +31,7 @@ public interface PlateRecModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R> recognize(String imagePath) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -41,6 +42,7 @@ public interface PlateRecModel extends AutoCloseable{ * @param inputStream * @return */ + @Deprecated default R> recognize(InputStream inputStream) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -51,6 +53,7 @@ public interface PlateRecModel extends AutoCloseable{ * @param base64Image * @return */ + @Deprecated default R> recognizeBase64(String base64Image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -61,6 +64,7 @@ public interface PlateRecModel extends AutoCloseable{ * @param image BufferedImage * @return */ + @Deprecated default R> recognize(BufferedImage image) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -71,6 +75,7 @@ public interface PlateRecModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default R> recognize(byte[] imageData) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -107,12 +112,21 @@ public interface PlateRecModel extends AutoCloseable{ * @param sourceImage * @return */ + @Deprecated default R recognizeAndDraw(BufferedImage sourceImage){ throw new UnsupportedOperationException("默认不支持该功能"); } + default R recognizeAndDraw(Image image){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + default GenericObjectPool> getPool() { throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java index be221e6..9e215d7 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/plate/Yolov5PlateDetModel.java @@ -4,7 +4,6 @@ import ai.djl.MalformedModelException; import ai.djl.engine.Engine; import ai.djl.inference.Predictor; import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.ImageFactory; import ai.djl.modality.cv.output.DetectedObjects; import ai.djl.ndarray.NDList; import ai.djl.ndarray.NDManager; @@ -12,17 +11,17 @@ import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelZoo; import ai.djl.repository.zoo.ZooModel; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.R; 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.common.utils.*; import cn.smartjavaai.ocr.config.OcrDetModelConfig; import cn.smartjavaai.ocr.config.PlateDetModelConfig; import cn.smartjavaai.ocr.entity.OcrBox; import cn.smartjavaai.ocr.entity.PlateInfo; import cn.smartjavaai.ocr.exception.OcrException; +import cn.smartjavaai.ocr.factory.OcrModelFactory; +import cn.smartjavaai.ocr.factory.PlateModelFactory; import cn.smartjavaai.ocr.model.common.detect.criteria.OcrCommonDetCriterialFactory; import cn.smartjavaai.ocr.model.plate.criteria.PlateDetCriterialFactory; import cn.smartjavaai.ocr.utils.OcrUtils; @@ -90,17 +89,15 @@ public class Yolov5PlateDetModel implements PlateDetModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + R> plateInfoList = detect(img); + return plateInfoList; } catch (IOException e) { throw new OcrException("无效的图片", e); + } finally { + ImageUtils.releaseOpenCVMat(img); } - DetectedObjects detectedObjects = detect(img); - if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ - return R.fail(R.Status.NO_OBJECT_DETECTED); - } - List plateInfoList = OcrUtils.convertToPlateInfo(detectedObjects, img); - ((Mat)img.getWrappedImage()).release(); - return R.ok(plateInfoList); + } @Override @@ -114,17 +111,13 @@ public class Yolov5PlateDetModel implements PlateDetModel{ @Override public R> detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); - DetectedObjects detectedObjects = detect(img); - if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ - return R.fail(R.Status.NO_OBJECT_DETECTED); - } - List plateInfoList = OcrUtils.convertToPlateInfo(detectedObjects, img); - ((Mat)img.getWrappedImage()).release(); - return R.ok(plateInfoList); + Image img = SmartImageFactory.getInstance().fromBufferedImage(image); + R> plateInfoList = detect(img); + ImageUtils.releaseOpenCVMat(img); + return plateInfoList; } @Override @@ -136,7 +129,7 @@ public class Yolov5PlateDetModel implements PlateDetModel{ } @Override - public DetectedObjects detect(Image image) { + public DetectedObjects detectCore(Image image) { Predictor predictor = null; try { predictor = detPredictorPool.borrowObject(); @@ -164,14 +157,15 @@ public class Yolov5PlateDetModel implements PlateDetModel{ if(Objects.isNull(inputStream)){ return R.fail(R.Status.INVALID_IMAGE); } + Image img = null; try { - Image img = ImageFactory.getInstance().fromInputStream(inputStream); - DetectedObjects detection = detect(img); - List plateInfoList = OcrUtils.convertToPlateInfo(detection, img); - ((Mat)img.getWrappedImage()).release(); - return R.ok(plateInfoList); + img = SmartImageFactory.getInstance().fromInputStream(inputStream); + R> plateInfoList = detect(img); + return plateInfoList; } catch (IOException e) { throw new OcrException("无效图片输入流", e); + } finally { + ImageUtils.releaseOpenCVMat(img); } } @@ -181,8 +175,8 @@ public class Yolov5PlateDetModel implements PlateDetModel{ return R.fail(R.Status.FILE_NOT_FOUND); } try { - Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); - DetectedObjects detectedObjects = detect(img); + Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ return R.fail(R.Status.NO_FACE_DETECTED); } @@ -198,11 +192,11 @@ public class Yolov5PlateDetModel implements PlateDetModel{ @Override public R detectAndDraw(BufferedImage sourceImage) { - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ return R.fail(R.Status.INVALID_IMAGE); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); - DetectedObjects detectedObjects = detect(img); + Image img = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); + DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ return R.fail(R.Status.NO_FACE_DETECTED); } @@ -219,6 +213,27 @@ public class Yolov5PlateDetModel implements PlateDetModel{ } } + @Override + public R> detect(Image image) { + DetectedObjects detectedObjects = detectCore(image); + if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ + return R.fail(R.Status.NO_OBJECT_DETECTED); + } + List plateInfoList = OcrUtils.convertToPlateInfo(detectedObjects, image); + return R.ok(plateInfoList); + } + + @Override + public Image detectAndDraw(Image image) { + DetectedObjects detectedObjects = detectCore(image); + if (Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ + throw new OcrException("未检测到车牌"); + } + Image img = ImageUtils.copy(image); + img.drawBoundingBoxes(detectedObjects); + return img; + } + @Override public GenericObjectPool> getPool() { return detPredictorPool; @@ -226,6 +241,9 @@ public class Yolov5PlateDetModel implements PlateDetModel{ @Override public void close() throws Exception { + if (fromFactory) { + PlateModelFactory.removeDetModelFromCache(config.getModelEnum()); + } try { if (detPredictorPool != null) { detPredictorPool.close(); @@ -241,4 +259,14 @@ public class Yolov5PlateDetModel implements PlateDetModel{ log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java index 55c2d4e..470a018 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/CommonTableStructureModel.java @@ -11,8 +11,10 @@ import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelZoo; import ai.djl.repository.zoo.ZooModel; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.pool.PredictorFactory; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.OpenCVUtils; @@ -22,6 +24,8 @@ import cn.smartjavaai.ocr.entity.OcrBox; import cn.smartjavaai.ocr.entity.OcrItem; import cn.smartjavaai.ocr.entity.TableStructureResult; import cn.smartjavaai.ocr.exception.OcrException; +import cn.smartjavaai.ocr.factory.PlateModelFactory; +import cn.smartjavaai.ocr.factory.TableRecModelFactory; import cn.smartjavaai.ocr.model.table.criteria.StructureCriteriaFactory; import cn.smartjavaai.ocr.utils.OcrUtils; import lombok.extern.slf4j.Slf4j; @@ -49,8 +53,11 @@ public class CommonTableStructureModel implements TableStructureModel{ private GenericObjectPool> predictorPool; + private TableStructureConfig config; + @Override public void loadModel(TableStructureConfig config) { + this.config = config; if(StringUtils.isBlank(config.getModelPath())){ throw new OcrException("modelPath is null"); } @@ -74,19 +81,17 @@ public class CommonTableStructureModel implements TableStructureModel{ @Override public R detect(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } Image img = null; try { - img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + img = SmartImageFactory.getInstance().fromBufferedImage(image); return detect(img); } catch (Exception e) { throw new OcrException(e); } finally { - if(Objects.nonNull(img)){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -97,14 +102,12 @@ public class CommonTableStructureModel implements TableStructureModel{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); return detect(img); } catch (IOException e) { throw new OcrException("无效的图片", e); } finally { - if (Objects.nonNull(img)){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -153,6 +156,9 @@ public class CommonTableStructureModel implements TableStructureModel{ @Override public void close() throws Exception { + if (fromFactory) { + TableRecModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -168,4 +174,14 @@ public class CommonTableStructureModel implements TableStructureModel{ log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java index e064cb6..bb252b9 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableRecognizer.java @@ -1,18 +1,12 @@ package cn.smartjavaai.ocr.model.table; import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.ImageFactory; -import ai.djl.modality.cv.output.BoundingBox; -import ai.djl.modality.cv.output.DetectedObjects; -import ai.djl.modality.cv.output.Rectangle; -import ai.djl.translate.TranslateException; -import ai.djl.util.JsonUtils; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; -import cn.smartjavaai.common.utils.OpenCVUtils; -import cn.smartjavaai.ocr.config.OcrRecModelConfig; import cn.smartjavaai.ocr.config.OcrRecOptions; import cn.smartjavaai.ocr.entity.OcrBox; import cn.smartjavaai.ocr.entity.OcrInfo; @@ -28,7 +22,6 @@ import org.apache.commons.collections.CollectionUtils; import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.tuple.Pair; import org.apache.poi.hssf.usermodel.HSSFWorkbook; -import org.opencv.core.Mat; import javax.imageio.ImageIO; import java.awt.*; @@ -107,19 +100,17 @@ public class TableRecognizer { * @return */ public R recognize(BufferedImage image) { - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ return R.fail(R.Status.INVALID_IMAGE); } Image img = null; try { - img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); + img = SmartImageFactory.getInstance().fromBufferedImage(image); return recognize(img); } catch (Exception e) { throw new OcrException(e); } finally { - if(Objects.nonNull(img)){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -134,14 +125,12 @@ public class TableRecognizer { } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); return recognize(img); } catch (IOException e) { throw new OcrException("无效的图片", e); } finally { - if (Objects.nonNull(img)){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -176,9 +165,13 @@ public class TableRecognizer { for (int i = 0; i < tableStructureResult.getOcrItemList().size(); i++){ OcrItem item = tableStructureResult.getOcrItemList().get(i); DetectionRectangle detectionRectangle = item.getOcrBox().toDetectionRectangle(); - ImageUtils.drawImageRectWithText(image, detectionRectangle, i + "", Color.RED); + BufferedImageUtils.drawRectAndText(image, detectionRectangle, i + "", Color.RED); + } + try { + BufferedImageUtils.saveImage(image, savePath); + } catch (IOException e) { + throw new OcrException(e); } - ImageUtils.saveImage(image, savePath); } @@ -195,7 +188,25 @@ public class TableRecognizer { for (int i = 0; i < tableStructureResult.getOcrItemList().size(); i++){ OcrItem item = tableStructureResult.getOcrItemList().get(i); DetectionRectangle detectionRectangle = item.getOcrBox().toDetectionRectangle(); - ImageUtils.drawImageRectWithText(image, detectionRectangle, i + "", Color.RED); + BufferedImageUtils.drawRectAndText(image, detectionRectangle, i + "", Color.RED); + } + return image; + } + + /** + * 绘制表格 + * @param tableStructureResult + * @param image + * @return + */ + public Image drawTable(TableStructureResult tableStructureResult, Image image){ + if(Objects.isNull(tableStructureResult) || CollectionUtils.isEmpty(tableStructureResult.getTableTagList())){ + throw new OcrException("表格结构为空"); + } + for (int i = 0; i < tableStructureResult.getOcrItemList().size(); i++){ + OcrItem item = tableStructureResult.getOcrItemList().get(i); + DetectionRectangle detectionRectangle = item.getOcrBox().toDetectionRectangle(); + ImageUtils.drawRectAndText(image, detectionRectangle, i + ""); } return image; } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java index 098d593..143c5e9 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/model/table/TableStructureModel.java @@ -31,6 +31,7 @@ public interface TableStructureModel extends AutoCloseable{ * @param image * @return */ + @Deprecated default R detect(BufferedImage image){ throw new UnsupportedOperationException("默认不支持该功能"); } @@ -40,6 +41,7 @@ public interface TableStructureModel extends AutoCloseable{ * @param imagePath 图片路径 * @return */ + @Deprecated default R detect(String imagePath) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -49,6 +51,7 @@ public interface TableStructureModel extends AutoCloseable{ * @param imageData 图片字节数组 * @return */ + @Deprecated default R detect(byte[] imageData) { throw new UnsupportedOperationException("默认不支持该功能"); } @@ -66,4 +69,10 @@ public interface TableStructureModel extends AutoCloseable{ default GenericObjectPool> getPool() { throw new UnsupportedOperationException("默认不支持该功能"); } + + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + + } diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java b/ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java deleted file mode 100644 index 0673f9d..0000000 --- a/ocr/src/main/java/cn/smartjavaai/ocr/opencv/OcrNDArrayUtils.java +++ /dev/null @@ -1,228 +0,0 @@ -package cn.smartjavaai.ocr.opencv; - -import ai.djl.ndarray.NDArray; -import org.opencv.core.CvType; -import org.opencv.core.Mat; -import org.opencv.core.MatOfPoint; -import org.opencv.core.Point; - -import java.util.ArrayList; -import java.util.List; -/** - * NDArray Utils - * - */ -public class OcrNDArrayUtils { - /** - * Mat To MatOfPoint - * @param mat - * @return - */ - public static MatOfPoint matToMatOfPoint(Mat mat) { - int rows = mat.rows(); - MatOfPoint matOfPoint = new MatOfPoint(); - - List list = new ArrayList<>(); - for (int i = 0; i < rows; i++) { - Point point = new Point((float) mat.get(i, 0)[0], (float) mat.get(i, 1)[0]); - list.add(point); - } - matOfPoint.fromList(list); - - return matOfPoint; - } - - /** - * float NDArray To float[][] Array - * @param ndArray - * @return - */ - public static float[][] floatNDArrayToArray(NDArray ndArray) { - int rows = (int) (ndArray.getShape().get(0)); - int cols = (int) (ndArray.getShape().get(1)); - float[][] arr = new float[rows][cols]; - - float[] arrs = ndArray.toFloatArray(); - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - arr[i][j] = arrs[i * cols + j]; - } - } - return arr; - } - - /** - * Mat To double[][] Array - * @param mat - * @return - */ - public static double[][] matToDoubleArray(Mat mat) { - int rows = mat.rows(); - int cols = mat.cols(); - - double[][] doubles = new double[rows][cols]; - - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - doubles[i][j] = mat.get(i, j)[0]; - } - } - - return doubles; - } - - /** - * Mat To float[][] Array - * @param mat - * @return - */ - public static float[][] matToFloatArray(Mat mat) { - int rows = mat.rows(); - int cols = mat.cols(); - - float[][] floats = new float[rows][cols]; - - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - floats[i][j] = (float) mat.get(i, j)[0]; - } - } - - return floats; - } - - /** - * Mat To byte[][] Array - * @param mat - * @return - */ - public static byte[][] matToUint8Array(Mat mat) { - int rows = mat.rows(); - int cols = mat.cols(); - - byte[][] bytes = new byte[rows][cols]; - - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - bytes[i][j] = (byte) mat.get(i, j)[0]; - } - } - - return bytes; - } - - /** - * float NDArray To float[][] Array - * @param ndArray - * @param cvType - * @return - */ - public static Mat floatNDArrayToMat(NDArray ndArray, int cvType) { - int rows = (int) (ndArray.getShape().get(0)); - int cols = (int) (ndArray.getShape().get(1)); - Mat mat = new Mat(rows, cols, cvType); - - float[] arrs = ndArray.toFloatArray(); - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - mat.put(i, j, arrs[i * cols + j]); - } - } - return mat; - } - - /** - * float NDArray To Mat - * @param ndArray - * @return - */ - public static Mat floatNDArrayToMat(NDArray ndArray) { - int rows = (int) (ndArray.getShape().get(0)); - int cols = (int) (ndArray.getShape().get(1)); - Mat mat = new Mat(rows, cols, CvType.CV_32F); - - float[] arrs = ndArray.toFloatArray(); - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - mat.put(i, j, arrs[i * cols + j]); - } - } - - return mat; - - } - - /** - * uint8 NDArray To Mat - * @param ndArray - * @return - */ - public static Mat uint8NDArrayToMat(NDArray ndArray) { - int rows = (int) (ndArray.getShape().get(0)); - int cols = (int) (ndArray.getShape().get(1)); - Mat mat = new Mat(rows, cols, CvType.CV_8U); - - byte[] arrs = ndArray.toByteArray(); - - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - mat.put(i, j, arrs[i * cols + j]); - } - } - return mat; - } - - /** - * float[][] Array To Mat - * @param arr - * @return - */ - public static Mat floatArrayToMat(float[][] arr) { - int rows = arr.length; - int cols = arr[0].length; - Mat mat = new Mat(rows, cols, CvType.CV_32F); - - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - mat.put(i, j, arr[i][j]); - } - } - - return mat; - } - - /** - * byte[][] Array To Mat - * @param arr - * @return - */ - public static Mat uint8ArrayToMat(byte[][] arr) { - int rows = arr.length; - int cols = arr[0].length; - Mat mat = new Mat(rows, cols, CvType.CV_8U); - - for (int i = 0; i < rows; i++) { - for (int j = 0; j < cols; j++) { - mat.put(i, j, arr[i][j]); - } - } - - return mat; - } - - /** - * List To Mat - * @param points - * @return - */ - public static Mat toMat(List points) { - Mat mat = new Mat(points.size(), 2, CvType.CV_32F); - for (int i = 0; i < points.size(); i++) { - ai.djl.modality.cv.output.Point point = points.get(i); - mat.put(i, 0, (float) point.getX()); - mat.put(i, 1, (float) point.getY()); - } - - return mat; - } -} diff --git a/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java b/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java index 2c87c7a..83dc2b1 100644 --- a/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java +++ b/ocr/src/main/java/cn/smartjavaai/ocr/utils/OcrUtils.java @@ -1,7 +1,6 @@ package cn.smartjavaai.ocr.utils; 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; @@ -10,16 +9,14 @@ import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; import ai.djl.ndarray.NDManager; import ai.djl.opencv.OpenCVImageFactory; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.entity.Point; -import cn.smartjavaai.common.utils.ImageUtils; -import cn.smartjavaai.common.utils.OpenCVUtils; -import cn.smartjavaai.common.utils.PointUtils; +import cn.smartjavaai.common.utils.*; 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; import lombok.extern.slf4j.Slf4j; import org.apache.commons.collections.CollectionUtils; import org.opencv.core.Mat; @@ -28,7 +25,6 @@ import org.opencv.imgproc.Imgproc; import java.awt.*; import java.awt.image.BufferedImage; -import java.math.BigDecimal; import java.util.*; import java.util.List; @@ -79,33 +75,6 @@ public class OcrUtils { - /** - * 图片旋转 - * - * @param manager - * @param image - * @return - */ - public static Image rotateImg(NDManager manager, Image image) { - NDArray rotated = NDImageUtils.rotate90(image.toNDArray(manager), 1); - return ImageFactory.getInstance().fromNDArray(rotated); - } - - /** - * 逆时针旋转图片 - * - * @param image - * @param times - * @return - */ - public static Image rotateImg(Image image, int times) { - try (NDManager manager = NDManager.newBaseManager()) { - NDArray rotated = NDImageUtils.rotate90(image.toNDArray(manager), times); - return OpenCVImageFactory.getInstance().fromNDArray(rotated); - } - } - - /** * 逆时针旋转图片 * @@ -199,66 +168,21 @@ public class OcrUtils { /** * 透视变换 + 裁剪 * @param srcMat - * @param landMarks - * @return - */ - public static Image transformAndCrop(Mat srcMat, List landMarks){ - if (landMarks == null || landMarks.size() != 4) { - throw new IllegalArgumentException("必须提供4个关键点"); - } - - // 步骤 1:排序为 左上、右上、右下、左下 - List ordered = PointUtils.orderPoints(landMarks); - - ai.djl.modality.cv.output.Point lt = ordered.get(0); - ai.djl.modality.cv.output.Point rt = ordered.get(1); - ai.djl.modality.cv.output.Point rb = ordered.get(2); - ai.djl.modality.cv.output.Point lb = ordered.get(3); - - // 步骤 2:计算目标图像尺寸(宽、高) - int img_crop_width = (int) Math.max( - PointUtils.distance(lt, rt), - PointUtils.distance(rb, lb) - ); - int img_crop_height = (int) Math.max( - PointUtils.distance(lt, lb), - PointUtils.distance(rt, rb) - ); - - // 步骤 3:构造目标坐标点 - List dstPoints = Arrays.asList( - new ai.djl.modality.cv.output.Point(0, 0), - new ai.djl.modality.cv.output.Point(img_crop_width, 0), - new ai.djl.modality.cv.output.Point(img_crop_width, img_crop_height), - new ai.djl.modality.cv.output.Point(0, img_crop_height) - ); - - // 步骤 4:透视变换 - Mat srcPoint2f = OcrNDArrayUtils.toMat(ordered); - Mat dstPoint2f = OcrNDArrayUtils.toMat(dstPoints); - Mat cvMat = OpenCVUtils.perspectiveTransform(srcMat, srcPoint2f, dstPoint2f); - - // 步骤 5:转为 DJL Image + 裁剪 - Image subImg = OpenCVImageFactory.getInstance().fromImage(cvMat); - subImg = subImg.getSubImage(0, 0, img_crop_width, img_crop_height); - - // 释放资源 - cvMat.release(); - srcPoint2f.release(); - dstPoint2f.release(); - - return subImg; - } - - - - /** - * 透视变换+裁剪 - * @param srcMat * @param box * @return */ public static Image transformAndCrop(Mat srcMat, OcrBox box){ + Mat subImg = transformAndCropToMat(srcMat, box); + return SmartImageFactory.getInstance().fromMat(subImg); + } + + /** + * 透视变换 + 裁剪 + * @param srcMat + * @param box + * @return + */ + public static Mat transformAndCropToMat(Mat srcMat, OcrBox box){ float[] pointsArr = box.toFloatArray(); float[] lt = java.util.Arrays.copyOfRange(pointsArr, 0, 2); float[] rt = java.util.Arrays.copyOfRange(pointsArr, 2, 4); @@ -276,18 +200,15 @@ public class OcrUtils { dstPoints.add(new ai.djl.modality.cv.output.Point(img_crop_width, 0)); dstPoints.add(new ai.djl.modality.cv.output.Point(img_crop_width, img_crop_height)); dstPoints.add(new ai.djl.modality.cv.output.Point(0, img_crop_height)); - Mat srcPoint2f = OcrNDArrayUtils.toMat(srcPoints); - Mat dstPoint2f = OcrNDArrayUtils.toMat(dstPoints); + Mat srcPoint2f = DJLCommonUtils.toMat(srcPoints); + Mat dstPoint2f = DJLCommonUtils.toMat(dstPoints); //透视变换 Mat cvMat = OpenCVUtils.perspectiveTransform(srcMat, srcPoint2f, dstPoint2f); - Image subImg = OpenCVImageFactory.getInstance().fromImage(cvMat); - //ImageUtils.saveImage(subImg, i + ".png", "build/output"); - //变换后裁剪 - subImg = subImg.getSubImage(0, 0, img_crop_width, img_crop_height); - cvMat.release(); srcPoint2f.release(); dstPoint2f.release(); - return subImg; + Mat result = OpenCVUtils.getSubImage(cvMat, 0, 0, img_crop_width, img_crop_height); + cvMat.release(); + return result; } /** @@ -307,65 +228,116 @@ public class OcrUtils { /** - * 绘制文本框及文本 - * @param image - * @param ocrInfo + * 将 OCR 结果转换为多边形标签列表 + * @param ocrItemList OCR 识别结果 + * @return PolygonLabel 列表 */ - public static void drawRectWithText(BufferedImage image, OcrInfo ocrInfo, int fontSize) { - // 将绘制图像转换为Graphics2D - Graphics2D g = (Graphics2D) image.getGraphics(); - try { - Font font = new Font("楷体", Font.PLAIN, fontSize); - g.setFont(font); - g.setColor(new Color(0, 0, 255)); - // 声明画笔属性 :粗 细(单位像素)末端无修饰 折线处呈尖角 - BasicStroke bStroke = new BasicStroke(2, BasicStroke.CAP_BUTT, BasicStroke.JOIN_MITER); - g.setStroke(bStroke); - List ocrItemList = ocrInfo.getOcrItemList(); - if(CollectionUtils.isNotEmpty(ocrInfo.getLineList())){ - ocrItemList = ocrInfo.flattenLines(); - } - for(OcrItem item : ocrItemList){ - OcrBox box = item.getOcrBox(); - int[] xPoints = { - (int)box.getTopLeft().getX(), - (int)box.getTopRight().getX(), - (int)box.getBottomRight().getX(), - (int)box.getBottomLeft().getX(), - (int)box.getTopLeft().getX() - }; - int[] yPoints = { - (int)box.getTopLeft().getY(), - (int)box.getTopRight().getY(), - (int)box.getBottomRight().getY(), - (int)box.getBottomLeft().getY(), - (int)box.getTopLeft().getY() - }; - g.drawPolyline(xPoints, yPoints, 5); - g.drawString(item.getText(), xPoints[0], yPoints[0]); - } - } finally { - g.dispose(); + public static List toPolygonLabelList(List ocrItemList) { + List polygonLabelList = new ArrayList<>(); + if (ocrItemList == null || ocrItemList.isEmpty()) { + return polygonLabelList; } + for (OcrItem item : ocrItemList) { + if (item.getOcrBox() == null) continue; + List points = Arrays.asList( + item.getOcrBox().getTopLeft(), + item.getOcrBox().getTopRight(), + item.getOcrBox().getBottomRight(), + item.getOcrBox().getBottomLeft() + ); + String text = null; + //角度 + if(item.getAngle() != null){ + text = item.getAngle().getValue(); + }else{ + text = item.getText(); + } + PolygonLabel label = new PolygonLabel(points, text); + polygonLabelList.add(label); + } + return polygonLabelList; } + /** + * 将 OCR 结果转换为多边形标签列表 + * @param boxList + * @return PolygonLabel 列表 + */ + public static List ocrBoxtoPolygonLabelList(List boxList) { + List polygonLabelList = new ArrayList<>(); + if (boxList == null || boxList.isEmpty()) { + return polygonLabelList; + } + for (OcrBox item : boxList) { + List points = Arrays.asList( + item.getTopLeft(), + item.getTopRight(), + item.getBottomRight(), + item.getBottomLeft() + ); + PolygonLabel label = new PolygonLabel(points); + polygonLabelList.add(label); + } + return polygonLabelList; + } + + + /** + * OCR 结果绘制 + */ + public static void drawOcrResult(BufferedImage image, OcrInfo ocrInfo, int fontSize) { + List ocrItemList = ocrInfo.getOcrItemList(); + if (CollectionUtils.isNotEmpty(ocrInfo.getLineList())) { + ocrItemList = ocrInfo.flattenLines(); + } + List polygonLabelList = toPolygonLabelList(ocrItemList); + BufferedImageUtils.drawPolygonWithText(image, polygonLabelList, fontSize); + } + + /** + * OCR 结果绘制 + */ + public static void drawOcrResult(BufferedImage image, List ocrItemList, int fontSize) { + List polygonLabelList = toPolygonLabelList(ocrItemList); + BufferedImageUtils.drawPolygonWithText(image, polygonLabelList, fontSize); + } + + /** + * OCR 结果绘制 + */ + public static void drawOcrDetResult(Image image, List ocrBoxList, int fontSize) { + List polygonLabelList = ocrBoxtoPolygonLabelList(ocrBoxList); + ImageUtils.drawPolygonWithText(image, polygonLabelList, fontSize); + } + + /** + * OCR 结果绘制 + */ + public static void drawOcrResult(Image image, List ocrItemList, int fontSize) { + List polygonLabelList = toPolygonLabelList(ocrItemList); + ImageUtils.drawPolygonWithText(image, polygonLabelList, fontSize); + } + + + /** * 绘制文本框及文本 * @param srcMat * @param itemList */ - public static void drawRectWithText(Mat srcMat, List itemList) { - for(OcrItem item : itemList){ - OcrBox ocrBox = item.getOcrBox(); - Imgproc.line(srcMat, ocrBox.getTopLeft().toCvPoint(), ocrBox.getTopRight().toCvPoint(), new Scalar(0, 255, 0), 1); - Imgproc.line(srcMat, ocrBox.getTopRight().toCvPoint(), ocrBox.getBottomRight().toCvPoint(), new Scalar(0, 255, 0),1); - Imgproc.line(srcMat, ocrBox.getBottomRight().toCvPoint(), ocrBox.getBottomLeft().toCvPoint(), new Scalar(0, 255, 0),1); - Imgproc.line(srcMat, ocrBox.getBottomLeft().toCvPoint(), ocrBox.getTopLeft().toCvPoint(), new Scalar(0, 255, 0), 1); - // 中文乱码 - Imgproc.putText(srcMat, item.getAngle().getValue(), ocrBox.getTopLeft().toCvPoint(), Imgproc.FONT_HERSHEY_SCRIPT_SIMPLEX, 1.0, new Scalar(0, 255, 0), 1); - } - } +// public static void drawRectWithText(Mat srcMat, List itemList) { +// for(OcrItem item : itemList){ +// OcrBox ocrBox = item.getOcrBox(); +// Imgproc.line(srcMat, ocrBox.getTopLeft().toCvPoint(), ocrBox.getTopRight().toCvPoint(), new Scalar(0, 255, 0), 1); +// Imgproc.line(srcMat, ocrBox.getTopRight().toCvPoint(), ocrBox.getBottomRight().toCvPoint(), new Scalar(0, 255, 0),1); +// Imgproc.line(srcMat, ocrBox.getBottomRight().toCvPoint(), ocrBox.getBottomLeft().toCvPoint(), new Scalar(0, 255, 0),1); +// Imgproc.line(srcMat, ocrBox.getBottomLeft().toCvPoint(), ocrBox.getTopLeft().toCvPoint(), new Scalar(0, 255, 0), 1); +// // 中文乱码 +// Imgproc.putText(srcMat, item.getAngle().getValue(), ocrBox.getTopLeft().toCvPoint(), Imgproc.FONT_HERSHEY_SIMPLEX, 1.0, new Scalar(0, 255, 0), 1); +// } +// } + public static List convertToPlateInfo(DetectedObjects detectedObjects, Image image) { List plateInfoList = new ArrayList<>(); @@ -414,7 +386,7 @@ public class OcrUtils { Imgproc.line(srcMat, ocrBox.getBottomRight().toCvPoint(), ocrBox.getBottomLeft().toCvPoint(), new Scalar(0, 0, 255),1); Imgproc.line(srcMat, ocrBox.getBottomLeft().toCvPoint(), ocrBox.getTopLeft().toCvPoint(), new Scalar(0, 0, 255), 1); // 中文乱码 - ImageUtils.putTextWithBackground(srcMat, plateInfo.getPlateNumber() + " " + plateInfo.getPlateColor(), ocrBox.getTopLeft().toCvPoint(), new Scalar(255, 255, 255), new Scalar(0, 0, 0), 1); + OpenCVUtils.putTextWithBackground(srcMat, plateInfo.getPlateNumber() + " " + plateInfo.getPlateColor(), ocrBox.getTopLeft().toCvPoint(), new Scalar(255, 255, 255), new Scalar(0, 0, 0), 1); } } @@ -424,33 +396,19 @@ public class OcrUtils { public static void drawPlateInfo(BufferedImage image, List plateInfoList) { // 将绘制图像转换为Graphics2D Graphics2D graphics = (Graphics2D) image.getGraphics(); - try { - graphics.setColor(Color.RED);// 边框颜色 - graphics.setStroke(new BasicStroke(2)); // 线宽2像素 - graphics.setRenderingHint(RenderingHints.KEY_ANTIALIASING, - RenderingHints.VALUE_ANTIALIAS_ON); // 抗锯齿 - int stroke = 2; - for(PlateInfo plateInfo : plateInfoList){ - DetectionRectangle rectangle = plateInfo.getDetectionRectangle(); - graphics.setColor(Color.RED);// 边框颜色 - //绘制车牌框 - graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); - graphics.setColor(Color.BLACK);// 字体颜色 - ImageUtils.drawText(graphics, plateInfo.getPlateNumber() + " " + plateInfo.getPlateColor(), rectangle.getX(), rectangle.getY(), stroke, 4); - OcrBox ocrBox = plateInfo.getBox(); - //绘制关键点 - graphics.setColor(Color.BLUE); - graphics.drawRect((int)ocrBox.getTopLeft().getX(), (int)ocrBox.getTopLeft().getY(), 2, 2); - graphics.setColor(Color.GREEN); - graphics.drawRect((int)ocrBox.getTopRight().getX(), (int)ocrBox.getTopRight().getY(), 2, 2); - graphics.setColor(Color.RED); - graphics.drawRect((int)ocrBox.getBottomLeft().getX(), (int)ocrBox.getBottomLeft().getY(), 2, 2); - graphics.setColor(Color.CYAN); - graphics.drawRect((int)ocrBox.getBottomRight().getX(), (int)ocrBox.getBottomRight().getY(), 2, 2); - } - } finally { - graphics.dispose(); + for(PlateInfo plateInfo : plateInfoList){ + DetectionRectangle rectangle = plateInfo.getDetectionRectangle(); + OcrBox ocrBox = plateInfo.getBox(); + String text = plateInfo.getPlateNumber() + " " + plateInfo.getPlateColor(); + BufferedImageUtils.drawRectAndText(graphics, rectangle, text, Color.BLACK); + List keyPoints = new ArrayList(); + keyPoints.add(ocrBox.getTopLeft()); + keyPoints.add(ocrBox.getTopRight()); + keyPoints.add(ocrBox.getBottomRight()); + keyPoints.add(ocrBox.getBottomLeft()); + BufferedImageUtils.drawKeyPoints(graphics, keyPoints, Color.GREEN); } + graphics.dispose(); } diff --git a/pom.xml b/pom.xml index 831d1f7..64b93fd 100644 --- a/pom.xml +++ b/pom.xml @@ -7,7 +7,7 @@ SmartJavaAI cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 pom SmartJavaAI @@ -19,6 +19,7 @@ ocr bom speech + @@ -205,6 +206,41 @@ 1.5.10 + + ch.qos.logback + logback-classic + 1.2.3 + runtime + + + + com.alibaba + fastjson + 1.2.83 + runtime + + + + + + + + + + + + + + + + + + junit + junit + 4.13.2 + test + + diff --git a/speech/pom.xml b/speech/pom.xml index d554c42..5ed9427 100644 --- a/speech/pom.xml +++ b/speech/pom.xml @@ -6,7 +6,7 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 speech @@ -51,7 +51,7 @@ - 1.0.24 + 1.0.25 speech SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java b/speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java index c26564b..bab3611 100644 --- a/speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java +++ b/speech/src/main/java/cn/smartjavaai/speech/asr/factory/SpeechRecognizerFactory.java @@ -89,6 +89,7 @@ public class SpeechRecognizerFactory { throw new AsrException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -100,4 +101,26 @@ public class SpeechRecognizerFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(AsrModelEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java b/speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java index d158610..63b7ace 100644 --- a/speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java +++ b/speech/src/main/java/cn/smartjavaai/speech/asr/model/SpeechRecognizer.java @@ -45,4 +45,7 @@ public interface SpeechRecognizer extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java b/speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java index f03acef..cbb4fcc 100644 --- a/speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java +++ b/speech/src/main/java/cn/smartjavaai/speech/asr/model/VoskRecognizer.java @@ -9,6 +9,7 @@ import cn.smartjavaai.speech.asr.audio.SmartAudioFactory; import cn.smartjavaai.speech.asr.config.AsrModelConfig; import cn.smartjavaai.speech.asr.entity.*; import cn.smartjavaai.speech.asr.exception.AsrException; +import cn.smartjavaai.speech.asr.factory.SpeechRecognizerFactory; import cn.smartjavaai.speech.asr.pool.WhisperStatePool; import cn.smartjavaai.speech.utils.AudioUtils; import com.google.gson.JsonArray; @@ -50,8 +51,11 @@ public class VoskRecognizer implements SpeechRecognizer{ private Model model; + private AsrModelConfig config; + @Override public void loadModel(AsrModelConfig config) { + this.config = config; if(StringUtils.isBlank(config.getModelPath())){ throw new AsrException("modelPath is null"); } @@ -320,6 +324,9 @@ public class VoskRecognizer implements SpeechRecognizer{ @Override public void close() throws Exception { + if (fromFactory) { + SpeechRecognizerFactory.removeFromCache(config.getModelEnum()); + } if(model != null){ model.close(); } @@ -339,4 +346,13 @@ public class VoskRecognizer implements SpeechRecognizer{ } + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java b/speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java index 67e3d8b..af9cc1a 100644 --- a/speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java +++ b/speech/src/main/java/cn/smartjavaai/speech/asr/model/WhisperRecognizer.java @@ -12,6 +12,7 @@ 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.exception.AsrException; +import cn.smartjavaai.speech.asr.factory.SpeechRecognizerFactory; import cn.smartjavaai.speech.asr.pool.WhisperStatePool; import io.github.givimad.whisperjni.*; import lombok.extern.slf4j.Slf4j; @@ -41,8 +42,11 @@ public class WhisperRecognizer implements SpeechRecognizer{ private WhisperStatePool statePool; + private AsrModelConfig config; + @Override public void loadModel(AsrModelConfig config) { + this.config = config; if(StringUtils.isBlank(config.getModelPath())){ throw new AsrException("modelPath is null"); } @@ -185,6 +189,9 @@ public class WhisperRecognizer implements SpeechRecognizer{ @Override public void close() throws Exception { + if (fromFactory) { + SpeechRecognizerFactory.removeFromCache(config.getModelEnum()); + } if(statePool != null){ statePool.close(); } @@ -216,4 +223,14 @@ public class WhisperRecognizer implements SpeechRecognizer{ false ); } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/translate/pom.xml b/translate/pom.xml index 7e16e26..fbd96d3 100644 --- a/translate/pom.xml +++ b/translate/pom.xml @@ -25,7 +25,7 @@ - 1.0.24 + 1.0.25 translate SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java b/translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java index e13ad9e..3756ed1 100644 --- a/translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java +++ b/translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java @@ -90,6 +90,7 @@ public class TranslationModelFactory { throw new TranslationException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -102,4 +103,26 @@ public class TranslationModelFactory { log.debug("缓存目录:{}", Config.getCachePath()); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(TranslationModeEnum modelEnum) { + modelMap.remove(modelEnum); + } + } diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java b/translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java index 1b9f22b..0319f36 100644 --- a/translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java @@ -24,6 +24,7 @@ import cn.smartjavaai.translation.config.NllbSearchConfig; import cn.smartjavaai.translation.entity.GreedyBatchTensorList; import cn.smartjavaai.translation.entity.TranslateParam; import cn.smartjavaai.translation.exception.TranslationException; +import cn.smartjavaai.translation.factory.TranslationModelFactory; import cn.smartjavaai.translation.model.translator.NllbDecoder2Translator; import cn.smartjavaai.translation.model.translator.NllbDecoderTranslator; import cn.smartjavaai.translation.model.translator.NllbEncoderTranslator; @@ -274,6 +275,9 @@ public class NllbModel implements TranslationModel{ @Override public void close() throws Exception { + if (fromFactory) { + TranslationModelFactory.removeFromCache(config.getModelEnum()); + } try { if (nllbModel != null) { nllbModel.close(); @@ -310,4 +314,14 @@ public class NllbModel implements TranslationModel{ log.warn("关闭 decode2PredictorPool 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java b/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java index 60f6608..83f9c54 100644 --- a/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/OpusMtModel.java @@ -35,6 +35,7 @@ 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.factory.TranslationModelFactory; import cn.smartjavaai.translation.model.translator.NllbDecoder2Translator; import cn.smartjavaai.translation.model.translator.NllbDecoderTranslator; import cn.smartjavaai.translation.model.translator.NllbEncoderTranslator; @@ -403,6 +404,9 @@ public class OpusMtModel implements TranslationModel{ @Override public void close() throws Exception { + if (fromFactory) { + TranslationModelFactory.removeFromCache(config.getModelEnum()); + } try { if (model != null) { model.close(); @@ -439,4 +443,14 @@ public class OpusMtModel implements TranslationModel{ log.warn("关闭 decode2PredictorPool 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java b/translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java index a9d55ec..cef271c 100644 --- a/translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java +++ b/translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java @@ -38,7 +38,9 @@ public interface TranslationModel extends AutoCloseable{ } - + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/vision/pom.xml b/vision/pom.xml index 67e9b20..4ee455c 100644 --- a/vision/pom.xml +++ b/vision/pom.xml @@ -6,11 +6,11 @@ cn.smartjavaai smartjavaai-parent - 1.0.24 + 1.0.25 vision - 1.0.24 + 1.0.25 vision SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java b/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java index 87f2188..d4c547f 100644 --- a/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java +++ b/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModel.java @@ -34,4 +34,8 @@ public interface ActionRecModel extends AutoCloseable{ } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } + } diff --git a/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModelFactory.java b/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModelFactory.java index 9ddd820..27d43cf 100644 --- a/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModelFactory.java +++ b/vision/src/main/java/cn/smartjavaai/action/model/ActionRecModelFactory.java @@ -76,6 +76,7 @@ public class ActionRecModelFactory { throw new DetectionException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -90,14 +91,6 @@ public class ActionRecModelFactory { } - /** - * 移除缓存的模型 - * @param modelEnum - */ - public static void removeFromCache(ActionRecModelEnum modelEnum) { - modelMap.remove(modelEnum); - } - // 初始化默认算法 static { @@ -107,5 +100,27 @@ public class ActionRecModelFactory { registerAlgorithm(ActionRecModelEnum.VIT_BASE_PATCH16_224_DJL, CommonActionRecModel.class); log.debug("缓存目录:{}", Config.getCachePath()); } + + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + + /** + * 移除缓存的模型 + * @param modelEnum + */ + public static void removeFromCache(ActionRecModelEnum modelEnum) { + modelMap.remove(modelEnum); + } } diff --git a/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java b/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java index 34808fd..0f1e328 100644 --- a/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java +++ b/vision/src/main/java/cn/smartjavaai/action/model/CommonActionRecModel.java @@ -110,6 +110,9 @@ public class CommonActionRecModel implements ActionRecModel{ @Override public void close() throws Exception { + if (fromFactory) { + ActionRecModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -125,4 +128,14 @@ public class CommonActionRecModel implements ActionRecModel{ log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java b/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java index d5c3ec5..c771cb4 100644 --- a/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java @@ -8,10 +8,12 @@ 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.action.model.ActionRecModelFactory; 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.common.utils.ImageUtils; import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig; import cn.smartjavaai.instanceseg.criteria.InstanceSegCriteriaFactory; import cn.smartjavaai.instanceseg.exception.InstanceSegException; @@ -124,8 +126,9 @@ public class CommonInstanceSegModel implements InstanceSegModel { @Override public R detectAndDraw(String imagePath, String outputPath) { + Image img = null; try { - Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ return R.fail(R.Status.NO_OBJECT_DETECTED); @@ -136,11 +139,16 @@ public class CommonInstanceSegModel implements InstanceSegModel { return R.ok(detectionResponse); } catch (IOException e) { throw new InstanceSegException(e); + } finally { + ImageUtils.releaseOpenCVMat(img); } } @Override public void close() throws Exception { + if (fromFactory) { + InstanceSegModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -156,4 +164,14 @@ public class CommonInstanceSegModel implements InstanceSegModel { log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java b/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java index 1ad627f..9f5d94b 100644 --- a/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModel.java @@ -43,5 +43,8 @@ public interface InstanceSegModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModelFactory.java b/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModelFactory.java index e5dbc2a..ca52395 100644 --- a/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModelFactory.java +++ b/vision/src/main/java/cn/smartjavaai/instanceseg/model/InstanceSegModelFactory.java @@ -75,6 +75,7 @@ public class InstanceSegModelFactory { throw new DetectionException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -97,6 +98,19 @@ public class InstanceSegModelFactory { modelMap.remove(modelEnum); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } // 初始化默认算法 static { diff --git a/vision/src/main/java/cn/smartjavaai/obb/model/CommonObbDetModel.java b/vision/src/main/java/cn/smartjavaai/obb/model/CommonObbDetModel.java index 3201935..f373ed1 100644 --- a/vision/src/main/java/cn/smartjavaai/obb/model/CommonObbDetModel.java +++ b/vision/src/main/java/cn/smartjavaai/obb/model/CommonObbDetModel.java @@ -12,6 +12,8 @@ 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.common.utils.ImageUtils; +import cn.smartjavaai.instanceseg.model.InstanceSegModelFactory; import cn.smartjavaai.obb.config.ObbDetModelConfig; import cn.smartjavaai.obb.criteria.ObbDetCriteriaFactory; import cn.smartjavaai.obb.entity.ObbResult; @@ -125,8 +127,9 @@ public class CommonObbDetModel implements ObbDetModel { @Override public R detectAndDraw(String imagePath, String outputPath) { + Image img = null; try { - Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); ObbResult obbResult = detectCore(img); if(Objects.isNull(obbResult) || CollectionUtils.isEmpty(obbResult.getRotatedBoxeList())){ return R.fail(R.Status.NO_OBJECT_DETECTED); @@ -137,11 +140,16 @@ public class CommonObbDetModel implements ObbDetModel { return R.ok(detectionResponse); } catch (IOException e) { throw new ObbDetException(e); + } finally { + ImageUtils.releaseOpenCVMat(img); } } @Override public void close() throws Exception { + if (fromFactory) { + ObbDetModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -157,4 +165,14 @@ public class CommonObbDetModel implements ObbDetModel { log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModel.java b/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModel.java index f455fce..ae0cd00 100644 --- a/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModel.java +++ b/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModel.java @@ -57,4 +57,7 @@ public interface ObbDetModel extends AutoCloseable{ } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModelFactory.java b/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModelFactory.java index c375f72..2cb4f79 100644 --- a/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModelFactory.java +++ b/vision/src/main/java/cn/smartjavaai/obb/model/ObbDetModelFactory.java @@ -75,6 +75,7 @@ public class ObbDetModelFactory { throw new DetectionException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -97,6 +98,19 @@ public class ObbDetModelFactory { modelMap.remove(modelEnum); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } // 初始化默认算法 static { diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java b/vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java index 4b55bf4..daa60d2 100644 --- a/vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/model/DetectorModel.java @@ -4,22 +4,20 @@ 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.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.common.utils.BufferedImageUtils; import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.OpenCVUtils; import cn.smartjavaai.objectdetection.config.DetectorModelConfig; import cn.smartjavaai.objectdetection.criteria.CriteriaBuilderFactory; import cn.smartjavaai.objectdetection.exception.DetectionException; -import cn.smartjavaai.vision.utils.CategoryMaskFilter; import cn.smartjavaai.vision.utils.DetectedObjectsFilter; import cn.smartjavaai.vision.utils.DetectorUtils; import lombok.extern.slf4j.Slf4j; @@ -31,8 +29,6 @@ import javax.imageio.ImageIO; import java.awt.image.BufferedImage; import java.io.*; import java.nio.file.Paths; -import java.util.ArrayList; -import java.util.List; import java.util.Objects; /** @@ -86,23 +82,21 @@ public class DetectorModel implements AutoCloseable{ * @return * @throws Exception */ + @Deprecated public DetectionResponse detect(String imagePath){ if(!FileUtils.isFileExists(imagePath)){ throw new DetectionException("图像文件不存在"); } Image image = null; try { - image = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); - DetectedObjects detectedObjects = detect(image); + image = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detectedObjects = detectCore(image); return DetectorUtils.convertToDetectionResponse(detectedObjects, image); } catch (Exception e) { throw new DetectionException(e); } finally { - if (image != null){ - ((Mat)image.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(image); } - } @@ -117,8 +111,8 @@ public class DetectorModel implements AutoCloseable{ } Image img = null; try { - img = ImageFactory.getInstance().fromFile(Paths.get(imagePath)); - DetectedObjects detectedObjects = detect(img); + img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath)); + DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ throw new DetectionException("未检测到图片中的物体"); } @@ -129,9 +123,7 @@ public class DetectorModel implements AutoCloseable{ } catch (IOException e) { throw new DetectionException(e); } finally { - if (img != null){ - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -140,6 +132,7 @@ public class DetectorModel implements AutoCloseable{ * @param imageData * @return */ + @Deprecated public DetectionResponse detect(byte[] imageData){ if(Objects.isNull(imageData)){ throw new DetectionException("图像无效"); @@ -159,21 +152,20 @@ public class DetectorModel implements AutoCloseable{ * @param image * @return */ + @Deprecated public DetectionResponse detect(BufferedImage image){ - if(!ImageUtils.isImageValid(image)){ + if(!BufferedImageUtils.isImageValid(image)){ throw new DetectionException("图像无效"); } Image img = null; try { - img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image)); - DetectedObjects detectedObjects = detect(img); + img = SmartImageFactory.getInstance().fromBufferedImage(image); + DetectedObjects detectedObjects = detectCore(img); return DetectorUtils.convertToDetectionResponse(detectedObjects, img); } catch (Exception e) { throw new DetectionException(e); } finally { - if (img != null) { - ((Mat)img.getWrappedImage()).release(); - } + ImageUtils.releaseOpenCVMat(img); } } @@ -183,30 +175,20 @@ public class DetectorModel implements AutoCloseable{ * @param sourceImage * @return */ + @Deprecated public BufferedImage detectAndDraw(BufferedImage sourceImage){ - if(!ImageUtils.isImageValid(sourceImage)){ + if(!BufferedImageUtils.isImageValid(sourceImage)){ throw new DetectionException("图像无效"); } - Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage)); - DetectedObjects detectedObjects = detect(img); + Image img = SmartImageFactory.getInstance().fromBufferedImage(sourceImage); + DetectedObjects detectedObjects = detectCore(img); if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ throw new DetectionException("未检测到图片中的物体"); } img.drawBoundingBoxes(detectedObjects); - try { - ByteArrayOutputStream outputStream = new ByteArrayOutputStream(); - // 调用 save 方法将 Image 写入字节流 - img.save(outputStream, "png"); - // 将字节流转换为 BufferedImage - byte[] imageBytes = outputStream.toByteArray(); - return ImageIO.read(new ByteArrayInputStream(imageBytes)); - } catch (IOException e) { - throw new DetectionException("导出图片失败", e); - } finally { - if (img != null) { - ((Mat)img.getWrappedImage()).release(); - } - } + BufferedImage drawnImage = ImageUtils.toBufferedImage(img); + ImageUtils.releaseOpenCVMat(img); + return drawnImage; } /** @@ -214,7 +196,35 @@ public class DetectorModel implements AutoCloseable{ * @param image * @return */ - public DetectedObjects detect(Image image){ + public DetectionResponse detect(Image image){ + DetectedObjects detectedObjects = detectCore(image); + return DetectorUtils.convertToDetectionResponse(detectedObjects, image); + } + + + /** + * 检测并绘制 + * @param image + * @return + */ + public DetectionResponse detectAndDraw(Image image){ + DetectedObjects detectedObjects = detectCore(image); + if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ + throw new DetectionException("未检测到图片中的物体"); + } + Image img = ImageUtils.copy(image); + DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, img); + img.drawBoundingBoxes(detectedObjects); + detectionResponse.setDrawnImage(img); + return detectionResponse; + } + + /** + * 目标检测 + * @param image + * @return + */ + public DetectedObjects detectCore(Image image){ Predictor predictor = null; try { predictor = predictorPool.borrowObject(); 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 index 4a4367e..7ed1a68 100644 --- a/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/PersonDetModel.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/model/person/PersonDetModel.java @@ -58,5 +58,8 @@ public interface PersonDetModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java b/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java index 604ae76..c906efe 100644 --- a/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java +++ b/vision/src/main/java/cn/smartjavaai/objectdetection/stream/StreamDetector.java @@ -9,6 +9,7 @@ import ai.djl.modality.cv.output.Mask; import ai.djl.modality.cv.output.Rectangle; import cn.hutool.core.date.LocalDateTimeUtil; import cn.hutool.core.lang.UUID; +import cn.smartjavaai.common.cv.SmartImageFactory; import cn.smartjavaai.common.entity.*; import cn.smartjavaai.common.enums.VideoSourceType; import cn.smartjavaai.common.utils.ImageUtils; @@ -230,7 +231,7 @@ public class StreamDetector implements AutoCloseable{ mat = converterToMat.convert(frame); if (mat == null) return; - Image image = ImageFactory.getInstance().fromImage(mat); + Image image = SmartImageFactory.getInstance().fromMat(mat); DetectedObjects detectedObjects = predictor.predict(image); // log.info("内部检测结果:{}", detectedObjects.toString()); DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image); diff --git a/vision/src/main/java/cn/smartjavaai/pose/criteria/PoseCriteriaFactory.java b/vision/src/main/java/cn/smartjavaai/pose/criteria/PoseCriteriaFactory.java index 1d54be5..12fcbd7 100644 --- a/vision/src/main/java/cn/smartjavaai/pose/criteria/PoseCriteriaFactory.java +++ b/vision/src/main/java/cn/smartjavaai/pose/criteria/PoseCriteriaFactory.java @@ -4,6 +4,7 @@ import ai.djl.Device; import ai.djl.modality.Classifications; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.output.Joints; +import ai.djl.modality.cv.translator.YoloPoseTranslatorFactory; import ai.djl.repository.zoo.Criteria; import ai.djl.training.util.ProgressBar; import ai.djl.translate.Translator; @@ -20,6 +21,7 @@ import cn.smartjavaai.obb.exception.ObbDetException; import cn.smartjavaai.obb.translator.YoloV11OddTranslator; import cn.smartjavaai.objectdetection.constant.DetectorConstant; import cn.smartjavaai.pose.config.PoseModelConfig; +import cn.smartjavaai.pose.enums.PoseModelEnum; import cn.smartjavaai.pose.exception.PoseException; import org.apache.commons.lang3.StringUtils; @@ -58,19 +60,26 @@ public class PoseCriteriaFactory { * @return */ public static Criteria createDJLCriteria(PoseModelConfig config, Device device) { - if(StringUtils.isNotBlank(config.getModelPath()) - && DJLCommonUtils.isServingPropertiesExists(Paths.get(config.getModelPath()))){ - throw new PoseException("模型所在目录未找到 serving.properties 文件"); +// if(StringUtils.isNotBlank(config.getModelPath()) +// && !DJLCommonUtils.isServingPropertiesExists(Paths.get(config.getModelPath()))){ +// throw new PoseException("模型所在目录未找到 serving.properties 文件"); +// } + + Criteria criteria = null; + if(config.getModelEnum() == PoseModelEnum.YOLO11N_POSE_PT || config.getModelEnum() == PoseModelEnum.YOLO11N_POSE_ONNX + || config.getModelEnum() == PoseModelEnum.YOLOV8N_POSE_PT || config.getModelEnum() == PoseModelEnum.YOLOV8N_POSE_ONNX){ + criteria = + Criteria.builder() + .setTypes(Image.class, Joints[].class) + .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : config.getModelEnum().getModelUri()) + .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) + .optDevice(device) + .optTranslatorFactory(new YoloPoseTranslatorFactory()) + .optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : null) + .optProgress(new ProgressBar()) + .build(); } - Criteria criteria = - Criteria.builder() - .setTypes(Image.class, Joints[].class) - .optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null : config.getModelEnum().getModelUri()) - .optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null) - .optDevice(device) - .optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : null) - .optProgress(new ProgressBar()) - .build(); + return criteria; } diff --git a/vision/src/main/java/cn/smartjavaai/pose/model/CommonPoseModel.java b/vision/src/main/java/cn/smartjavaai/pose/model/CommonPoseModel.java index 463897d..a452707 100644 --- a/vision/src/main/java/cn/smartjavaai/pose/model/CommonPoseModel.java +++ b/vision/src/main/java/cn/smartjavaai/pose/model/CommonPoseModel.java @@ -14,6 +14,7 @@ import cn.smartjavaai.common.entity.R; import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.obb.exception.ObbDetException; +import cn.smartjavaai.obb.model.ObbDetModelFactory; import cn.smartjavaai.objectdetection.config.PersonDetModelConfig; import cn.smartjavaai.objectdetection.criteria.PersonDetCriteriaFactory; import cn.smartjavaai.objectdetection.exception.DetectionException; @@ -121,6 +122,7 @@ public class CommonPoseModel implements PoseModel { } // 调用 save 方法将 Image 写入字节流 img.save(Files.newOutputStream(Paths.get(outputPath)), "png"); + ImageUtils.releaseOpenCVMat(img); return allJoints; } catch (IOException e) { throw new ObbDetException(e); @@ -129,6 +131,9 @@ public class CommonPoseModel implements PoseModel { @Override public void close() throws Exception { + if (fromFactory) { + PoseDetModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -144,4 +149,14 @@ public class CommonPoseModel implements PoseModel { log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/vision/src/main/java/cn/smartjavaai/pose/model/PoseDetModelFactory.java b/vision/src/main/java/cn/smartjavaai/pose/model/PoseDetModelFactory.java index a5e02a3..8b5cc18 100644 --- a/vision/src/main/java/cn/smartjavaai/pose/model/PoseDetModelFactory.java +++ b/vision/src/main/java/cn/smartjavaai/pose/model/PoseDetModelFactory.java @@ -76,6 +76,7 @@ public class PoseDetModelFactory { throw new DetectionException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -99,6 +100,20 @@ public class PoseDetModelFactory { } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } + // 初始化默认算法 static { registerAlgorithm(PoseModelEnum.YOLOV8N_POSE_ONNX, CommonPoseModel.class); diff --git a/vision/src/main/java/cn/smartjavaai/pose/model/PoseModel.java b/vision/src/main/java/cn/smartjavaai/pose/model/PoseModel.java index 0017216..684a136 100644 --- a/vision/src/main/java/cn/smartjavaai/pose/model/PoseModel.java +++ b/vision/src/main/java/cn/smartjavaai/pose/model/PoseModel.java @@ -49,5 +49,8 @@ public interface PoseModel extends AutoCloseable{ throw new UnsupportedOperationException("默认不支持该功能"); } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java b/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java index 5db097d..cc00cd2 100644 --- a/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java +++ b/vision/src/main/java/cn/smartjavaai/semseg/model/CommonSemSegModel.java @@ -18,6 +18,7 @@ import cn.smartjavaai.common.utils.Base64ImageUtils; import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.instanceseg.exception.InstanceSegException; import cn.smartjavaai.objectdetection.exception.DetectionException; +import cn.smartjavaai.pose.model.PoseDetModelFactory; import cn.smartjavaai.semseg.config.SemSegModelConfig; import cn.smartjavaai.semseg.criteria.SemSegCriteriaFactory; import cn.smartjavaai.vision.utils.CategoryMaskFilter; @@ -122,6 +123,7 @@ public class CommonSemSegModel implements SemSegModel { } ImageUtils.drawMask(categoryMask, img, 180, 0); img.save(Files.newOutputStream(Paths.get(outputPath)), "png"); + ImageUtils.releaseOpenCVMat(img); return R.ok(categoryMask); } catch (IOException e) { throw new InstanceSegException(e); @@ -141,6 +143,9 @@ public class CommonSemSegModel implements SemSegModel { @Override public void close() throws Exception { + if (fromFactory) { + SemSegModelFactory.removeFromCache(config.getModelEnum()); + } try { if (predictorPool != null) { predictorPool.close(); @@ -156,4 +161,14 @@ public class CommonSemSegModel implements SemSegModel { log.warn("关闭 model 失败", e); } } + + private boolean fromFactory = false; + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + public boolean isFromFactory() { + return fromFactory; + } } diff --git a/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java b/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java index 37445a8..6b1dc57 100644 --- a/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java +++ b/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModel.java @@ -38,5 +38,8 @@ public interface SemSegModel extends AutoCloseable{ } + default void setFromFactory(boolean fromFactory){ + throw new UnsupportedOperationException("默认不支持该功能"); + } } diff --git a/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModelFactory.java b/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModelFactory.java index 0099541..285ad1c 100644 --- a/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModelFactory.java +++ b/vision/src/main/java/cn/smartjavaai/semseg/model/SemSegModelFactory.java @@ -75,6 +75,7 @@ public class SemSegModelFactory { throw new DetectionException(e); } model.loadModel(config); + model.setFromFactory(true); return model; } @@ -97,6 +98,19 @@ public class SemSegModelFactory { modelMap.remove(modelEnum); } + /** + * 关闭所有已加载的模型 + */ + public void closeAll() { + modelMap.values().forEach(model -> { + try { + model.close(); + } catch (Exception e) { + e.printStackTrace(); + } + }); + modelMap.clear(); + } // 初始化默认算法 static { diff --git a/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java b/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java index aeb8904..c689b17 100644 --- a/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java +++ b/vision/src/main/java/cn/smartjavaai/vision/utils/DetectorUtils.java @@ -60,13 +60,13 @@ public class DetectorUtils { 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){ + if(box instanceof Mask){ Mask mask = (Mask)box; InstanceSegInfo instanceSegInfo = new InstanceSegInfo(className, mask.getProbDist()); detectionInfo.setInstanceSegInfo(instanceSegInfo); + }else if(box instanceof Rectangle){ + ObjectDetInfo objectDetInfo = new ObjectDetInfo(className); + detectionInfo.setObjectDetInfo(objectDetInfo); } detectionInfoList.add(detectionInfo); index++; @@ -106,9 +106,9 @@ public class DetectorUtils { 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.line(srcMat, points.get(3).toCvPoint(), points.get(0).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); + Imgproc.putText(srcMat, box.className, points.get(0).toCvPoint(), Imgproc.FONT_HERSHEY_SIMPLEX, 1.0, new Scalar(0, 255, 0), 1); } }