1、FaceNet 特征提取新增人脸对齐

2、人脸检测新5点人脸关键点定位
3、特征提取接口支持多人脸和最佳人脸提取
4、修复人脸框边界精度问题
5、更新 Maven 发布的 groupId
This commit is contained in:
dengwenjie
2025-04-28 16:13:34 +08:00
parent 2fdc20f610
commit 42d2943a94
50 changed files with 3413 additions and 554 deletions

View File

@@ -36,12 +36,13 @@
<tr> <tr>
<td> <td>
<div align="left"> <div align="left">
<p>人脸检测</p> <p>人脸检测</p>
- 5点人脸关键点定位 <br>
</div> </div>
</td> </td>
<td> <td>
<div align="center"> <div align="center">
<img src="https://cdn.jsdelivr.net/gh/geekwenjie/SmartJavaAI-Site/images/face.jpg" width = "500px"/> <img src="https://cdn.jsdelivr.net/gh/geekwenjie/SmartJavaAI-Site/images/face5point.png" width = "500px"/>
</div> </div>
</td> </td>
</tr> </tr>
@@ -49,6 +50,7 @@
<td> <td>
<div align="left"> <div align="left">
<p>人脸比对11</p> <p>人脸比对11</p>
- 人脸对齐 <br>
</div> </div>
</td> </td>
<td> <td>
@@ -72,7 +74,8 @@
<tr> <tr>
<td> <td>
<div align="left"> <div align="left">
<p>人脸比对1N</p> <p>人脸比对1N</p>
- 人脸对齐 <br>
- 人脸注册 <br> - 人脸注册 <br>
- 人脸库查询<br> - 人脸库查询<br>
- 人脸库删除<br> - 人脸库删除<br>
@@ -100,6 +103,7 @@
</div> </div>
## 🚀 核心亮点 ## 🚀 核心亮点
针对 Java 开发者在集成智能算法时常见的两大痛点: 针对 Java 开发者在集成智能算法时常见的两大痛点:
@@ -176,9 +180,9 @@
```xml ```xml
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-all</artifactId> <artifactId>smartjavaai-all</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</dependency> </dependency>
``` ```
### 3、完整示例代码 ### 3、完整示例代码
@@ -193,6 +197,7 @@
- **[Seetaface6JNI](https://gitee.com/cnsugar/seetaface6JNI)** - **[Seetaface6JNI](https://gitee.com/cnsugar/seetaface6JNI)**
- **[Deep Java Library](https://docs.djl.ai)** - **[Deep Java Library](https://docs.djl.ai)**
- **[AIAS](https://gitee.com/mymagicpower/AIAS)**
@@ -216,7 +221,12 @@
## 更新日志 ## 更新日志
## [v1.0.11] - 2025-04-28
- FaceNet 特征提取新增人脸对齐
- 人脸检测新5点人脸关键点定位
- 特征提取接口支持多人脸和最佳人脸提取
- 修复人脸框边界精度问题
- 更新 Maven 发布的 groupId
## [v1.0.10] - 2025-04-19 ## [v1.0.10] - 2025-04-19
- 兼容 SeetaFace6 在 Linux 系统下的运行 - 兼容 SeetaFace6 在 Linux 系统下的运行
- 新增全局缓存路径设置功能 - 新增全局缓存路径设置功能

26
pom.xml
View File

@@ -4,9 +4,9 @@
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd"> xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion> <modelVersion>4.0.0</modelVersion>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId> <artifactId>smartjavaai-parent</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
<packaging>pom</packaging> <packaging>pom</packaging>
<description>SmartJavaAI</description> <description>SmartJavaAI</description>
<modules> <modules>
@@ -14,6 +14,7 @@
<module>smartjavaai-common</module> <module>smartjavaai-common</module>
<module>smartjavaai-objectdetection</module> <module>smartjavaai-objectdetection</module>
<module>smartjavaai-all</module> <module>smartjavaai-all</module>
<module>smartjavaai-ocr</module>
</modules> </modules>
<properties> <properties>
@@ -21,6 +22,7 @@
<maven.compiler.target>11</maven.compiler.target> <maven.compiler.target>11</maven.compiler.target>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding> <project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
<djl.version>0.32.0</djl.version> <djl.version>0.32.0</djl.version>
</properties> </properties>
<dependencyManagement> <dependencyManagement>
@@ -34,15 +36,15 @@
</dependency> </dependency>
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-common</artifactId> <artifactId>smartjavaai-common</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</dependency> </dependency>
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-face</artifactId> <artifactId>smartjavaai-face</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</dependency> </dependency>
</dependencies> </dependencies>
@@ -173,6 +175,12 @@
<groupId>ai.djl.onnxruntime</groupId> <groupId>ai.djl.onnxruntime</groupId>
<artifactId>onnxruntime-engine</artifactId> <artifactId>onnxruntime-engine</artifactId>
</dependency> </dependency>
<dependency>
<groupId>ai.djl.opencv</groupId>
<artifactId>opencv</artifactId>
</dependency>
<dependency> <dependency>
<groupId>ai.djl.ml.xgboost</groupId> <groupId>ai.djl.ml.xgboost</groupId>
<artifactId>xgboost</artifactId> <artifactId>xgboost</artifactId>
@@ -196,6 +204,12 @@
<version>5.8.16</version> <version>5.8.16</version>
</dependency> </dependency>
<dependency>
<groupId>gov.nist.math</groupId>
<artifactId>jama</artifactId>
<version>1.0.3</version>
</dependency>
</dependencies> </dependencies>

View File

@@ -4,13 +4,13 @@
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd"> xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion> <modelVersion>4.0.0</modelVersion>
<parent> <parent>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId> <artifactId>smartjavaai-parent</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</parent> </parent>
<artifactId>smartjavaai-all</artifactId> <artifactId>smartjavaai-all</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
<name>smartjavaai-all</name> <name>smartjavaai-all</name>
<description>SmartJavaAI</description> <description>SmartJavaAI</description>
<url>https://github.com/geekwenjie/SmartJavaAI</url> <url>https://github.com/geekwenjie/SmartJavaAI</url>
@@ -23,13 +23,13 @@
<dependencies> <dependencies>
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-face</artifactId> <artifactId>smartjavaai-face</artifactId>
<version>${project.version}</version> <version>${project.version}</version>
</dependency> </dependency>
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-objectdetection</artifactId> <artifactId>smartjavaai-objectdetection</artifactId>
<version>${project.version}</version> <version>${project.version}</version>
</dependency> </dependency>

View File

@@ -4,9 +4,9 @@
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd"> xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion> <modelVersion>4.0.0</modelVersion>
<parent> <parent>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId> <artifactId>smartjavaai-parent</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</parent> </parent>
<artifactId>smartjavaai-common</artifactId> <artifactId>smartjavaai-common</artifactId>

View File

@@ -1,5 +1,7 @@
package cn.smartjavaai.common.entity; package cn.smartjavaai.common.entity;
import java.util.List;
/** /**
* 检测结果-矩形区域 * 检测结果-矩形区域
* @author dwj * @author dwj
@@ -11,9 +13,13 @@ public class DetectionRectangle {
public int width; public int width;
public int height; public int height;
public float score; public float score;
public String className; public String className;
/**
* 人脸关键点
*/
private List<Point> keyPoints;
public DetectionRectangle() { public DetectionRectangle() {
} }
@@ -81,4 +87,12 @@ public class DetectionRectangle {
public void setClassName(String className) { public void setClassName(String className) {
this.className = className; this.className = className;
} }
public List<Point> getKeyPoints() {
return keyPoints;
}
public void setKeyPoints(List<Point> keyPoints) {
this.keyPoints = keyPoints;
}
} }

View File

@@ -10,27 +10,27 @@ import java.io.Serializable;
*/ */
public class Point implements Serializable { public class Point implements Serializable {
private static final long serialVersionUID = 1L; private static final long serialVersionUID = 1L;
private int x; private double x;
private int y; private double y;
public Point(int x, int y) { public Point(double x, double y) {
this.x = x; this.x = x;
this.y = y; this.y = y;
} }
public int getX() { public double getX() {
return x; return x;
} }
public void setX(int x) { public void setX(double x) {
this.x = x; this.x = x;
} }
public int getY() { public double getY() {
return y; return y;
} }
public void setY(int y) { public void setY(double y) {
this.y = y; this.y = y;
} }

View File

@@ -1,13 +1,23 @@
package cn.smartjavaai.common.utils; package cn.smartjavaai.common.utils;
import ai.djl.modality.cv.BufferedImageFactory;
import ai.djl.modality.cv.Image;
import ai.djl.ndarray.NDArray;
import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.DetectionResponse;
import javax.imageio.ImageIO; import javax.imageio.ImageIO;
import java.awt.*; import java.awt.*;
import java.awt.image.BufferedImage; import java.awt.image.BufferedImage;
//import java.awt.image.ColorConvertOp; //import java.awt.image.ColorConvertOp;
import java.awt.image.ComponentSampleModel; import java.awt.image.ComponentSampleModel;
import java.awt.image.ImageObserver;
import java.io.File; import java.io.File;
import java.io.IOException; import java.io.IOException;
import java.util.ArrayList;
import java.util.Arrays; import java.util.Arrays;
import java.util.List;
import java.util.Objects;
/** /**
* 图片处理工具类 * 图片处理工具类
@@ -88,6 +98,86 @@ public class ImageUtils {
return image != null && image.getWidth() > 0 && image.getHeight() > 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 drawImageRect(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();
}
}
/**
* 画检测框
*
* @param image
* @param x
* @param y
* @param width
* @param height
*/
public static void drawImageRect(Image image, DetectionResponse detectionResponse) {
if(Objects.nonNull(detectionResponse) && Objects.nonNull(detectionResponse.getRectangleList()) && !detectionResponse.getRectangleList().isEmpty()){
// 将绘制图像转换为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);
for(DetectionRectangle detectionRectangle : detectionResponse.getRectangleList()){
g.drawRect(detectionRectangle.getX(), detectionRectangle.getY(), detectionRectangle.getWidth(), detectionRectangle.getHeight());
}
} finally {
g.dispose();
}
}
}

View File

@@ -0,0 +1,119 @@
package cn.smartjavaai.common.utils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDManager;
import org.opencv.core.CvType;
import org.opencv.core.Mat;
import org.opencv.core.Point;
import org.opencv.core.Scalar;
import org.opencv.imgproc.Imgproc;
import java.awt.image.BufferedImage;
import java.awt.image.DataBufferByte;
/**
* OpenCV 工具类
*/
public class OpenCVUtils {
/**
* canny算法边缘检测
*
* @param src
* @return
*/
public static Mat canny(Mat src) {
Mat mat = src.clone();
Imgproc.Canny(src, mat, 100, 200);
return mat;
}
/**
* 画线
*
* @param mat
* @param point1
* @param point2
*/
public static void line(Mat mat, Point point1, Point point2) {
Imgproc.line(mat, point1, point2, new Scalar(255, 255, 255), 1);
}
/**
* NDArray to opencv_core.Mat
*
* @param manager
* @param srcPoints
* @param dstPoints
* @return
*/
public static Mat toOpenCVMat(NDManager manager, NDArray srcPoints, NDArray dstPoints) {
NDArray svdMat = SVDUtils.transformationFromPoints(manager, srcPoints, dstPoints);
double[] doubleArray = svdMat.toDoubleArray();
Mat newSvdMat = new Mat(2, 3, CvType.CV_64F);
for (int i = 0; i < 2; i++) {
for (int j = 0; j < 3; j++) {
newSvdMat.put(i, j, doubleArray[i * 3 + j]);
}
}
return newSvdMat;
}
/**
* double[][] points array to Mat
* @param points
* @return
*/
public static Mat toOpenCVMat(double[][] points) {
Mat mat = new Mat(5, 2, CvType.CV_64F);
for (int i = 0; i < 5; i++) {
for (int j = 0; j < 2; j++) {
mat.put(i, j, points[i * 5 + j]);
}
}
return mat;
}
/**
* 变换矩阵的逆矩阵
*
* @param src
* @return
*/
public static Mat invertAffineTransform(Mat src) {
Mat dst = src.clone();
Imgproc.invertAffineTransform(src, dst);
return dst;
}
/**
* Mat to BufferedImage
*
* @param mat
* @return
*/
public static BufferedImage mat2Image(Mat mat) {
int width = mat.width();
int height = mat.height();
byte[] data = new byte[width * height * (int) mat.elemSize()];
Imgproc.cvtColor(mat, mat, 4);
mat.get(0, 0, data);
BufferedImage ret = new BufferedImage(width, height, 5);
ret.getRaster().setDataElements(0, 0, width, height, data);
return ret;
}
/**
* BufferedImage to Mat
*
* @param img
* @return
*/
public static Mat image2Mat(BufferedImage img) {
int width = img.getWidth();
int height = img.getHeight();
byte[] data = ((DataBufferByte) img.getRaster().getDataBuffer()).getData();
Mat mat = new Mat(height, width, CvType.CV_8UC3);
mat.put(0, 0, data);
return mat;
}
}

View File

@@ -0,0 +1,119 @@
package cn.smartjavaai.common.utils;
import Jama.Matrix;
import Jama.SingularValueDecomposition;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDManager;
/**
* 仿射变换处理工具
*/
public class SVDUtils {
/**
* 计算仿射变换矩阵
* Calculate affine transformation matrix
*
* @param manager
* @param points1
* @param points2
* @return
*/
public static NDArray transformationFromPoints(
NDManager manager, NDArray points1, NDArray points2) {
// 按列计算均值
// Calculate column-wise mean
NDArray c1 = points1.mean(new int[]{0}); // axis=0 列操作 - axis=0 column operation
NDArray c2 = points2.mean(new int[]{0}); // axis=0 列操作 - axis=0 column operation
// 按列减去均值
// Subtract column-wise mean
points1 = points1.sub(c1);
points2 = points2.sub(c2);
// 计算全局标准差
// Calculate global standard deviation
double s1 = std(points1);
double s2 = std(points2);
// 矩阵除以全局标准差
// Matrix divided by global standard deviation
NDArray djl_s1 = manager.create(s1);
NDArray djl_s2 = manager.create(s2);
points1 = points1.div(djl_s1);
points2 = points2.div(djl_s2);
double[] points1D = points1.toDoubleArray();
double[] points2D = points2.toDoubleArray();
// DJL 格式转换成Jamma格式
// Convert DJL format to Jama format
double[][] m1 = new double[5][2];
double[][] m2 = new double[5][2];
for (int i = 0; i < 5; i++) {
for (int j = 0; j < 2; j++) {
m1[i][j] = points1D[i * 2 + j];
}
}
for (int i = 0; i < 5; i++) {
for (int j = 0; j < 2; j++) {
m2[i][j] = points2D[i * 2 + j];
}
}
Matrix p1 = new Matrix(m1);
Matrix p2 = new Matrix(m2);
// 进行奇异值分解
// Perform singular value decomposition
Matrix p3 = p1.transpose().times(p2);
SingularValueDecomposition s = p3.svd();
Matrix U = s.getU();
Matrix S = s.getS();
Matrix V = s.getV();
// TODO 为什么第2列的符号是反的
// Why is the sign of the second column opposite?
m1 = U.getArray();
m1[0][1] = -m1[0][1];
m1[1][1] = -m1[1][1];
m2 = V.getArray();
m2[0][1] = -m2[0][1];
m2[1][1] = -m2[1][1];
Matrix R = (U.times(V)).transpose();
double[][] rArray = R.getArray();
NDArray newR = manager.create(rArray);
// np.vstack([np.hstack(((s2 / s1) * R, c2.T - (s2 / s1) * R * c1.T)), np.matrix([0.,0., 1.])])
// (s2 / s1) * R
NDArray leftPart = djl_s2.div(djl_s1).mul(newR);
// c2.T - (s2 / s1) * R * c1.T)
NDArray rightPart = c2.reshape(2, 1).sub(leftPart.matMul(c1.reshape(2, 1)));
// numpy.hstack(((s2 / s1) * R, c2.T - (s2 / s1) * R * c1.T))
NDArray upPart = leftPart.concat(rightPart, 1);
// np.matrix([0.,0., 1.])
double[] downArray = {0d, 0d, 1d};
NDArray downPart = manager.create(downArray).reshape(1, 3);
NDArray all = upPart.concat(downPart, 0);
// System.out.println("all: " + all);
return upPart;
}
/**
* 计算全局标准差
* Calculate global standard deviation
*
* @param points
* @return
*/
public static double std(NDArray points) {
points = points.square();
double[] doubleResult = points.toDoubleArray();
double std = 0;
for (int i = 0; i < doubleResult.length; i++) {
std = std + doubleResult[i];
}
std = (float) Math.sqrt(std / doubleResult.length);
return std;
}
}

View File

@@ -4,13 +4,13 @@
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd"> xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion> <modelVersion>4.0.0</modelVersion>
<parent> <parent>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId> <artifactId>smartjavaai-parent</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</parent> </parent>
<artifactId>smartjavaai-face</artifactId> <artifactId>smartjavaai-face</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
<name>smartjavaai-face</name> <name>smartjavaai-face</name>
<description>SmartJavaAI</description> <description>SmartJavaAI</description>
<url>https://github.com/geekwenjie/SmartJavaAI</url> <url>https://github.com/geekwenjie/SmartJavaAI</url>
@@ -31,7 +31,7 @@
<dependencies> <dependencies>
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-common</artifactId> <artifactId>smartjavaai-common</artifactId>
<version>${project.version}</version> <version>${project.version}</version>
</dependency> </dependency>
@@ -47,6 +47,7 @@
<artifactId>sqlite-jdbc</artifactId> <artifactId>sqlite-jdbc</artifactId>
<version>3.36.0.3</version> <version>3.36.0.3</version>
</dependency> </dependency>
</dependencies> </dependencies>
<build> <build>

View File

@@ -5,6 +5,7 @@ import cn.smartjavaai.face.entity.FaceResult;
import java.awt.image.BufferedImage; import java.awt.image.BufferedImage;
import java.io.InputStream; import java.io.InputStream;
import java.util.List;
/** /**
* 人脸识别算法 * 人脸识别算法
@@ -46,16 +47,6 @@ public abstract class AbstractFaceModel implements FaceModel {
throw new UnsupportedOperationException("默认不支持该功能"); throw new UnsupportedOperationException("默认不支持该功能");
} }
@Override
public float[] featureExtraction(String imagePath) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] featureExtraction(InputStream inputStream) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override @Override
public float calculSimilar(float[] feature1, float[] feature2) { public float calculSimilar(float[] feature1, float[] feature2) {
throw new UnsupportedOperationException("默认不支持该功能"); throw new UnsupportedOperationException("默认不支持该功能");
@@ -111,15 +102,6 @@ public abstract class AbstractFaceModel implements FaceModel {
throw new UnsupportedOperationException("默认不支持该功能"); throw new UnsupportedOperationException("默认不支持该功能");
} }
@Override
public float[] featureExtraction(BufferedImage sourceImage) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] featureExtraction(byte[] imageData) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override @Override
public float featureComparison(BufferedImage sourceImage1, BufferedImage sourceImag2) { public float featureComparison(BufferedImage sourceImage1, BufferedImage sourceImag2) {
@@ -140,4 +122,64 @@ public abstract class AbstractFaceModel implements FaceModel {
public FaceResult search(byte[] imageData) { public FaceResult search(byte[] imageData) {
throw new UnsupportedOperationException("默认不支持该功能"); throw new UnsupportedOperationException("默认不支持该功能");
} }
@Override
public List<float[]> extractFeatures(String imagePath) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public List<float[]> extractFeatures(byte[] imageData) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public List<float[]> extractFeatures(BufferedImage image) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public List<float[]> extractFeatures(BufferedImage image, FaceExtractConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public List<float[]> extractFeatures(String imagePath, FaceExtractConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public List<float[]> extractFeatures(byte[] imageData, FaceExtractConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] extractTopFaceFeature(BufferedImage image) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] extractTopFaceFeature(String imagePath) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] extractTopFaceFeature(byte[] imageData) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] extractTopFaceFeature(BufferedImage image, FaceExtractConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] extractTopFaceFeature(String imagePath, FaceExtractConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public float[] extractTopFaceFeature(byte[] imageData, FaceExtractConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
} }

View File

@@ -0,0 +1,38 @@
package cn.smartjavaai.face;
import lombok.Data;
/**
* 人脸特征提取配置
* @author dwj
* @date 2025/4/24
*/
@Data
public class FaceExtractConfig {
/**
* 是否裁剪人脸
*/
private boolean cropFace = true;
/**
* 是否对齐人脸
*/
private boolean align = true;
/**
* 人脸检测模型配置
*/
private FaceModelConfig detectModelConfig;
public FaceExtractConfig() {
}
public FaceExtractConfig(boolean cropFace, boolean align, FaceModelConfig detectModelConfig) {
this.cropFace = cropFace;
this.align = align;
this.detectModelConfig = detectModelConfig;
}
}

View File

@@ -5,6 +5,7 @@ import cn.smartjavaai.face.entity.FaceResult;
import java.awt.image.BufferedImage; import java.awt.image.BufferedImage;
import java.io.InputStream; import java.io.InputStream;
import java.util.List;
/** /**
* 人脸识别算法 * 人脸识别算法
@@ -61,34 +62,6 @@ public interface FaceModel {
*/ */
BufferedImage detectAndDraw(BufferedImage sourceImage); BufferedImage detectAndDraw(BufferedImage sourceImage);
/**
* 特征提取
* @param imagePath 图片路径
* @return
*/
float[] featureExtraction(String imagePath);
/**
* 特征提取
* @param inputStream 输入流
* @return
*/
float[] featureExtraction(InputStream inputStream);
/**
* 特征提取
* @param sourceImage BufferedImage图片数据
* @return
*/
float[] featureExtraction(BufferedImage sourceImage);
/**
* 特征提取
* @param imageData 图片字节流
* @return
*/
float[] featureExtraction(byte[] imageData);
/** /**
* 计算相似度 * 计算相似度
* @param feature1 图1特征 * @param feature1 图1特征
@@ -203,4 +176,94 @@ public interface FaceModel {
*/ */
long clearFace(); long clearFace();
/**
* 特征提取(使用默认配置)
* @param imagePath 图片路径
* @return
*/
List<float[]> extractFeatures(String imagePath);
/**
* 特征提取(使用默认配置)
* @param imageData 图片字节流
* @return
*/
List<float[]> extractFeatures(byte[] imageData);
/**
* 特征提取(使用默认配置)
* @param image BufferedImage
* @return
*/
List<float[]> extractFeatures(BufferedImage image);
/**
* 特征提取(使用自定义配置)
* 强制裁剪操作
* @param image BufferedImage
* @param config
* @return
*/
List<float[]> extractFeatures(BufferedImage image, FaceExtractConfig config);
/**
* 特征提取(使用自定义配置)
* @param imagePath 图片路径
* @param config
* @return
*/
List<float[]> extractFeatures(String imagePath, FaceExtractConfig config);
/**
* 特征提取(使用自定义配置)
* @param imageData 图片字节流
* @param config
* @return
*/
List<float[]> extractFeatures(byte[] imageData, FaceExtractConfig config);
/**
* 提取分数最高人脸特征(使用默认配置)
* @param image BufferedImage
* @return
*/
float[] extractTopFaceFeature(BufferedImage image);
/**
* 提取分数最高人脸特征(使用默认配置)
* @param imagePath 图片路径
* @return
*/
float[] extractTopFaceFeature(String imagePath);
/**
* 提取分数最高人脸特征(使用默认配置)
* @param imageData 图片字节流
* @return
*/
float[] extractTopFaceFeature(byte[] imageData);
/**
* 提取分数最高人脸特征(使用自定义配置)
* @param image BufferedImage
* @return
*/
float[] extractTopFaceFeature(BufferedImage image, FaceExtractConfig config);
/**
* 提取分数最高人脸特征(使用自定义配置)
* @param imagePath 图片路径
* @param config
* @return
*/
float[] extractTopFaceFeature(String imagePath, FaceExtractConfig config);
/**
* 提取分数最高人脸特征(使用自定义配置)
* @param imageData 图片字节流
* @param config
* @return
*/
float[] extractTopFaceFeature(byte[] imageData, FaceExtractConfig config);
} }

View File

@@ -50,6 +50,15 @@ public class FaceModelConfig {
*/ */
private int gpuId = 0; private int gpuId = 0;
public FaceModelConfig() {
}
public FaceModelConfig(FaceModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
public FaceModelConfig(FaceModelEnum modelEnum, String modelPath) {
this.modelEnum = modelEnum;
this.modelPath = modelPath;
}
} }

View File

@@ -1,34 +0,0 @@
package cn.smartjavaai.face.factory;
import ai.djl.inference.Predictor;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.face.model.SeetaFace6Model;
import com.seetaface.SeetaFace6JNI;
import org.apache.commons.pool2.BasePooledObjectFactory;
import org.apache.commons.pool2.PooledObject;
import org.apache.commons.pool2.impl.DefaultPooledObject;
/**
* Predictor 工厂类
* @author dwj
* @date 2025/4/8
*/
public class SeetaFace6Factory extends BasePooledObjectFactory<SeetaFace6JNI> {
@Override
public SeetaFace6JNI create() {
return new SeetaFace6JNI();
}
@Override
public PooledObject<SeetaFace6JNI> wrap(SeetaFace6JNI obj) {
return new DefaultPooledObject<>(obj);
}
@Override
public void destroyObject(PooledObject<SeetaFace6JNI> p) {
//p.getObject().dispose(); // 如果需要释放 native 资源
SeetaFace6JNI object = p.getObject();
object = null;
}
}

View File

@@ -5,31 +5,41 @@ import ai.djl.MalformedModelException;
import ai.djl.inference.Predictor; import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image; import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory; import ai.djl.modality.cv.ImageFactory;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDManager;
import ai.djl.opencv.OpenCVImageFactory;
import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel; import ai.djl.repository.zoo.ZooModel;
import ai.djl.training.util.ProgressBar; import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.face.AbstractFaceModel; import cn.smartjavaai.face.*;
import cn.smartjavaai.face.FaceModelConfig;
import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.exception.FaceException;
import cn.smartjavaai.face.translator.FaceFeatureTranslator; import cn.smartjavaai.face.translator.FaceFeatureTranslator;
import cn.smartjavaai.face.utils.FaceAlignUtils;
import cn.smartjavaai.face.utils.FaceUtils;
import cn.smartjavaai.face.utils.OpenCVUtils;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.ObjectPool;
import org.apache.commons.pool2.impl.GenericObjectPool; import org.apache.commons.pool2.impl.GenericObjectPool;
import org.apache.commons.pool2.impl.GenericObjectPoolConfig; import org.apache.commons.pool2.impl.GenericObjectPoolConfig;
import org.opencv.core.Mat;
import javax.imageio.ImageIO; import javax.imageio.ImageIO;
import java.awt.image.BufferedImage; import java.awt.image.BufferedImage;
import java.io.ByteArrayInputStream; import java.io.ByteArrayInputStream;
import java.io.File;
import java.io.IOException; import java.io.IOException;
import java.io.InputStream; import java.io.InputStream;
import java.nio.file.Paths; import java.nio.file.Paths;
import java.time.Duration; import java.time.Duration;
import java.util.ArrayList;
import java.util.Arrays; import java.util.Arrays;
import java.util.List; import java.util.List;
import java.util.Objects; import java.util.Objects;
@@ -117,65 +127,6 @@ public class FeatureExtractionModel extends AbstractFaceModel implements AutoClo
} }
/**
* 特征提取
* @param imagePath 图片路径
* @return
*/
@Override
public float[] featureExtraction(String imagePath) {
if(!FileUtils.isFileExists(imagePath)){
throw new FaceException("图像文件不存在");
}
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
} catch (IOException e) {
throw new FaceException("无效图片", e);
}
return featureExtraction(img);
}
/**
* 特征提取
* @param inputStream 输入流
* @return
*/
@Override
public float[] featureExtraction(InputStream inputStream) {
if(Objects.isNull(inputStream)){
throw new FaceException("图像输入流无效");
}
Image img = null;
try {
img = ImageFactory.getInstance().fromInputStream(inputStream);
} catch (IOException e) {
throw new FaceException("无效图片输入流", e);
}
return featureExtraction(img);
}
@Override
public float[] featureExtraction(BufferedImage sourceImage) {
if(!ImageUtils.isImageValid(sourceImage)){
throw new FaceException("图像无效");
}
Image img = ImageFactory.getInstance().fromImage(sourceImage);
return featureExtraction(img);
}
@Override
public float[] featureExtraction(byte[] imageData) {
if(Objects.isNull(imageData)){
throw new FaceException("图像无效");
}
try {
return featureExtraction(ImageIO.read(new ByteArrayInputStream(imageData)));
} catch (IOException e) {
throw new FaceException("无效图片字节流", e);
}
}
/** /**
* 计算相似度 * 计算相似度
* @param feature1 图1特征 * @param feature1 图1特征
@@ -207,34 +158,19 @@ public class FeatureExtractionModel extends AbstractFaceModel implements AutoClo
if(!FileUtils.isFileExists(imagePath1) || !FileUtils.isFileExists(imagePath2)){ if(!FileUtils.isFileExists(imagePath1) || !FileUtils.isFileExists(imagePath2)){
throw new FaceException("图像文件不存在"); throw new FaceException("图像文件不存在");
} }
float[] feature1 = featureExtraction(imagePath1); float[] feature1 = extractTopFaceFeature(imagePath1);
float[] feature2 = featureExtraction(imagePath2); float[] feature2 = extractTopFaceFeature(imagePath2);
return calculSimilar(feature1, feature2); return calculSimilar(feature1, feature2);
} }
/**
* 特征比较
* @param inputStream1 图1输入流
* @param inputStream2 图2输入流
* @return
*/
@Override
public float featureComparison(InputStream inputStream1, InputStream inputStream2) {
if(Objects.isNull(inputStream1) || Objects.isNull(inputStream2)){
throw new FaceException("图像输入流无效");
}
float[] feature1 = featureExtraction(inputStream1);
float[] feature2 = featureExtraction(inputStream2);
return calculSimilar(feature1, feature2);
}
@Override @Override
public float featureComparison(BufferedImage sourceImage1, BufferedImage sourceImag2) { public float featureComparison(BufferedImage sourceImage1, BufferedImage sourceImag2) {
if(!ImageUtils.isImageValid(sourceImage1) || !ImageUtils.isImageValid(sourceImag2)){ if(!ImageUtils.isImageValid(sourceImage1) || !ImageUtils.isImageValid(sourceImag2)){
throw new FaceException("图像无效"); throw new FaceException("图像无效");
} }
float[] feature1 = featureExtraction(sourceImage1); float[] feature1 = extractTopFaceFeature(sourceImage1);
float[] feature2 = featureExtraction(sourceImag2); float[] feature2 = extractTopFaceFeature(sourceImag2);
return calculSimilar(feature1, feature2); return calculSimilar(feature1, feature2);
} }
@@ -243,11 +179,203 @@ public class FeatureExtractionModel extends AbstractFaceModel implements AutoClo
if(Objects.isNull(imageData1) || Objects.isNull(imageData2)){ if(Objects.isNull(imageData1) || Objects.isNull(imageData2)){
throw new FaceException("图像无效"); throw new FaceException("图像无效");
} }
float[] feature1 = featureExtraction(imageData1); float[] feature1 = extractTopFaceFeature(imageData1);
float[] feature2 = featureExtraction(imageData2); float[] feature2 = extractTopFaceFeature(imageData2);
return calculSimilar(feature1, feature2); return calculSimilar(feature1, feature2);
} }
/**
* 获取默认特征提取配置
* @return
*/
private FaceExtractConfig getDefaultConfig() {
FaceExtractConfig config = new FaceExtractConfig();
FaceModelConfig detectModelConfig = new FaceModelConfig();
detectModelConfig.setModelEnum(FaceModelEnum.ULTRA_LIGHT_FAST_GENERIC_FACE);
config.setDetectModelConfig(detectModelConfig);
return config;
}
@Override
public List<float[]> extractFeatures(String imagePath) {
return extractFeatures(imagePath, getDefaultConfig());
}
@Override
public List<float[]> extractFeatures(byte[] imageData) {
return extractFeatures(imageData, getDefaultConfig());
}
@Override
public List<float[]> extractFeatures(BufferedImage image) {
return extractFeatures(image, getDefaultConfig());
}
@Override
public List<float[]> extractFeatures(BufferedImage image, FaceExtractConfig config) {
if(Objects.isNull(config)){
throw new FaceException("config为null");
}
List<float[]> featureList = new ArrayList<float[]>();
if(Objects.isNull(config.getDetectModelConfig())){
throw new FaceException("config.detectModelConfig为null");
}
FaceModel faceModel = FaceModelFactory.getInstance().getModel(config.getDetectModelConfig());
DetectionResponse detectedResult = faceModel.detect(image);
if(Objects.isNull(detectedResult) || Objects.isNull(detectedResult.getRectangleList()) || detectedResult.getRectangleList().isEmpty()){
throw new FaceException("未检测到人脸");
}
Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
NDManager manager = NDManager.newBaseManager();
for (DetectionRectangle rectangle : detectedResult.getRectangleList()){
float[] features = null;
//裁剪人脸
Image subImage = djlImage.getSubImage(rectangle.getX(), rectangle.getY() , rectangle.getWidth() , rectangle.getHeight());
//人脸对齐
if(config.isAlign()){
//获取子图中人脸关键点坐标
double[][] pointsArray = FaceUtils.facePoints(rectangle.getKeyPoints());
NDArray srcPoints = manager.create(pointsArray);
NDArray dstPoints = FaceUtils.faceTemplate512x512(manager);
// 5点仿射变换
Mat affine_matrix = OpenCVUtils.toOpenCVMat(manager, srcPoints, dstPoints);
/*Mat sourceMat = OpenCVUtils.image2Mat(image);
Mat mat = FaceAlignUtils.warpAffine(sourceMat, affine_matrix);
//OpenCVUtils.mat2Image(mat);
Image alignedImg = ImageFactory.getInstance().fromImage(mat);
features = featureExtraction(alignedImg);*/
Mat mat = FaceAlignUtils.warpAffine((Mat) djlImage.getWrappedImage(), affine_matrix);
Image alignedImg = OpenCVImageFactory.getInstance().fromImage(mat);
features = featureExtraction(alignedImg);
}else{
//不对齐人脸
features = featureExtraction(subImage);
}
if(Objects.nonNull(features)){
featureList.add(features);
}
}
return featureList;
}
@Override
public List<float[]> extractFeatures(String imagePath, FaceExtractConfig config) {
if(!FileUtils.isFileExists(imagePath)){
throw new FaceException("图像文件不存在");
}
// 将图片路径转换为 BufferedImage
BufferedImage image = null;
try {
image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString()));
} catch (IOException e) {
throw new FaceException("无效图片路径", e);
}
return extractFeatures(image, config);
}
@Override
public List<float[]> extractFeatures(byte[] imageData, FaceExtractConfig config) {
if(Objects.isNull(imageData)){
throw new FaceException("图像无效");
}
try {
return extractFeatures(ImageIO.read(new ByteArrayInputStream(imageData)), config);
} catch (IOException e) {
throw new FaceException("错误的图像", e);
}
}
@Override
public float[] extractTopFaceFeature(BufferedImage image) {
return extractTopFaceFeature(image, getDefaultConfig());
}
@Override
public float[] extractTopFaceFeature(String imagePath) {
return extractTopFaceFeature(imagePath, getDefaultConfig());
}
@Override
public float[] extractTopFaceFeature(byte[] imageData) {
return extractTopFaceFeature(imageData, getDefaultConfig());
}
@Override
public float[] extractTopFaceFeature(BufferedImage image, FaceExtractConfig config) {
if(Objects.isNull(config)){
throw new FaceException("config为null");
}
if(Objects.isNull(config.getDetectModelConfig())){
throw new FaceException("config.detectModelConfig为null");
}
Image djlImage = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
float[] features = null;
if(config.isCropFace()){
FaceModel faceModel = FaceModelFactory.getInstance().getModel(config.getDetectModelConfig());
DetectionResponse detectedResult = faceModel.detect(image);
if(Objects.isNull(detectedResult) || Objects.isNull(detectedResult.getRectangleList()) || detectedResult.getRectangleList().isEmpty()){
throw new FaceException("未检测到人脸");
}
//只取第一个人脸
DetectionRectangle rectangle = detectedResult.getRectangleList().get(0);
//裁剪人脸
Image subImage = djlImage.getSubImage(rectangle.getX(), rectangle.getY() , rectangle.getWidth() , rectangle.getHeight());
//人脸对齐
if(config.isAlign()){
NDManager manager = NDManager.newBaseManager();
//获取子图中人脸关键点坐标
double[][] pointsArray = FaceUtils.facePoints(rectangle.getKeyPoints());
NDArray srcPoints = manager.create(pointsArray);
NDArray dstPoints = FaceUtils.faceTemplate512x512(manager);
// 5点仿射变换
Mat affine_matrix = OpenCVUtils.toOpenCVMat(manager, srcPoints, dstPoints);
/*Mat sourceMat = OpenCVUtils.image2Mat(image);
Mat mat = FaceAlignUtils.warpAffine(sourceMat, affine_matrix);
OpenCVUtils.mat2Image(mat);
Image alignedImg = ImageFactory.getInstance().fromImage(OpenCVUtils.mat2Image(mat));
features = featureExtraction(alignedImg);*/
Mat mat = FaceAlignUtils.warpAffine((Mat) djlImage.getWrappedImage(), affine_matrix);
Image alignedImg = OpenCVImageFactory.getInstance().fromImage(mat);
features = featureExtraction(alignedImg);
}else{
//不对齐人脸
features = featureExtraction(subImage);
}
}else{
//不裁剪人脸直接提取特征
features = featureExtraction(djlImage);
}
return features;
}
@Override
public float[] extractTopFaceFeature(String imagePath, FaceExtractConfig config) {
if(!FileUtils.isFileExists(imagePath)){
throw new FaceException("图像文件不存在");
}
// 将图片路径转换为 BufferedImage
BufferedImage image = null;
try {
image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString()));
} catch (IOException e) {
throw new FaceException("无效图片路径", e);
}
return extractTopFaceFeature(image, config);
}
@Override
public float[] extractTopFaceFeature(byte[] imageData, FaceExtractConfig config) {
if(Objects.isNull(imageData)){
throw new FaceException("图像无效");
}
try {
return extractTopFaceFeature(ImageIO.read(new ByteArrayInputStream(imageData)), config);
} catch (IOException e) {
throw new FaceException("错误的图像", e);
}
}
@Override @Override
public void close() { public void close() {
if (predictorPool != null) { if (predictorPool != null) {
@@ -255,4 +383,6 @@ public class FeatureExtractionModel extends AbstractFaceModel implements AutoClo
} }
} }
} }

View File

@@ -19,6 +19,7 @@ import cn.smartjavaai.face.*;
import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.exception.FaceException;
import cn.smartjavaai.face.translator.FaceDetectionTranslator; import cn.smartjavaai.face.translator.FaceDetectionTranslator;
import cn.smartjavaai.face.utils.FaceUtils; import cn.smartjavaai.face.utils.FaceUtils;
import cn.smartjavaai.face.utils.OpenCVUtils;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.ObjectPool;
@@ -144,7 +145,7 @@ public class RetinaFaceModel extends AbstractFaceModel implements AutoCloseable{
if(!ImageUtils.isImageValid(image)){ if(!ImageUtils.isImageValid(image)){
throw new FaceException("图像无效"); throw new FaceException("图像无效");
} }
Image img = ImageFactory.getInstance().fromImage(image); Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
DetectedObjects detection = detect(img); DetectedObjects detection = detect(img);
return FaceUtils.convertToDetectionResponse(detection,img); return FaceUtils.convertToDetectionResponse(detection,img);
} }
@@ -186,7 +187,7 @@ public class RetinaFaceModel extends AbstractFaceModel implements AutoCloseable{
if(!ImageUtils.isImageValid(sourceImage)){ if(!ImageUtils.isImageValid(sourceImage)){
throw new FaceException("图像无效"); throw new FaceException("图像无效");
} }
Image img = ImageFactory.getInstance().fromImage(sourceImage); Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage));
DetectedObjects detectedObjects = detect(img); DetectedObjects detectedObjects = detect(img);
if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){
throw new FaceException("未识别到人脸"); throw new FaceException("未识别到人脸");

View File

@@ -1,29 +1,23 @@
package cn.smartjavaai.face.model; package cn.smartjavaai.face.model;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.face.AbstractFaceModel; import cn.smartjavaai.face.AbstractFaceModel;
import cn.smartjavaai.face.FaceExtractConfig;
import cn.smartjavaai.face.FaceModelConfig; import cn.smartjavaai.face.FaceModelConfig;
import cn.smartjavaai.face.dao.FaceDao; import cn.smartjavaai.face.dao.FaceDao;
import cn.smartjavaai.face.entity.FaceData; import cn.smartjavaai.face.entity.FaceData;
import cn.smartjavaai.face.entity.FaceResult; import cn.smartjavaai.face.entity.FaceResult;
import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.exception.FaceException;
import cn.smartjavaai.face.utils.FaceAlignUtils;
import cn.smartjavaai.face.utils.FaceUtils; import cn.smartjavaai.face.utils.FaceUtils;
import com.seeta.pool.*; import com.seeta.pool.*;
import com.seeta.sdk.*; import com.seeta.sdk.*;
import com.seetaface.NativeLoader; import cn.smartjavaai.face.seetaface.NativeLoader;
import com.seetaface.SeetaFace6JNI;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool;
import org.apache.commons.pool2.impl.GenericObjectPool;
import javax.imageio.ImageIO; import javax.imageio.ImageIO;
import java.awt.image.BufferedImage; import java.awt.image.BufferedImage;
@@ -31,6 +25,7 @@ import java.io.*;
import java.nio.file.Path; import java.nio.file.Path;
import java.nio.file.Paths; import java.nio.file.Paths;
import java.sql.SQLException; import java.sql.SQLException;
import java.util.ArrayList;
import java.util.List; import java.util.List;
import java.util.Objects; import java.util.Objects;
@@ -154,10 +149,19 @@ public class SeetaFace6Model extends AbstractFaceModel {
SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3);
imageData.data = ImageUtils.getMatrixBGR(image); imageData.data = ImageUtils.getMatrixBGR(image);
FaceDetector predictor = null; FaceDetector predictor = null;
FaceLandmarker faceLandmarker = null;
try { try {
predictor = faceDetectorPool.borrowObject(); predictor = faceDetectorPool.borrowObject();
faceLandmarker = faceLandmarkerPool.borrowObject();
SeetaRect[] seetaResult = predictor.Detect(imageData); SeetaRect[] seetaResult = predictor.Detect(imageData);
return FaceUtils.convertToDetectionResponse(seetaResult, config); List<SeetaPointF[]> seetaPointFSList = new ArrayList<SeetaPointF[]>();
for(SeetaRect seetaRect : seetaResult){
//提取人脸的5点人脸标识
SeetaPointF[] pointFS = new SeetaPointF[faceLandmarker.number()];
faceLandmarker.mark(imageData, seetaRect, pointFS);
seetaPointFSList.add(pointFS);
}
return FaceUtils.convertToDetectionResponse(seetaResult, config, seetaPointFSList);
} catch (Exception e) { } catch (Exception e) {
throw new FaceException("目标检测错误", e); throw new FaceException("目标检测错误", e);
}finally { }finally {
@@ -168,6 +172,13 @@ public class SeetaFace6Model extends AbstractFaceModel {
log.warn("归还Predictor失败", e); log.warn("归还Predictor失败", e);
} }
} }
if (faceLandmarker != null) {
try {
faceLandmarkerPool.returnObject(faceLandmarker); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
}
}
} }
} }
@@ -295,104 +306,6 @@ public class SeetaFace6Model extends AbstractFaceModel {
} }
} }
@Override
public float[] featureExtraction(BufferedImage image) {
if(!ImageUtils.isImageValid(image)){
throw new FaceException("图像无效");
}
SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3);
imageData.data = ImageUtils.getMatrixBGR(image);
FaceDetector faceDetector = null;
FaceLandmarker faceLandmarker = null;
FaceRecognizer faceRecognizer = null;
try {
faceDetector = faceDetectorPool.borrowObject();
faceLandmarker = faceLandmarkerPool.borrowObject();
faceRecognizer = faceRecognizerPool.borrowObject();
//检测人脸
SeetaRect[] seetaResult = faceDetector.Detect(imageData);
if(Objects.isNull(seetaResult) || seetaResult.length == 0){
throw new FaceException("未检测到人脸");
}
//提取第一个人脸的5点人脸标识
SeetaPointF[] pointFS = new SeetaPointF[faceLandmarker.number()];
faceLandmarker.mark(imageData, seetaResult[0], pointFS);
//提取特征
float[] features = new float[faceRecognizer.GetExtractFeatureSize()];
boolean isSuccess = faceRecognizer.Extract(imageData, pointFS, features);
if(!isSuccess){
throw new FaceException("人脸特征提取失败");
}
return 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 float[] featureExtraction(String imagePath) {
if(!FileUtils.isFileExists(imagePath)){
throw new FaceException("图像文件不存在");
}
BufferedImage image = null;
try {
image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString()));
} catch (IOException e) {
throw new FaceException("无效图片路径", e);
}
return featureExtraction(image);
}
@Override
public float[] featureExtraction(InputStream inputStream) {
if(Objects.isNull(inputStream)){
throw new FaceException("图像输入流无效");
}
BufferedImage image = null;
try {
image = ImageIO.read(inputStream);
} catch (IOException e) {
throw new FaceException("无效图片输入流", e);
}
return featureExtraction(image);
}
@Override
public float[] featureExtraction(byte[] imageData) {
if(Objects.isNull(imageData)){
throw new FaceException("图像无效");
}
try {
return featureExtraction(ImageIO.read(new ByteArrayInputStream(imageData)));
} catch (IOException e) {
throw new FaceException("错误的图像", e);
}
}
@Override @Override
public float calculSimilar(float[] feature1, float[] feature2) { public float calculSimilar(float[] feature1, float[] feature2) {
@@ -897,4 +810,185 @@ public class SeetaFace6Model extends AbstractFaceModel {
pageNo++; pageNo++;
} }
} }
@Override
public List<float[]> extractFeatures(String imagePath) {
if(!FileUtils.isFileExists(imagePath)){
throw new FaceException("图像文件不存在");
}
// 将图片路径转换为 BufferedImage
BufferedImage image = null;
try {
image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString()));
} catch (IOException e) {
throw new FaceException("无效图片路径", e);
}
return extractFeatures(image);
}
@Override
public List<float[]> extractFeatures(byte[] imageData) {
if(Objects.isNull(imageData)){
throw new FaceException("图像无效");
}
try {
return extractFeatures(ImageIO.read(new ByteArrayInputStream(imageData)));
} catch (IOException e) {
throw new FaceException("错误的图像", e);
}
}
@Override
public List<float[]> extractFeatures(BufferedImage image) {
if(!ImageUtils.isImageValid(image)){
throw new FaceException("图像无效");
}
List<float[]> featureList = new ArrayList<float[]>();
SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3);
imageData.data = ImageUtils.getMatrixBGR(image);
FaceDetector faceDetector = null;
FaceLandmarker faceLandmarker = null;
FaceRecognizer faceRecognizer = null;
try {
faceDetector = faceDetectorPool.borrowObject();
faceLandmarker = faceLandmarkerPool.borrowObject();
faceRecognizer = faceRecognizerPool.borrowObject();
//检测人脸
SeetaRect[] seetaResult = faceDetector.Detect(imageData);
if(Objects.isNull(seetaResult) || seetaResult.length == 0){
throw new FaceException("未检测到人脸");
}
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){
throw new FaceException("人脸特征提取失败");
}
featureList.add(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);
}
}
}
return featureList;
}
@Override
public float[] extractTopFaceFeature(BufferedImage image) {
if(!ImageUtils.isImageValid(image)){
throw new FaceException("图像无效");
}
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 {
faceDetector = faceDetectorPool.borrowObject();
faceLandmarker = faceLandmarkerPool.borrowObject();
faceRecognizer = faceRecognizerPool.borrowObject();
//检测人脸
SeetaRect[] seetaResult = faceDetector.Detect(imageData);
if(Objects.isNull(seetaResult) || seetaResult.length == 0){
throw new FaceException("未检测到人脸");
}
//提取人脸的5点人脸标识
SeetaPointF[] pointFS = new SeetaPointF[faceLandmarker.number()];
faceLandmarker.mark(imageData, seetaResult[0], pointFS);
//提取特征
features = new float[faceRecognizer.GetExtractFeatureSize()];
//CropFaceV2 + ExtractCroppedFace 已包含裁剪+人脸对齐
boolean isSuccess = faceRecognizer.Extract(imageData, pointFS, features);
if(!isSuccess){
throw new FaceException("人脸特征提取失败");
}
} 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);
}
}
}
return features;
}
@Override
public float[] extractTopFaceFeature(String imagePath) {
if(!FileUtils.isFileExists(imagePath)){
throw new FaceException("图像文件不存在");
}
// 将图片路径转换为 BufferedImage
BufferedImage image = null;
try {
image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString()));
} catch (IOException e) {
throw new FaceException("无效图片路径", e);
}
return extractTopFaceFeature(image);
}
@Override
public float[] extractTopFaceFeature(byte[] imageData) {
if(Objects.isNull(imageData)){
throw new FaceException("图像无效");
}
try {
return extractTopFaceFeature(ImageIO.read(new ByteArrayInputStream(imageData)));
} catch (IOException e) {
throw new FaceException("错误的图像", e);
}
}
} }

View File

@@ -19,6 +19,7 @@ import cn.smartjavaai.face.*;
import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.exception.FaceException;
import cn.smartjavaai.face.translator.FaceDetectionTranslator; import cn.smartjavaai.face.translator.FaceDetectionTranslator;
import cn.smartjavaai.face.utils.FaceUtils; import cn.smartjavaai.face.utils.FaceUtils;
import cn.smartjavaai.face.utils.OpenCVUtils;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool; import org.apache.commons.pool2.ObjectPool;
@@ -133,7 +134,7 @@ public class UltraLightFastGenericFaceModel extends AbstractFaceModel implements
@Override @Override
public DetectionResponse detect(BufferedImage image) { public DetectionResponse detect(BufferedImage image) {
Image img = ImageFactory.getInstance().fromImage(image); Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
DetectedObjects detection = detect(img); DetectedObjects detection = detect(img);
return FaceUtils.convertToDetectionResponse(detection,img); return FaceUtils.convertToDetectionResponse(detection,img);
} }
@@ -175,7 +176,7 @@ public class UltraLightFastGenericFaceModel extends AbstractFaceModel implements
if(!ImageUtils.isImageValid(sourceImage)){ if(!ImageUtils.isImageValid(sourceImage)){
throw new FaceException("图像无效"); throw new FaceException("图像无效");
} }
Image img = ImageFactory.getInstance().fromImage(sourceImage); Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage));
DetectedObjects detectedObjects = detect(img); DetectedObjects detectedObjects = detect(img);
if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){ if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){
throw new FaceException("未识别到人脸"); throw new FaceException("未识别到人脸");
@@ -222,6 +223,9 @@ public class UltraLightFastGenericFaceModel extends AbstractFaceModel implements
} }
@Override @Override
public void close() { public void close() {
if (predictorPool != null) { if (predictorPool != null) {

View File

@@ -1,4 +1,4 @@
package com.seetaface; package cn.smartjavaai.face.seetaface;
import cn.hutool.core.io.FileUtil; import cn.hutool.core.io.FileUtil;
@@ -9,18 +9,11 @@ import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.common.enums.DeviceEnum; import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.face.FaceModelConfig; import cn.smartjavaai.face.FaceModelConfig;
import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.exception.FaceException;
import com.seeta.sdk.SeetaDevice;
import com.seeta.sdk.util.DllItem; import com.seeta.sdk.util.DllItem;
import com.seeta.sdk.util.LoadNativeCore; import com.seeta.sdk.util.LoadNativeCore;
import jdk.dynalink.linker.support.Lookup;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils;
import java.io.*; import java.io.*;
import java.lang.invoke.MethodHandles;
import java.lang.invoke.VarHandle;
import java.lang.reflect.Field;
import java.lang.reflect.Method;
import java.nio.file.Files; import java.nio.file.Files;
import java.nio.file.Path; import java.nio.file.Path;
import java.nio.file.Paths; import java.nio.file.Paths;
@@ -37,18 +30,9 @@ public class NativeLoader {
private static Path seetaface6NativePath; private static Path seetaface6NativePath;
private static final String[] WIN_LIBS = {"tennis.dll","tennis_haswell.dll","tennis_pentium.dll","tennis_sandy_bridge.dll","SeetaAuthorize.dll","SeetaFaceAntiSpoofingX600.dll","SeetaFaceDetector600.dll","SeetaFaceLandmarker600.dll","SeetaFaceRecognizer610.dll","SeetaFace6JNI.dll"};
//private static final String[] WIN_LIBS = {"tennis","tennis_haswell","tennis_pentium","tennis_sandy_bridge","SeetaAuthorize","SeetaFaceAntiSpoofingX600","SeetaFaceDetector600","SeetaFaceLandmarker600","SeetaFaceRecognizer610","SeetaFace6JNI"};
private static final String[] LINUX_CENTOS_LIBS = {"libSeetaAuthorize.so","libtennis.so","libtennis_haswell.so","libtennis_pentium.so","libtennis_sandy_bridge.so","libSeetaFaceDetector600.so","libSeetaAgePredictor600.so","libSeetaEyeStateDetector200.so","libSeetaFaceAntiSpoofingX600.so","libSeetaFaceLandmarker600.so","libSeetaFaceRecognizer610.so","libSeetaGenderPredictor600.so","libSeetaMaskDetector200.so","libSeetaPoseEstimation600.so","libSeetaFaceTracking600.so","libSeetaQualityAssessor300.so"};
private static final String[] LINUX_UBUNTU_LIBS = {"libSeetaAuthorize.so","libtennis.so","libtennis_haswell.so","libtennis_pentium.so","libtennis_sandy_bridge.so","libSeetaFaceDetector600.so","libSeetaAgePredictor600.so","libSeetaEyeStateDetector200.so","libSeetaFaceAntiSpoofingX600.so","libSeetaFaceLandmarker600.so","libSeetaFaceRecognizer610.so","libSeetaGenderPredictor600.so","libSeetaMaskDetector200.so","libSeetaPoseEstimation600.so","libSeetaFaceTracking600.so","libSeetaQualityAssessor300.so"};
private static final String SEETAFACE_LIB_DIR = "seetaface6"; private static final String SEETAFACE_LIB_DIR = "seetaface6";
public static SeetaFace6JNI seetaFace6SDK;
public static final String AMD64 = "amd64";
public static final String x86_64 = "amd64";
/** /**
* 定义dll 路径和加载顺序的文件 * 定义dll 路径和加载顺序的文件
@@ -204,7 +188,7 @@ public class NativeLoader {
*/ */
private static File extractLibrary(String libPath) throws IOException { private static File extractLibrary(String libPath) throws IOException {
String resourcePath = libPath; String resourcePath = libPath;
try (InputStream in = com.seetaface.NativeLoader.class.getResourceAsStream(resourcePath)) { try (InputStream in = NativeLoader.class.getResourceAsStream(resourcePath)) {
if (in == null) throw new FileNotFoundException(resourcePath); if (in == null) throw new FileNotFoundException(resourcePath);
Path path = Paths.get(resourcePath); Path path = Paths.get(resourcePath);
String fileName = path.getFileName().toString(); String fileName = path.getFileName().toString();

View File

@@ -0,0 +1,68 @@
package cn.smartjavaai.face.utils;
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.ImageUtils;
import com.seeta.sdk.SeetaImageData;
import com.seeta.sdk.SeetaPointF;
import org.opencv.core.Mat;
import org.opencv.core.Scalar;
import org.opencv.core.Size;
import org.opencv.imgproc.Imgproc;
import java.awt.image.BufferedImage;
/**
* 人脸对齐
* @author Calvin
*/
public class FaceAlignUtils {
/**
* 根据目标点,进行旋转仿射变换
* Perform rotation and affine transformation based on the target 5 points
*
* @param src
* @param rot_mat
* @return
*/
public static Mat warpAffine(Mat src, Mat rot_mat) {
Mat rot = new Mat();
// 进行仿射变换变换后大小为src的大小
// Perform affine transformation, the size after transformation is the same as the size of src
Scalar scalar = new Scalar(135, 133, 132);
Size size = new Size(512, 512);
Imgproc.warpAffine(src, rot, rot_mat, size, 0, 0, scalar);
return rot;
}
public static Mat warpAffine(Mat src, Mat rot_mat, int width, int height) {
Mat rot = new Mat();
Size size = new Size(width, height);
Imgproc.warpAffine(src, rot, rot_mat, size);
return rot;
}
public static Mat warpAffine(Mat src, Mat rot_mat, int width, int height, int flags) {
Mat rot = new Mat();
Size size = new Size(width, height);
Imgproc.warpAffine(src, rot, rot_mat, size, flags);
return rot;
}
public static SeetaImageData faceAlign(BufferedImage sourceImage, SeetaPointF[] pointFS) {
NDManager manager = NDManager.newBaseManager();
//获取子图中人脸关键点坐标
double[][] pointsArray = FaceUtils.facePoints(pointFS);
NDArray srcPoints = manager.create(pointsArray);
NDArray dstPoints = FaceUtils.faceTemplate512x512(manager);
// 5点仿射变换
Mat affine_matrix = OpenCVUtils.toOpenCVMat(manager, srcPoints, dstPoints);
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);
return imageData;
}
}

View File

@@ -3,12 +3,16 @@ package cn.smartjavaai.face.utils;
import ai.djl.modality.cv.Image; import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.BoundingBox; import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects; import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDManager;
import cn.smartjavaai.common.entity.DetectionResponse; import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.Point;
import cn.smartjavaai.common.entity.DetectionRectangle; import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.Point;
import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.face.FaceModelConfig; import cn.smartjavaai.face.FaceModelConfig;
import cn.smartjavaai.face.exception.FaceException; import cn.smartjavaai.face.exception.FaceException;
import com.seeta.sdk.SeetaImageData;
import com.seeta.sdk.SeetaPointF;
import com.seeta.sdk.SeetaRect; import com.seeta.sdk.SeetaRect;
import javax.imageio.ImageIO; import javax.imageio.ImageIO;
@@ -16,10 +20,8 @@ import java.awt.*;
import java.awt.image.BufferedImage; import java.awt.image.BufferedImage;
import java.io.File; import java.io.File;
import java.io.IOException; import java.io.IOException;
import java.util.ArrayList; import java.util.*;
import java.util.Iterator;
import java.util.List; import java.util.List;
import java.util.Objects;
import java.util.stream.Collectors; import java.util.stream.Collectors;
/** /**
@@ -49,11 +51,21 @@ public class FaceUtils {
while(iterator.hasNext()) { while(iterator.hasNext()) {
DetectedObjects.DetectedObject result = (DetectedObjects.DetectedObject)iterator.next(); DetectedObjects.DetectedObject result = (DetectedObjects.DetectedObject)iterator.next();
BoundingBox box = result.getBoundingBox(); BoundingBox box = result.getBoundingBox();
int x = (int)(box.getBounds().getX() * (double)img.getWidth()); List<Point> keyPoints = new ArrayList<Point>();
int y = (int)(box.getBounds().getY() * (double)img.getHeight()); box.getBounds().getPath().forEach(point -> {
int width = (int)(box.getBounds().getWidth() * (double)img.getWidth()); keyPoints.add(new Point(point.getX(), point.getY()));
int height = (int)(box.getBounds().getHeight() * (double)img.getHeight()); });
int x = (int)(box.getBounds().getX() * img.getWidth());
int y = (int)(box.getBounds().getY() * img.getHeight());
int width = (int)(box.getBounds().getWidth() * img.getWidth());
int height = (int)(box.getBounds().getHeight() * img.getHeight());
// 修正边界,防止越界
if (x < 0) x = 0;
if (y < 0) y = 0;
if (x + width > img.getWidth()) width = img.getWidth() - x;
if (y + height > img.getHeight()) height = img.getHeight() - y;
DetectionRectangle rectangle = new DetectionRectangle(x, y, width, height, detection.getProbabilities().get(index).floatValue()); DetectionRectangle rectangle = new DetectionRectangle(x, y, width, height, detection.getProbabilities().get(index).floatValue());
rectangle.setKeyPoints(keyPoints);
rectangleList.add(rectangle); rectangleList.add(rectangle);
index++; index++;
} }
@@ -66,18 +78,24 @@ public class FaceUtils {
* @param seetaResult * @param seetaResult
* @return * @return
*/ */
public static DetectionResponse convertToDetectionResponse(SeetaRect[] seetaResult, FaceModelConfig config){ public static DetectionResponse convertToDetectionResponse(SeetaRect[] seetaResult, FaceModelConfig config,List<SeetaPointF[]> seetaPointFSList){
if(Objects.isNull(seetaResult) || seetaResult.length == 0){ if(Objects.isNull(seetaResult) || seetaResult.length == 0){
return null; return null;
} }
DetectionResponse detectionResponse = new DetectionResponse(); DetectionResponse detectionResponse = new DetectionResponse();
List<DetectionRectangle> rectangleList = new ArrayList<DetectionRectangle>(); List<DetectionRectangle> rectangleList = new ArrayList<DetectionRectangle>();
for(SeetaRect rect : seetaResult){ for(int i = 0; i < seetaResult.length; i++){
SeetaRect rect = seetaResult[i];
SeetaPointF[] seetaPointFS = seetaPointFSList.get(i);
//过滤置信度 //过滤置信度
/*if(config.getConfidenceThreshold() > 0){ /*if(config.getConfidenceThreshold() > 0){
continue; continue;
}*/ }*/
DetectionRectangle rectangle = new DetectionRectangle(rect.x, rect.y, rect.width, rect.height, 0); DetectionRectangle rectangle = new DetectionRectangle(rect.x, rect.y, rect.width, rect.height, 0);
List<Point> keyPoints = Arrays.stream(seetaPointFS)
.map(p -> new Point(p.x, p.y))
.collect(Collectors.toList());
rectangle.setKeyPoints(keyPoints);
rectangleList.add(rectangle); rectangleList.add(rectangle);
} }
detectionResponse.setRectangleList(rectangleList); detectionResponse.setRectangleList(rectangleList);
@@ -108,6 +126,10 @@ public class FaceUtils {
graphics.setColor(Color.RED);// 边框颜色 graphics.setColor(Color.RED);// 边框颜色
graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight());
drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4); drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4);
//绘制人脸关键点
if(rectangle.getKeyPoints() != null){
drawLandmarks(graphics, rectangle.getKeyPoints());
}
} }
graphics.dispose(); graphics.dispose();
ImageIO.write(sourceImage, "jpg", new File(savePath)); ImageIO.write(sourceImage, "jpg", new File(savePath));
@@ -136,6 +158,10 @@ public class FaceUtils {
graphics.setColor(Color.RED);// 边框颜色 graphics.setColor(Color.RED);// 边框颜色
graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight()); graphics.drawRect(rectangle.getX(), rectangle.getY(), rectangle.getWidth(), rectangle.getHeight());
drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4); drawText(graphics, "face", rectangle.getX(), rectangle.getY(), stroke, 4);
//绘制人脸关键点
if(rectangle.getKeyPoints() != null){
drawLandmarks(graphics, rectangle.getKeyPoints());
}
} }
graphics.dispose(); graphics.dispose();
return sourceImage; return sourceImage;
@@ -163,5 +189,124 @@ public class FaceUtils {
g.drawString(text, x + padding, y + ascent); g.drawString(text, x + padding, y + ascent);
} }
/**
* 修正检测框
* @param rectangle
* @param imageWidth
* @param imageHeight
* @return
*/
public static DetectionRectangle correctRect(DetectionRectangle rectangle, int imageWidth, int imageHeight) {
int x = rectangle.getX();
int y = rectangle.getY();
int width = rectangle.getWidth();
int height = rectangle.getHeight();
// 修正x, y防止越界
if (x < 0) x = 0;
if (y < 0) y = 0;
// 宽高不能超出图片范围
if (x + width > imageWidth) {
width = imageWidth - x;
}
if (y + height > imageHeight) {
height = imageHeight - y;
}
// 防止最终 width 或 height 为负或为 0
if (width <= 0 || height <= 0) {
return null; // 无效区域
}
return new DetectionRectangle(x, y, width, height, rectangle.score);
}
/**
* 子图中人脸关键点坐标 - Coordinates of key points in the image
*
* @param points
* @return
*/
public static double[][] facePoints(List<Point> points) {
// 图中关键点坐标 - 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 (Point point : points) {
pointsArray[i][0] = point.getX();
pointsArray[i][1] = point.getY();
i++;
}
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
* standard 5 landmarks for FFHQ faces with 512 x 512
*
* @param manager
* @return
*/
public static NDArray faceTemplate512x512(NDManager manager) {
double[][] coord5point = {
{192.98138, 239.94708}, // 512x512的目标点 - Target point of 512x512
{318.90277, 240.1936},
{256.63416, 314.01935},
{201.26117, 371.41043},
{313.08905, 371.15118}
};
NDArray points = manager.create(coord5point);
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<Point> 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);
}
}
} }

View File

@@ -0,0 +1,119 @@
package cn.smartjavaai.face.utils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDManager;
import org.opencv.core.CvType;
import org.opencv.core.Mat;
import org.opencv.core.Point;
import org.opencv.core.Scalar;
import org.opencv.imgproc.Imgproc;
import java.awt.image.BufferedImage;
import java.awt.image.DataBufferByte;
/**
* OpenCV 工具类
*/
public class OpenCVUtils {
/**
* canny算法边缘检测
*
* @param src
* @return
*/
public static Mat canny(Mat src) {
Mat mat = src.clone();
Imgproc.Canny(src, mat, 100, 200);
return mat;
}
/**
* 画线
*
* @param mat
* @param point1
* @param point2
*/
public static void line(Mat mat, Point point1, Point point2) {
Imgproc.line(mat, point1, point2, new Scalar(255, 255, 255), 1);
}
/**
* NDArray to opencv_core.Mat
*
* @param manager
* @param srcPoints
* @param dstPoints
* @return
*/
public static Mat toOpenCVMat(NDManager manager, NDArray srcPoints, NDArray dstPoints) {
NDArray svdMat = SVDUtils.transformationFromPoints(manager, srcPoints, dstPoints);
double[] doubleArray = svdMat.toDoubleArray();
Mat newSvdMat = new Mat(2, 3, CvType.CV_64F);
for (int i = 0; i < 2; i++) {
for (int j = 0; j < 3; j++) {
newSvdMat.put(i, j, doubleArray[i * 3 + j]);
}
}
return newSvdMat;
}
/**
* double[][] points array to Mat
* @param points
* @return
*/
public static Mat toOpenCVMat(double[][] points) {
Mat mat = new Mat(5, 2, CvType.CV_64F);
for (int i = 0; i < 5; i++) {
for (int j = 0; j < 2; j++) {
mat.put(i, j, points[i * 5 + j]);
}
}
return mat;
}
/**
* 变换矩阵的逆矩阵
*
* @param src
* @return
*/
public static Mat invertAffineTransform(Mat src) {
Mat dst = src.clone();
Imgproc.invertAffineTransform(src, dst);
return dst;
}
/**
* Mat to BufferedImage
*
* @param mat
* @return
*/
public static BufferedImage mat2Image(Mat mat) {
int width = mat.width();
int height = mat.height();
byte[] data = new byte[width * height * (int) mat.elemSize()];
Imgproc.cvtColor(mat, mat, 4);
mat.get(0, 0, data);
BufferedImage ret = new BufferedImage(width, height, 5);
ret.getRaster().setDataElements(0, 0, width, height, data);
return ret;
}
/**
* BufferedImage to Mat
*
* @param img
* @return
*/
public static Mat image2Mat(BufferedImage img) {
int width = img.getWidth();
int height = img.getHeight();
byte[] data = ((DataBufferByte) img.getRaster().getDataBuffer()).getData();
Mat mat = new Mat(height, width, CvType.CV_8UC3);
mat.put(0, 0, data);
return mat;
}
}

View File

@@ -0,0 +1,119 @@
package cn.smartjavaai.face.utils;
import Jama.Matrix;
import Jama.SingularValueDecomposition;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDManager;
/**
* 仿射变换处理工具
*/
public class SVDUtils {
/**
* 计算仿射变换矩阵
* Calculate affine transformation matrix
*
* @param manager
* @param points1
* @param points2
* @return
*/
public static NDArray transformationFromPoints(
NDManager manager, NDArray points1, NDArray points2) {
// 按列计算均值
// Calculate column-wise mean
NDArray c1 = points1.mean(new int[]{0}); // axis=0 列操作 - axis=0 column operation
NDArray c2 = points2.mean(new int[]{0}); // axis=0 列操作 - axis=0 column operation
// 按列减去均值
// Subtract column-wise mean
points1 = points1.sub(c1);
points2 = points2.sub(c2);
// 计算全局标准差
// Calculate global standard deviation
double s1 = std(points1);
double s2 = std(points2);
// 矩阵除以全局标准差
// Matrix divided by global standard deviation
NDArray djl_s1 = manager.create(s1);
NDArray djl_s2 = manager.create(s2);
points1 = points1.div(djl_s1);
points2 = points2.div(djl_s2);
double[] points1D = points1.toDoubleArray();
double[] points2D = points2.toDoubleArray();
// DJL 格式转换成Jamma格式
// Convert DJL format to Jama format
double[][] m1 = new double[5][2];
double[][] m2 = new double[5][2];
for (int i = 0; i < 5; i++) {
for (int j = 0; j < 2; j++) {
m1[i][j] = points1D[i * 2 + j];
}
}
for (int i = 0; i < 5; i++) {
for (int j = 0; j < 2; j++) {
m2[i][j] = points2D[i * 2 + j];
}
}
Matrix p1 = new Matrix(m1);
Matrix p2 = new Matrix(m2);
// 进行奇异值分解
// Perform singular value decomposition
Matrix p3 = p1.transpose().times(p2);
SingularValueDecomposition s = p3.svd();
Matrix U = s.getU();
Matrix S = s.getS();
Matrix V = s.getV();
// TODO 为什么第2列的符号是反的
// Why is the sign of the second column opposite?
m1 = U.getArray();
m1[0][1] = -m1[0][1];
m1[1][1] = -m1[1][1];
m2 = V.getArray();
m2[0][1] = -m2[0][1];
m2[1][1] = -m2[1][1];
Matrix R = (U.times(V)).transpose();
double[][] rArray = R.getArray();
NDArray newR = manager.create(rArray);
// np.vstack([np.hstack(((s2 / s1) * R, c2.T - (s2 / s1) * R * c1.T)), np.matrix([0.,0., 1.])])
// (s2 / s1) * R
NDArray leftPart = djl_s2.div(djl_s1).mul(newR);
// c2.T - (s2 / s1) * R * c1.T)
NDArray rightPart = c2.reshape(2, 1).sub(leftPart.matMul(c1.reshape(2, 1)));
// numpy.hstack(((s2 / s1) * R, c2.T - (s2 / s1) * R * c1.T))
NDArray upPart = leftPart.concat(rightPart, 1);
// np.matrix([0.,0., 1.])
double[] downArray = {0d, 0d, 1d};
NDArray downPart = manager.create(downArray).reshape(1, 3);
NDArray all = upPart.concat(downPart, 0);
// System.out.println("all: " + all);
return upPart;
}
/**
* 计算全局标准差
* Calculate global standard deviation
*
* @param points
* @return
*/
public static double std(NDArray points) {
points = points.square();
double[] doubleResult = points.toDoubleArray();
double std = 0;
for (int i = 0; i < doubleResult.length; i++) {
std = std + doubleResult[i];
}
std = (float) Math.sqrt(std / doubleResult.length);
return std;
}
}

View File

@@ -1,135 +0,0 @@
package com.seetaface;
import com.seetaface.model.RecognizeResult;
import com.seetaface.model.SeetaImageData;
import com.seetaface.model.SeetaPointF;
import com.seetaface.model.SeetaRect;
/**
* seetaface6 sdk
* @author dwj
*/
public class SeetaFace6JNI {
/**
* 初始化,指定人脸识别模型文件目录
*
* @param modelDir
* @return
*/
public native boolean initModel(String modelDir);
/**
* 检测人脸
*
* @param img
* @return
*/
public native SeetaRect[] detect(SeetaImageData img);
/**
* 根据人脸检测关键点
* 关键定定位输入的是原始图片和人脸检测结果,给出指定人脸上的关键点的依次坐标。
* 这里检测到的5点坐标循序依次为左眼中心、右眼中心、鼻尖、左嘴角和右嘴角。
* 注意这里的左右是基于图片内容的左右,并不是图片中人的左右,即左眼中心就是图片中左边的眼睛的中心。
*
* @param img
* @param faces
* @return
*/
public native SeetaPointF[] mark(SeetaImageData img, SeetaRect faces);
/**
* 1 v 1 人脸比对
*
* @param img1
* @param img2
* @return 相似度范围在0~1,返回负数表示出错
*/
public native float compare(SeetaImageData img1, SeetaImageData img2);
/**
* 提取人脸区域特性
* @param face crop方法返回的人脸图像
* @return
*/
public native float[] extractCroppedFace(byte[] face);
/**
* 提取一个图像中最大人脸的特征
* @param img
* @return
*/
public native float[] extractMaxFace(SeetaImageData img);
/**
* 计算两个特性的相似度
* @param features1
* @param features2
* @return
*/
public native float calculateSimilarity(float[] features1, float[] features2);
/**
* 注册人脸
*
* @param img
* @return The returned value is the index of face database. Reture -1 if failed
*/
public native long register(SeetaImageData img);
/**
* 注册裁剪后的人脸,推荐使用该方法
* @param bytes
* @return
*/
public native long registerCroppedFace(byte[] bytes);
/**
* 从人脸库中搜索,返回相似度最高的索引
*
* @param img
* @return index saves the index of face databese, which is same as the retured value by Register. similar saves the most similar.
*/
public native RecognizeResult query(SeetaImageData img);
/**
* 用裁剪后的人脸进行搜索
* @param bytes
* @return
*/
public native RecognizeResult queryByCroppedFace(byte[] bytes);
/**
* 将人脸从数据库中删除
* @param index -1: 删除所有
* @return 返回删除记录数
*/
public native long delete(long[] index);
/**
* 人脸提取
*
* @param img
* @return The returned value is face data. Reture null if failed
*/
public native byte[][] crop(SeetaImageData img);
/**
* 图片活体检测
* @param img
* @return
*/
public native int predictImage(SeetaImageData img);
public native void dispose();
@Override
protected void finalize() throws Throwable {
super.finalize();
this.dispose();
}
}

View File

@@ -1,10 +0,0 @@
package com.seetaface.model;
/**
* 人脸识别结果
*/
public class RecognizeResult {
public int index;
public float similar;
}

View File

@@ -1,27 +0,0 @@
package com.seetaface.model;
/**
* 人脸识别参数
* @author dwj
*/
public class SeetaImageData {
public SeetaImageData() {
}
public SeetaImageData(int width, int height, int channels) {
this.data = new byte[width * height * channels];
this.width = width;
this.height = height;
this.channels = channels;
}
public SeetaImageData(int width, int height) {
this(width, height, 3);
}
public byte[] data;
public int width;
public int height;
public int channels;
}

View File

@@ -1,18 +0,0 @@
package com.seetaface.model;
/**
* SeetaPointF
* @author dwj
*/
public class SeetaPointF {
public double x;
public double y;
@Override
public String toString() {
return "{" +
"x=" + x +
", y=" + y +
'}';
}
}

View File

@@ -1,24 +0,0 @@
package com.seetaface.model;
/**
* SeetaPointF
* @author dwj
*/
public class SeetaRect {
public int x;
public int y;
public int width;
public int height;
public float score;
@Override
public String toString() {
return "SeetaRect{" +
"x=" + x +
", y=" + y +
", width=" + width +
", height=" + height +
", score=" + score +
'}';
}
}

View File

@@ -4,13 +4,13 @@
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd"> xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion> <modelVersion>4.0.0</modelVersion>
<parent> <parent>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId> <artifactId>smartjavaai-parent</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
</parent> </parent>
<artifactId>smartjavaai-objectdetection</artifactId> <artifactId>smartjavaai-objectdetection</artifactId>
<version>1.0.10</version> <version>1.0.11</version>
<name>smartjavaai-objectdetection</name> <name>smartjavaai-objectdetection</name>
<description>SmartJavaAI</description> <description>SmartJavaAI</description>
<url>https://github.com/geekwenjie/SmartJavaAI</url> <url>https://github.com/geekwenjie/SmartJavaAI</url>
@@ -31,7 +31,7 @@
<dependencies> <dependencies>
<dependency> <dependency>
<groupId>ink.numberone</groupId> <groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-common</artifactId> <artifactId>smartjavaai-common</artifactId>
<version>${project.version}</version> <version>${project.version}</version>
</dependency> </dependency>

View File

@@ -26,4 +26,16 @@ public class DetectorModelConfig {
* 设备类型 * 设备类型
*/ */
private DeviceEnum device; private DeviceEnum device;
public DetectorModelConfig() {
}
public DetectorModelConfig(DetectorModelEnum modelEnum, DeviceEnum device) {
this.modelEnum = modelEnum;
this.device = device;
}
public DetectorModelConfig(DetectorModelEnum modelEnum) {
this.modelEnum = modelEnum;
}
} }

View File

@@ -7,6 +7,7 @@ import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image; import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory; import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.DetectedObjects; import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.opencv.OpenCVImageFactory;
import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException; import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel; import ai.djl.repository.zoo.ZooModel;
@@ -18,6 +19,7 @@ import cn.smartjavaai.common.pool.ModelPredictorPoolManager;
import cn.smartjavaai.common.pool.PredictorFactory; import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils; import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.common.utils.ImageUtils; import cn.smartjavaai.common.utils.ImageUtils;
import cn.smartjavaai.common.utils.OpenCVUtils;
import cn.smartjavaai.objectdetection.DetectorConfig; import cn.smartjavaai.objectdetection.DetectorConfig;
import cn.smartjavaai.objectdetection.DetectorModelConfig; import cn.smartjavaai.objectdetection.DetectorModelConfig;
import cn.smartjavaai.objectdetection.exception.DetectionException; import cn.smartjavaai.objectdetection.exception.DetectionException;
@@ -145,7 +147,7 @@ public class DetectorModel implements AutoCloseable{
if(!ImageUtils.isImageValid(image)){ if(!ImageUtils.isImageValid(image)){
throw new DetectionException("图像无效"); throw new DetectionException("图像无效");
} }
Image img = ImageFactory.getInstance().fromImage(image); Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(image));
DetectedObjects detectedObjects = detect(img); DetectedObjects detectedObjects = detect(img);
return DetectorUtils.convertToDetectionResponse(detectedObjects, img); return DetectorUtils.convertToDetectionResponse(detectedObjects, img);
} }
@@ -159,7 +161,7 @@ public class DetectorModel implements AutoCloseable{
if(!ImageUtils.isImageValid(sourceImage)){ if(!ImageUtils.isImageValid(sourceImage)){
throw new DetectionException("图像无效"); throw new DetectionException("图像无效");
} }
Image img = ImageFactory.getInstance().fromImage(sourceImage); Image img = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(sourceImage));
DetectedObjects detectedObjects = detect(img); DetectedObjects detectedObjects = detect(img);
img.drawBoundingBoxes(detectedObjects); img.drawBoundingBoxes(detectedObjects);
try { try {

149
smartjavaai-ocr/pom.xml Normal file
View File

@@ -0,0 +1,149 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<parent>
<groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-parent</artifactId>
<version>1.0.11</version>
</parent>
<artifactId>smartjavaai-ocr</artifactId>
<properties>
<maven.compiler.source>11</maven.compiler.source>
<maven.compiler.target>11</maven.compiler.target>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
<javacv.version>1.5.8</javacv.version>
<javacv.ffmpeg.version>5.1.2-1.5.8</javacv.ffmpeg.version>
</properties>
<dependencies>
<dependency>
<groupId>cn.smartjavaai</groupId>
<artifactId>smartjavaai-common</artifactId>
<version>${project.version}</version>
</dependency>
<dependency>
<groupId>org.bytedeco</groupId>
<artifactId>javacpp</artifactId>
<version>${javacv.version}</version>
</dependency>
<dependency>
<groupId>org.bytedeco</groupId>
<artifactId>ffmpeg</artifactId>
<version>${javacv.ffmpeg.version}</version>
</dependency>
<dependency>
<groupId>ai.djl.opencv</groupId>
<artifactId>opencv</artifactId>
</dependency>
</dependencies>
<version>1.0.11</version>
<name>smartjavaai-ocr</name>
<description>SmartJavaAI</description>
<url>https://github.com/geekwenjie/SmartJavaAI</url>
<licenses>
<license>
<name>MIT License</name>
<url>https://opensource.org/licenses/MIT</url>
</license>
</licenses>
<build>
<plugins>
<plugin>
<groupId>org.sonatype.central</groupId>
<artifactId>central-publishing-maven-plugin</artifactId>
<version>0.4.0</version>
<extensions>true</extensions>
<configuration>
<publishingServerId>dengwenjie</publishingServerId>
<tokenAuth>true</tokenAuth>
<deploymentName>${project.groupId}:${project.artifactId}:${project.version}</deploymentName>
</configuration>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-source-plugin</artifactId>
<version>3.1.0</version>
<executions>
<execution>
<id>attach-sources</id>
<goals>
<goal>jar-no-fork</goal>
</goals>
</execution>
</executions>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-javadoc-plugin</artifactId>
<version>3.1.0</version>
<configuration>
<javadocExecutable>${java.home}/bin/javadoc</javadocExecutable>
<doclint>none</doclint>
<additionalJOptions>
<additionalJOption>-Xdoclint:none</additionalJOption>
</additionalJOptions>
</configuration>
<executions>
<execution>
<id>attach-javadocs</id>
<goals>
<goal>jar</goal>
</goals>
</execution>
</executions>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-gpg-plugin</artifactId>
<version>3.1.0</version>
<executions>
<execution>
<id>sign-artifacts</id>
<phase>verify</phase>
<goals>
<goal>sign</goal>
</goals>
</execution>
</executions>
</plugin>
</plugins>
</build>
<!-- 必须添加SCM信息 -->
<scm>
<connection>scm:git:git://github.com/geekwenjie/SmartJavaAI.git</connection>
<developerConnection>scm:git:ssh://github.com/geekwenjie/SmartJavaAI.git</developerConnection>
<url>http://github.com/geekwenjie/SmartJavaAI/tree/master</url>
</scm>
<distributionManagement>
<snapshotRepository>
<id>dengwenjie</id>
<url>https://s01.oss.sonatype.org/content/repositories/snapshots</url>
</snapshotRepository>
<repository>
<id>dengwenjie</id>
<url>https://s01.oss.sonatype.org/service/local/staging/deploy/maven2/</url>
</repository>
</distributionManagement>
<developers>
<developer>
<name>dengwenjie</name>
<email>775747758@qq.com</email>
<roles>
<role>Project Manager</role>
<role>Architect</role>
</roles>
</developer>
</developers>
</project>

View File

@@ -0,0 +1,24 @@
package cn.smartjavaai.ocr;
import cn.smartjavaai.common.entity.DetectionResponse;
/**
* 人脸识别算法
* @author dwj
*/
public abstract class AbstractOcrModel implements OcrModel {
@Override
public void loadModel(OcrModelConfig config) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public DetectionResponse detect(String imagePath) {
throw new UnsupportedOperationException("默认不支持该功能");
}
@Override
public void detectAndDraw(String imagePath, String outputPath) {
throw new UnsupportedOperationException("默认不支持该功能");
}
}

View File

@@ -0,0 +1,35 @@
package cn.smartjavaai.ocr;
import cn.smartjavaai.common.entity.DetectionResponse;
import java.awt.image.BufferedImage;
import java.io.InputStream;
/**
* 人脸识别算法
* @author dwj
*/
public interface OcrModel {
/**
* 加载模型
* @param config
*/
void loadModel(OcrModelConfig config); // 加载模型
/**
* 人脸检测
* @param imagePath 图片路径
* @return
*/
DetectionResponse detect(String imagePath);
/**
* 检测并绘制结果
* @param imagePath 图片输入路径(包含文件名称)
* @param outputPath 图片输出路径(包含文件名称)
*/
void detectAndDraw(String imagePath, String outputPath);
}

View File

@@ -0,0 +1,28 @@
package cn.smartjavaai.ocr;
import cn.smartjavaai.common.enums.DeviceEnum;
import lombok.Data;
/**
* @author dwj
* @date 2025/4/22
*/
@Data
public class OcrModelConfig {
/**
* 模型名称
*/
private OcrModelEnum modelEnum;
/**
* 设备类型
*/
private DeviceEnum device;
/**
* 模型路径
*/
private String modelPath;
}

View File

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

View File

@@ -0,0 +1,95 @@
package cn.smartjavaai.ocr;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.model.PaddleOCRV4DetectModel;
import lombok.extern.slf4j.Slf4j;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 人脸算法工厂
* @author dwj
*/
@Slf4j
public class OcrModelFactory {
// 使用 volatile 和双重检查锁定来确保线程安全的单例模式
private static volatile OcrModelFactory instance;
private static final ConcurrentHashMap<String, OcrModel> modelMap = new ConcurrentHashMap<>();
/**
* 算法注册表
*/
private static final Map<String, Class<? extends OcrModel>> registry =
new ConcurrentHashMap<>();
public static OcrModelFactory getInstance() {
if (instance == null) {
synchronized (OcrModelFactory.class) {
if (instance == null) {
instance = new OcrModelFactory();
}
}
}
return instance;
}
/**
* 注册算法
* @param name
* @param clazz
*/
private static void registerModel(String name, Class<? extends OcrModel> clazz) {
registry.put(name.toLowerCase(), clazz);
}
/**
* 获取模型(通过配置)
* @param config
* @return
*/
public OcrModel getModel(OcrModelConfig config) {
if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){
throw new OcrException("未配置OCR模型");
}
return modelMap.computeIfAbsent(config.getModelEnum().name(), k -> {
return createFaceModel(config);
});
}
/**
* 使用ModelConfig创建算法
* @param config
* @return
*/
private OcrModel createFaceModel(OcrModelConfig config) {
Class<?> clazz = registry.get(config.getModelEnum().name().toLowerCase());
if(clazz == null){
throw new OcrException("Unsupported model");
}
OcrModel algorithm = null;
try {
algorithm = (OcrModel) clazz.newInstance();
} catch (InstantiationException | IllegalAccessException e) {
throw new OcrException(e);
}
algorithm.loadModel(config);
return algorithm;
}
// 初始化默认算法
static {
registerModel("PADDLEOCR_V4_DET_MODEL", PaddleOCRV4DetectModel.class);
log.info("缓存目录:{}", Config.getCachePath());
}
}

View File

@@ -0,0 +1,50 @@
package cn.smartjavaai.ocr.entity;
import ai.djl.modality.cv.Image;
import ai.djl.ndarray.NDArray;
/**
* 图像信息
*/
public class ImageInfo {
private String name;
private Double prob;
private Image image;
private NDArray box;
public ImageInfo(Image image, NDArray box) {
this.image = image;
this.box = box;
}
public String getName() {
return name;
}
public void setName(String name) {
this.name = name;
}
public Double getProb() {
return prob;
}
public void setProb(Double prob) {
this.prob = prob;
}
public Image getImage() {
return image;
}
public void setImage(Image image) {
this.image = image;
}
public NDArray getBox() {
return box;
}
public void setBox(NDArray box) {
this.box = box;
}
}

View File

@@ -0,0 +1,46 @@
package cn.smartjavaai.ocr.entity;
import ai.djl.ndarray.NDArray;
/**
* 旋转检测框
*/
public class RotatedBox implements Comparable<RotatedBox> {
private NDArray box;
private String text;
public RotatedBox(NDArray box, String text) {
this.box = box;
this.text = text;
}
/**
* 将左上角 Y 坐标升序排序
*
* @param o
* @return
*/
@Override
public int compareTo(RotatedBox o) {
NDArray lowBox = this.getBox();
NDArray highBox = o.getBox();
float lowY = lowBox.toFloatArray()[1];
float highY = highBox.toFloatArray()[1];
return (lowY < highY) ? -1 : 1;
}
public NDArray getBox() {
return box;
}
public void setBox(NDArray box) {
this.box = box;
}
public String getText() {
return text;
}
public void setText(String text) {
this.text = text;
}
}

View File

@@ -0,0 +1,46 @@
package cn.smartjavaai.ocr.entity;
import ai.djl.ndarray.NDArray;
/**
* 旋转检测框 - 支持左上角 X 坐标升序排序
*/
public class RotatedBoxCompX implements Comparable<RotatedBoxCompX> {
private NDArray box;
private String text;
public RotatedBoxCompX(NDArray box, String text) {
this.box = box;
this.text = text;
}
/**
* 将左上角 X 坐标升序排序
*
* @param o
* @return
*/
@Override
public int compareTo(RotatedBoxCompX o) {
NDArray leftBox = this.getBox();
NDArray rightBox = o.getBox();
float leftX = leftBox.toFloatArray()[0];
float rightX = rightBox.toFloatArray()[0];
return (leftX < rightX) ? -1 : 1;
}
public NDArray getBox() {
return box;
}
public void setBox(NDArray box) {
this.box = box;
}
public String getText() {
return text;
}
public void setText(String text) {
this.text = text;
}
}

View File

@@ -0,0 +1,30 @@
package cn.smartjavaai.ocr.exception;
/**
* 人脸检测异常
* @author dwj
* @date 2025/4/4
*/
public class OcrException extends RuntimeException{
public OcrException() {
super();
}
public OcrException(String message, Throwable cause, boolean enableSuppression, boolean writableStackTrace) {
super(message, cause, enableSuppression, writableStackTrace);
}
public OcrException(String message, Throwable cause) {
super(message, cause);
}
public OcrException(String message) {
super(message);
}
public OcrException(Throwable cause) {
super(cause);
}
}

View File

@@ -0,0 +1,134 @@
package cn.smartjavaai.ocr.model;
import ai.djl.MalformedModelException;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ModelZoo;
import ai.djl.repository.zoo.ZooModel;
import ai.djl.training.util.ProgressBar;
import ai.djl.translate.TranslateException;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.common.utils.FileUtils;
import cn.smartjavaai.ocr.AbstractOcrModel;
import cn.smartjavaai.ocr.OcrModelConfig;
import cn.smartjavaai.ocr.exception.OcrException;
import cn.smartjavaai.ocr.translator.PaddleOCRV4DetectionTranslator;
import cn.smartjavaai.ocr.utils.ImageUtils;
import cn.smartjavaai.ocr.utils.OcrUtils;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.lang3.StringUtils;
import org.apache.commons.pool2.ObjectPool;
import org.apache.commons.pool2.impl.GenericObjectPool;
import org.opencv.core.Mat;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* @author dwj
* @date 2025/4/21
*/
@Slf4j
public class PaddleOCRV4DetectModel extends AbstractOcrModel {
private ZooModel detectionModel;
private ObjectPool<Predictor<Image, NDList>> predictorPool;
@Override
public void loadModel(OcrModelConfig config){
if(StringUtils.isBlank(config.getModelPath())){
throw new OcrException("modelPath is null");
}
Criteria<Image, NDList> criteria =
Criteria.builder()
.optEngine("OnnxRuntime")
.setTypes(Image.class, NDList.class)
.optModelPath(Paths.get(config.getModelPath()))
.optTranslator(new PaddleOCRV4DetectionTranslator(new ConcurrentHashMap<String, String>()))
.optProgress(new ProgressBar())
.build();
try{
detectionModel = ModelZoo.loadModel(criteria);
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(detectionModel));
log.info("当前设备: " + detectionModel.getNDManager().getDevice());
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new OcrException("模型加载失败", e);
}
}
@Override
public DetectionResponse detect(String imagePath){
if(!FileUtils.isFileExists(imagePath)){
throw new OcrException("图像文件不存在");
}
Image img = null;
try {
img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
} catch (IOException e) {
throw new OcrException("无效的图片", e);
}
return detect(img);
}
/**
* 人脸检测
* @param image
* @return
*/
private DetectionResponse detect(Image image){
Predictor<Image, NDList> predictor = null;
try {
predictor = predictorPool.borrowObject();
NDList result = predictor.predict(image);
return OcrUtils.convertToDetectionResponse(result, image);
} catch (Exception e) {
throw new OcrException("OCR检测错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@Override
public void detectAndDraw(String imagePath, String outputPath) {
if(!FileUtils.isFileExists(imagePath)){
throw new OcrException("图像文件不存在");
}
try {
Image img = ImageFactory.getInstance().fromFile(Paths.get(imagePath));
DetectionResponse detectionResponse = detect(img);
if(Objects.isNull(detectionResponse) || Objects.isNull(detectionResponse.getRectangleList()) || detectionResponse.getRectangleList().isEmpty()){
throw new OcrException("未识别到文字");
}
ImageUtils.drawRect((Mat)img.getWrappedImage(), detectionResponse);
Path output = Paths.get(outputPath);
log.info("Saving to {}", output.toAbsolutePath().toString());
img.save(Files.newOutputStream(output), "png");
} catch (IOException e) {
throw new OcrException(e);
}
}
}

View File

@@ -0,0 +1,27 @@
package cn.smartjavaai.ocr.model;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelZoo;
import ai.djl.repository.zoo.ZooModel;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.ocr.translator.PaddleOCRV4DetectionTranslator;
import org.opencv.core.Mat;
import java.nio.file.Paths;
import java.util.concurrent.ConcurrentHashMap;
/**
* @author dwj
* @date 2025/4/21
*/
public class PaddleOCRV4Model {
public void loadModel(){
}
}

View File

@@ -0,0 +1,228 @@
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<Point> 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<ai.djl.modality.cv.output.Point> 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;
}
}

View File

@@ -0,0 +1,64 @@
package cn.smartjavaai.ocr.opencv;
import org.opencv.core.CvType;
import org.opencv.core.Mat;
import org.opencv.imgproc.Imgproc;
import java.awt.image.BufferedImage;
import java.awt.image.DataBufferByte;
/**
* OpenCV Utils
*
*/
public class OcrOpenCVUtils {
/**
* 透视变换
*
* @param src
* @param srcPoints
* @param dstPoints
* @return
*/
public static Mat perspectiveTransform(Mat src, Mat srcPoints, Mat dstPoints) {
Mat dst = src.clone();
Mat warp_mat = Imgproc.getPerspectiveTransform(srcPoints, dstPoints);
Imgproc.warpPerspective(src, dst, warp_mat, dst.size());
warp_mat.release();
return dst;
}
/**
* Mat to BufferedImage
*
* @param mat
* @return
*/
public static BufferedImage mat2Image(Mat mat) {
int width = mat.width();
int height = mat.height();
byte[] data = new byte[width * height * (int) mat.elemSize()];
Imgproc.cvtColor(mat, mat, 4);
mat.get(0, 0, data);
BufferedImage ret = new BufferedImage(width, height, 5);
ret.getRaster().setDataElements(0, 0, width, height, data);
return ret;
}
/**
* BufferedImage to Mat
*
* @param img
* @return
*/
public static Mat image2Mat(BufferedImage img) {
int width = img.getWidth();
int height = img.getHeight();
byte[] data = ((DataBufferByte) img.getRaster().getDataBuffer()).getData();
Mat mat = new Mat(height, width, CvType.CV_8UC3);
mat.put(0, 0, data);
return mat;
}
}

View File

@@ -0,0 +1,521 @@
package cn.smartjavaai.ocr.translator;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.util.NDImageUtils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDArrays;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.ndarray.index.NDIndex;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import ai.djl.translate.Batchifier;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;
import cn.smartjavaai.ocr.opencv.OcrNDArrayUtils;
import org.opencv.core.*;
import org.opencv.imgproc.Imgproc;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
/**
* 文字检测前后处理
*
* @author Calvin
* @mail 179209347@qq.com
* @website www.aias.top
*/
public class PaddleOCRV4DetectionTranslator implements Translator<Image, NDList> {
// det_algorithm == "DB"
private final float thresh = 0.3f;
private final boolean use_dilation = false;
private final String score_mode = "fast";
private final String box_type = "quad";
private final int limit_side_len;
private final int max_candidates;
private final int min_size;
private final float box_thresh;
private final float unclip_ratio;
private float ratio_h;
private float ratio_w;
private int img_height;
private int img_width;
public PaddleOCRV4DetectionTranslator(Map<String, ?> arguments) {
limit_side_len =
arguments.containsKey("limit_side_len")
? Integer.parseInt(arguments.get("limit_side_len").toString())
: 960;
max_candidates =
arguments.containsKey("max_candidates")
? Integer.parseInt(arguments.get("max_candidates").toString())
: 1000;
min_size =
arguments.containsKey("min_size")
? Integer.parseInt(arguments.get("min_size").toString())
: 3;
box_thresh =
arguments.containsKey("box_thresh")
? Float.parseFloat(arguments.get("box_thresh").toString())
: 0.6f; // 0.5f
unclip_ratio =
arguments.containsKey("unclip_ratio")
? Float.parseFloat(arguments.get("unclip_ratio").toString())
: 1.6f;
}
@Override
public NDList processOutput(TranslatorContext ctx, NDList list) {
NDManager manager = ctx.getNDManager();
NDArray pred = list.get(0);
pred = pred.squeeze();
NDArray segmentation = pred.gt(thresh); // thresh=0.3 .mul(255f)
segmentation = segmentation.toType(DataType.UINT8, true);
Shape shape = segmentation.getShape();
int rows = (int) shape.get(0);
int cols = (int) shape.get(1);
Mat newMask = new Mat();
if (this.use_dilation) {
Mat mask = new Mat();
//convert from NDArray to Mat
Mat srcMat = OcrNDArrayUtils.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}});
/**
* 膨胀说明: 图像的一部分区域与指定的核进行卷积, 求核的最`大`值并赋值给指定区域。 膨胀可以理解为图像中`高亮区域`的'领域扩大'。
* 意思是高亮部分会侵蚀不是高亮的部分,使高亮部分越来越多。
*/
Imgproc.dilate(srcMat, mask, dilation_kernel);
//destination Matrix
Scalar scalar = new Scalar(255);
Core.multiply(mask, scalar, newMask);
// release Mat
mask.release();
srcMat.release();
dilation_kernel.release();
} else {
Mat srcMat = OcrNDArrayUtils.uint8NDArrayToMat(segmentation);
//destination Matrix
Scalar scalar = new Scalar(255);
Core.multiply(srcMat, scalar, newMask);
// release Mat
srcMat.release();
}
NDArray boxes = boxes_from_bitmap(manager, pred, newMask);
//boxes[:, :, 0] = boxes[:, :, 0] / ratio_w
NDArray boxes1 = boxes.get(":, :, 0").div(ratio_w);
boxes.set(new NDIndex(":, :, 0"), boxes1);
//boxes[:, :, 1] = boxes[:, :, 1] / ratio_h
NDArray boxes2 = boxes.get(":, :, 1").div(ratio_h);
boxes.set(new NDIndex(":, :, 1"), boxes2);
NDList dt_boxes = this.filter_tag_det_res(boxes);
dt_boxes.detach();
// release Mat
newMask.release();
return dt_boxes;
}
private NDList filter_tag_det_res(NDArray dt_boxes) {
NDList boxesList = new NDList();
int num = (int) dt_boxes.getShape().get(0);
for (int i = 0; i < num; i++) {
NDArray box = dt_boxes.get(i);
box = order_points_clockwise(box);
box = clip_det_res(box);
float[] box0 = box.get(0).toFloatArray();
float[] box1 = box.get(1).toFloatArray();
float[] box3 = box.get(3).toFloatArray();
int rect_width = (int) Math.sqrt(Math.pow(box1[0] - box0[0], 2) + Math.pow(box1[1] - box0[1], 2));
int rect_height = (int) Math.sqrt(Math.pow(box3[0] - box0[0], 2) + Math.pow(box3[1] - box0[1], 2));
if (rect_width <= 3 || rect_height <= 3)
continue;
boxesList.add(box);
}
return boxesList;
}
private NDArray clip_det_res(NDArray points) {
for (int i = 0; i < points.getShape().get(0); i++) {
int value = Math.max((int) points.get(i, 0).toFloatArray()[0], 0);
value = Math.min(value, img_width - 1);
points.set(new NDIndex(i + ",0"), value);
value = Math.max((int) points.get(i, 1).toFloatArray()[0], 0);
value = Math.min(value, img_height - 1);
points.set(new NDIndex(i + ",1"), value);
}
return points;
}
/**
* sort the points based on their x-coordinates
* 顺时针
*
* @param pts
* @return
*/
private NDArray order_points_clockwise(NDArray pts) {
NDList list = new NDList();
long[] indexes = pts.get(":, 0").argSort().toLongArray();
// grab the left-most and right-most points from the sorted
// x-roodinate points
Shape s1 = pts.getShape();
NDArray leftMost1 = pts.get(indexes[0] + ",:");
NDArray leftMost2 = pts.get(indexes[1] + ",:");
NDArray leftMost = leftMost1.concat(leftMost2).reshape(2, 2);
NDArray rightMost1 = pts.get(indexes[2] + ",:");
NDArray rightMost2 = pts.get(indexes[3] + ",:");
NDArray rightMost = rightMost1.concat(rightMost2).reshape(2, 2);
// now, sort the left-most coordinates according to their
// y-coordinates so we can grab the top-left and bottom-left
// points, respectively
indexes = leftMost.get(":, 1").argSort().toLongArray();
NDArray lt = leftMost.get(indexes[0] + ",:");
NDArray lb = leftMost.get(indexes[1] + ",:");
indexes = rightMost.get(":, 1").argSort().toLongArray();
NDArray rt = rightMost.get(indexes[0] + ",:");
NDArray rb = rightMost.get(indexes[1] + ",:");
list.add(lt);
list.add(rt);
list.add(rb);
list.add(lb);
NDArray rect = NDArrays.concat(list).reshape(4, 2);
return rect;
}
/**
* Get boxes from the binarized image predicted by DB
*
* @param manager
* @param pred the binarized image predicted by DB.
* @param bitmap new 'pred' after threshold filtering.
*/
private NDArray boxes_from_bitmap(NDManager manager, NDArray pred, Mat bitmap) {
int dest_height = (int) pred.getShape().get(0);
int dest_width = (int) pred.getShape().get(1);
int height = bitmap.rows();
int width = bitmap.cols();
List<MatOfPoint> contours = new ArrayList<>();
Mat hierarchy = new Mat();
// 寻找轮廓
Imgproc.findContours(
bitmap,
contours,
hierarchy,
Imgproc.RETR_LIST,
Imgproc.CHAIN_APPROX_SIMPLE);
int num_contours = Math.min(contours.size(), max_candidates);
NDList boxList = new NDList();
float[] scores = new float[num_contours];
for (int index = 0; index < num_contours; index++) {
MatOfPoint contour = contours.get(index);
MatOfPoint2f newContour = new MatOfPoint2f(contour.toArray());
float[][] pointsArr = new float[4][2];
int sside = get_mini_boxes(newContour, pointsArr);
if (sside < this.min_size)
continue;
NDArray points = manager.create(pointsArr);
float score = box_score_fast(manager, pred, points);
if (score < this.box_thresh)
continue;
NDArray box = unclip(manager, points); // TODO get_mini_boxes(box)
// box[:, 0] = np.clip(np.round(box[:, 0] / width * dest_width), 0, dest_width)
NDArray boxes1 = box.get(":,0").div(width).mul(dest_width).round().clip(0, dest_width);
box.set(new NDIndex(":, 0"), boxes1);
// box[:, 1] = np.clip(np.round(box[:, 1] / height * dest_height), 0, dest_height)
NDArray boxes2 = box.get(":,1").div(height).mul(dest_height).round().clip(0, dest_height);
box.set(new NDIndex(":, 1"), boxes2);
boxList.add(box);
scores[index] = score;
// release memory
contour.release();
newContour.release();
}
NDArray boxes = NDArrays.stack(boxList);
// release
hierarchy.release();
return boxes;
}
/**
* Shrink or expand the boxaccording to 'unclip_ratio'
*
* @param points The predicted box.
* @return uncliped box
*/
private NDArray unclip(NDManager manager, NDArray points) {
points = order_points_clockwise(points);
float[] pointsArr = points.toFloatArray();
float[] lt = java.util.Arrays.copyOfRange(pointsArr, 0, 2);
float[] lb = java.util.Arrays.copyOfRange(pointsArr, 6, 8);
float[] rt = java.util.Arrays.copyOfRange(pointsArr, 2, 4);
float[] rb = java.util.Arrays.copyOfRange(pointsArr, 4, 6);
float width = distance(lt, rt);
float height = distance(lt, lb);
if (width > height) {
float k = (lt[1] - rt[1]) / (lt[0] - rt[0]); // y = k * x + b
float delta_dis = height;
float delta_x = (float) Math.sqrt((delta_dis * delta_dis) / (k * k + 1));
float delta_y = Math.abs(k * delta_x);
if (k > 0) {
pointsArr[0] = lt[0] - delta_x + delta_y;
pointsArr[1] = lt[1] - delta_y - delta_x;
pointsArr[2] = rt[0] + delta_x + delta_y;
pointsArr[3] = rt[1] + delta_y - delta_x;
pointsArr[4] = rb[0] + delta_x - delta_y;
pointsArr[5] = rb[1] + delta_y + delta_x;
pointsArr[6] = lb[0] - delta_x - delta_y;
pointsArr[7] = lb[1] - delta_y + delta_x;
} else {
pointsArr[0] = lt[0] - delta_x - delta_y;
pointsArr[1] = lt[1] + delta_y - delta_x;
pointsArr[2] = rt[0] + delta_x - delta_y;
pointsArr[3] = rt[1] - delta_y - delta_x;
pointsArr[4] = rb[0] + delta_x + delta_y;
pointsArr[5] = rb[1] - delta_y + delta_x;
pointsArr[6] = lb[0] - delta_x + delta_y;
pointsArr[7] = lb[1] + delta_y + delta_x;
}
} else {
float k = (lt[1] - rt[1]) / (lt[0] - rt[0]); // y = k * x + b
float delta_dis = width;
float delta_y = (float) Math.sqrt((delta_dis * delta_dis) / (k * k + 1));
float delta_x = Math.abs(k * delta_y);
if (k > 0) {
pointsArr[0] = lt[0] + delta_x - delta_y;
pointsArr[1] = lt[1] - delta_y - delta_x;
pointsArr[2] = rt[0] + delta_x + delta_y;
pointsArr[3] = rt[1] - delta_y + delta_x;
pointsArr[4] = rb[0] - delta_x + delta_y;
pointsArr[5] = rb[1] + delta_y + delta_x;
pointsArr[6] = lb[0] - delta_x - delta_y;
pointsArr[7] = lb[1] + delta_y - delta_x;
} else {
pointsArr[0] = lt[0] - delta_x - delta_y;
pointsArr[1] = lt[1] - delta_y + delta_x;
pointsArr[2] = rt[0] - delta_x + delta_y;
pointsArr[3] = rt[1] - delta_y - delta_x;
pointsArr[4] = rb[0] + delta_x + delta_y;
pointsArr[5] = rb[1] + delta_y - delta_x;
pointsArr[6] = lb[0] + delta_x - delta_y;
pointsArr[7] = lb[1] + delta_y + delta_x;
}
}
points = manager.create(pointsArr).reshape(4, 2);
return points;
}
private float distance(float[] point1, float[] point2) {
float disX = point1[0] - point2[0];
float disY = point1[1] - point2[1];
float dis = (float) Math.sqrt(disX * disX + disY * disY);
return dis;
}
/**
* Get boxes from the contour or box.
*
* @param contour The predicted contour.
* @param pointsArr The predicted box.
* @return smaller side of box
*/
private int get_mini_boxes(MatOfPoint2f contour, float[][] pointsArr) {
// https://blog.csdn.net/qq_37385726/article/details/82313558
// bounding_box[1] - rect 返回矩形的长和宽
RotatedRect rect = Imgproc.minAreaRect(contour);
Mat points = new Mat();
Imgproc.boxPoints(rect, points);
float[][] fourPoints = new float[4][2];
for (int row = 0; row < 4; row++) {
fourPoints[row][0] = (float) points.get(row, 0)[0];
fourPoints[row][1] = (float) points.get(row, 1)[0];
}
float[] tmpPoint = new float[2];
for (int i = 0; i < 4; i++) {
for (int j = i + 1; j < 4; j++) {
if (fourPoints[j][0] < fourPoints[i][0]) {
tmpPoint[0] = fourPoints[i][0];
tmpPoint[1] = fourPoints[i][1];
fourPoints[i][0] = fourPoints[j][0];
fourPoints[i][1] = fourPoints[j][1];
fourPoints[j][0] = tmpPoint[0];
fourPoints[j][1] = tmpPoint[1];
}
}
}
int index_1 = 0;
int index_2 = 1;
int index_3 = 2;
int index_4 = 3;
if (fourPoints[1][1] > fourPoints[0][1]) {
index_1 = 0;
index_4 = 1;
} else {
index_1 = 1;
index_4 = 0;
}
if (fourPoints[3][1] > fourPoints[2][1]) {
index_2 = 2;
index_3 = 3;
} else {
index_2 = 3;
index_3 = 2;
}
pointsArr[0] = fourPoints[index_1];
pointsArr[1] = fourPoints[index_2];
pointsArr[2] = fourPoints[index_3];
pointsArr[3] = fourPoints[index_4];
int height = rect.boundingRect().height;
int width = rect.boundingRect().width;
int sside = Math.min(height, width);
// release
points.release();
return sside;
}
/**
* Calculate the score of box.
*
* @param bitmap The binarized image predicted by DB.
* @param points The predicted box
* @return
*/
private float box_score_fast(NDManager manager, NDArray bitmap, NDArray points) {
NDArray box = points.get(":");
long h = bitmap.getShape().get(0);
long w = bitmap.getShape().get(1);
// xmin = np.clip(np.floor(box[:, 0].min()).astype(np.int), 0, w - 1)
int xmin = box.get(":, 0").min().floor().clip(0, w - 1).toType(DataType.INT32, true).toIntArray()[0];
int xmax = box.get(":, 0").max().ceil().clip(0, w - 1).toType(DataType.INT32, true).toIntArray()[0];
int ymin = box.get(":, 1").min().floor().clip(0, h - 1).toType(DataType.INT32, true).toIntArray()[0];
int ymax = box.get(":, 1").max().ceil().clip(0, h - 1).toType(DataType.INT32, true).toIntArray()[0];
NDArray mask = manager.zeros(new Shape(ymax - ymin + 1, xmax - xmin + 1), DataType.UINT8);
box.set(new NDIndex(":, 0"), box.get(":, 0").sub(xmin));
box.set(new NDIndex(":, 1"), box.get(":, 1").sub(ymin));
//mask - convert from NDArray to Mat
Mat maskMat = OcrNDArrayUtils.uint8NDArrayToMat(mask);
//mask - convert from NDArray to Mat - 4 rows, 2 cols
Mat boxMat = OcrNDArrayUtils.floatNDArrayToMat(box, CvType.CV_32S);
// boxMat.reshape(1, new int[]{1, 4, 2});
List<MatOfPoint> pts = new ArrayList<>();
MatOfPoint matOfPoint = OcrNDArrayUtils.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);
Scalar score = Core.mean(bitMapMat, maskMat);
float scoreValue = (float) score.val[0];
// release
maskMat.release();
boxMat.release();
bitMapMat.release();
return scoreValue;
}
@Override
public NDList processInput(TranslatorContext ctx, Image input) {
NDArray img = input.toNDArray(ctx.getNDManager());
int h = input.getHeight();
int w = input.getWidth();
img_height = h;
img_width = w;
// limit the max side
float ratio = 1.0f;
if (Math.max(h, w) > limit_side_len) {
if (h > w) {
ratio = (float) limit_side_len / (float) h;
} else {
ratio = (float) limit_side_len / (float) w;
}
}
int resize_h = (int) (h * ratio);
int resize_w = (int) (w * ratio);
resize_h = Math.round((float) resize_h / 32f) * 32;
resize_w = Math.round((float) resize_w / 32f) * 32;
ratio_h = resize_h / (float) h;
ratio_w = resize_w / (float) w;
img = NDImageUtils.resize(img, resize_w, resize_h);
img = NDImageUtils.toTensor(img);
img =
NDImageUtils.normalize(
img,
new float[]{0.485f, 0.456f, 0.406f},
new float[]{0.229f, 0.224f, 0.225f});
img = img.expandDims(0);
return new NDList(img);
}
@Override
public Batchifier getBatchifier() {
return null;
}
}

View File

@@ -0,0 +1,123 @@
package cn.smartjavaai.ocr.translator;
import ai.djl.Model;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.util.NDImageUtils;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.index.NDIndex;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import ai.djl.translate.Batchifier;
import ai.djl.translate.Translator;
import ai.djl.translate.TranslatorContext;
import ai.djl.util.Utils;
import java.io.IOException;
import java.io.InputStream;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
/**
* 文字识别前后处理
*
*/
public class PaddleOCRV4WordRecTranslator implements Translator<Image, String> {
private List<String> table;
private final boolean use_space_char;
public PaddleOCRV4WordRecTranslator(Map<String, ?> arguments) {
use_space_char =
arguments.containsKey("use_space_char")
? Boolean.parseBoolean(arguments.get("use_space_char").toString())
: true;
}
@Override
public void prepare(TranslatorContext ctx) throws IOException {
Model model = ctx.getModel();
try (InputStream is = model.getArtifact("dict.txt").openStream()) {
table = Utils.readLines(is, true);
table.add(0, "blank");
if(use_space_char){
table.add(" ");
table.add(" ");
}
else{
table.add("");
table.add("");
}
}
}
@Override
public String processOutput(TranslatorContext ctx, NDList list) throws IOException {
StringBuilder sb = new StringBuilder();
NDArray tokens = list.singletonOrThrow();
long[] indices = tokens.get(0).argMax(1).toLongArray();
boolean[] selection = new boolean[indices.length];
Arrays.fill(selection, true);
for (int i = 1; i < indices.length; i++) {
if (indices[i] == indices[i - 1]) {
selection[i] = false;
}
}
// 字符置信度
// float[] probs = new float[indices.length];
// for (int row = 0; row < indices.length; row++) {
// NDArray value = tokens.get(0).get(new NDIndex(""+ row +":" + (row + 1) +"," + indices[row] +":" + ( indices[row] + 1)));
// probs[row] = value.toFloatArray()[0];
// }
int lastIdx = 0;
for (int i = 0; i < indices.length; i++) {
if (selection[i] == true && indices[i] > 0 && !(i > 0 && indices[i] == lastIdx)) {
sb.append(table.get((int) indices[i]));
}
}
return sb.toString();
}
@Override
public NDList processInput(TranslatorContext ctx, Image input) {
NDArray img = input.toNDArray(ctx.getNDManager(), Image.Flag.COLOR);
int imgC = 3;
int imgH = 48;
int imgW = 320;
float max_wh_ratio = (float) imgW / (float) imgH;
int h = input.getHeight();
int w = input.getWidth();
float wh_ratio = (float) w / (float) h;
max_wh_ratio = Math.max(max_wh_ratio,wh_ratio);
imgW = (int)(imgH * max_wh_ratio);
int resized_w;
if (Math.ceil(imgH * wh_ratio) > imgW) {
resized_w = imgW;
} else {
resized_w = (int) (Math.ceil(imgH * wh_ratio));
}
NDArray resized_image = NDImageUtils.resize(img, resized_w, imgH);
resized_image = resized_image.transpose(2, 0, 1).toType(DataType.FLOAT32,false);
resized_image.divi(255f).subi(0.5f).divi(0.5f);
NDArray padding_im = ctx.getNDManager().zeros(new Shape(imgC, imgH, imgW), DataType.FLOAT32);
padding_im.set(new NDIndex(":,:,0:" + resized_w), resized_image);
padding_im = padding_im.flip(0);
padding_im = padding_im.expandDims(0);
return new NDList(padding_im);
}
@Override
public Batchifier getBatchifier() {
return null;
}
}

View File

@@ -0,0 +1,219 @@
package cn.smartjavaai.ocr.utils;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.ImageFactory;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.ndarray.NDArray;
import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.utils.OpenCVUtils;
import org.opencv.core.Mat;
import org.opencv.core.Point;
import org.opencv.core.Scalar;
import org.opencv.imgproc.Imgproc;
import java.awt.*;
import java.awt.image.BufferedImage;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.List;
/**
* 图像工具类
*/
public class ImageUtils {
/**
* 保存BufferedImage图片
*
* @param img
* @param name
* @param path
*/
public static void saveImage(BufferedImage img, String name, String path) {
Image djlImg = ImageFactory.getInstance().fromImage(OpenCVUtils.image2Mat(img)); // 支持多种图片格式,自动适配
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();
}
}
/**
* 保存DJL图片
*
* @param img
* @param name
* @param path
*/
public static void saveImage(Image img, String name, String path) {
Path outputDir = Paths.get(path);
Path imagePath = outputDir.resolve(name);
// OpenJDK 不能保存 jpg 图片的 alpha channel
try {
img.save(Files.newOutputStream(imagePath), "png");
} catch (IOException e) {
e.printStackTrace();
}
}
/**
* 保存图片,含检测框
*
* @param img
* @param detection
* @param name
* @param path
* @throws IOException
*/
public static void saveBoundingBoxImage(
Image img, DetectedObjects detection, String name, String path) throws IOException {
// Make image copy with alpha channel because original image was jpg
img.drawBoundingBoxes(detection);
Path outputDir = Paths.get(path);
Files.createDirectories(outputDir);
Path imagePath = outputDir.resolve(name);
// OpenJDK can't save jpg with alpha channel
img.save(Files.newOutputStream(imagePath), "png");
}
/**
* 画检测框(有倾斜角)
*
* @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 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 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 mat
* @param box
*/
public static void drawRect(Mat mat, DetectionResponse detectionResponse) {
for(DetectionRectangle detectionRectangle : detectionResponse.getRectangleList()){
// 左上角点
Point topLeft = new Point(detectionRectangle.getX(), detectionRectangle.getY());
// 右下角点
Point bottomRight = new Point(detectionRectangle.getX() + detectionRectangle.getWidth(), detectionRectangle.getY() + detectionRectangle.getHeight());
// 绘制矩形图像、左上角、右下角、颜色BGR、线宽
Imgproc.rectangle(mat, topLeft, bottomRight, new Scalar(0, 255, 0), 1);
}
}
}

View File

@@ -0,0 +1,64 @@
package cn.smartjavaai.ocr.utils;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.BoundingBox;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import cn.smartjavaai.common.entity.DetectionRectangle;
import cn.smartjavaai.common.entity.DetectionResponse;
import lombok.extern.slf4j.Slf4j;
import java.math.BigDecimal;
import java.util.ArrayList;
import java.util.Iterator;
import java.util.List;
import java.util.Objects;
/**
* @author dwj
* @date 2025/4/22
*/
@Slf4j
public class OcrUtils {
/**
* 转换为FaceDetectedResult
* @param dt_boxes
* @param img
* @return
*/
public static DetectionResponse convertToDetectionResponse(NDList dt_boxes, Image img){
if(Objects.isNull(dt_boxes) || dt_boxes.size() == 0){
return null;
}
DetectionResponse detectionResponse = new DetectionResponse();
List<DetectionRectangle> rectangleList = new ArrayList<DetectionRectangle>();
for(NDArray box : dt_boxes){
DetectionRectangle rectangle = new DetectionRectangle();
float[] points = box.toFloatArray();
log.info("points: {}", points);
int x = (int)points[0];
int y = (int)points[1];
int width = new BigDecimal(points[4]).subtract(new BigDecimal(points[6])).intValue();
int height = new BigDecimal(points[7]).subtract(new BigDecimal(points[1])).intValue();
// 修正边界,防止越界
if (x < 0) x = 0;
if (y < 0) y = 0;
if (x + width > img.getWidth()) width = img.getWidth() - x;
if (y + height > img.getHeight()) height = img.getHeight() - y;
rectangle.setX(x);
rectangle.setY(y);
rectangle.setHeight(height);
rectangle.setWidth(width);
rectangleList.add(rectangle);
}
detectionResponse.setRectangleList(rectangleList);
return detectionResponse;
}
}