mirror of
https://github.com/geekwenjie/SmartJavaAI.git
synced 2026-09-13 13:18:58 +00:00
1、目标检测:支持自己训练的模型推理
2、目标检测:支持yolo12模型 3、支持JDK8使用 4、引入离线依赖库,支持完全离线使用 5、优化FaceNet人脸比对速度 6、支持4通道图片检测
This commit is contained in:
@@ -6,11 +6,11 @@
|
||||
<parent>
|
||||
<groupId>cn.smartjavaai</groupId>
|
||||
<artifactId>smartjavaai-parent</artifactId>
|
||||
<version>1.0.12</version>
|
||||
<version>1.0.13</version>
|
||||
</parent>
|
||||
|
||||
<artifactId>smartjavaai-objectdetection</artifactId>
|
||||
<version>1.0.12</version>
|
||||
<version>1.0.13</version>
|
||||
<name>smartjavaai-objectdetection</name>
|
||||
<description>SmartJavaAI</description>
|
||||
<url>https://github.com/geekwenjie/SmartJavaAI</url>
|
||||
@@ -23,8 +23,8 @@
|
||||
|
||||
|
||||
<properties>
|
||||
<maven.compiler.source>11</maven.compiler.source>
|
||||
<maven.compiler.target>11</maven.compiler.target>
|
||||
<!-- <maven.compiler.source>11</maven.compiler.source>-->
|
||||
<!-- <maven.compiler.target>11</maven.compiler.target>-->
|
||||
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
|
||||
<maven.test.skip>true</maven.test.skip>
|
||||
</properties>
|
||||
@@ -72,7 +72,7 @@
|
||||
<artifactId>maven-javadoc-plugin</artifactId>
|
||||
<version>3.1.0</version>
|
||||
<configuration>
|
||||
<javadocExecutable>${java.home}/bin/javadoc</javadocExecutable>
|
||||
<!-- <javadocExecutable>${java.home}/bin/javadoc</javadocExecutable>-->
|
||||
<doclint>none</doclint>
|
||||
<additionalJOptions>
|
||||
<additionalJOption>-Xdoclint:none</additionalJOption>
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
package cn.smartjavaai.objectdetection;
|
||||
package cn.smartjavaai.objectdetection.config;
|
||||
|
||||
import cn.smartjavaai.common.enums.DeviceEnum;
|
||||
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
|
||||
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
|
||||
import lombok.Data;
|
||||
|
||||
/**
|
||||
@@ -20,13 +22,25 @@ public class DetectorModelConfig {
|
||||
/**
|
||||
* 置信度阈值
|
||||
*/
|
||||
private float threshold = DetectorConfig.DEFAULT_THRESHOLD;
|
||||
private float threshold = DetectorConstant.DEFAULT_THRESHOLD;
|
||||
|
||||
/**
|
||||
* 设备类型
|
||||
*/
|
||||
private DeviceEnum device;
|
||||
|
||||
/**
|
||||
* 模型路径
|
||||
*/
|
||||
private String modelPath;
|
||||
|
||||
/**
|
||||
* 候选框数量:默认为8400. 应设置0到8400之间的整数
|
||||
* 用于性能优化的关键参数,它通过限制模型后处理阶段需要处理的候选框(bounding boxes)数量来提高推理速度
|
||||
* 建议不低于1000
|
||||
*/
|
||||
private int maxBox;
|
||||
|
||||
public DetectorModelConfig() {
|
||||
}
|
||||
|
||||
@@ -1,14 +1,17 @@
|
||||
package cn.smartjavaai.objectdetection;
|
||||
package cn.smartjavaai.objectdetection.constant;
|
||||
|
||||
/**
|
||||
* @author dwj
|
||||
* @date 2025/4/7
|
||||
*/
|
||||
public class DetectorConfig {
|
||||
public class DetectorConstant {
|
||||
|
||||
/**
|
||||
* 置信度阈值
|
||||
*/
|
||||
public static final float DEFAULT_THRESHOLD = 0.5F;
|
||||
|
||||
|
||||
|
||||
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package cn.smartjavaai.objectdetection.criteria;
|
||||
|
||||
import ai.djl.modality.cv.Image;
|
||||
import ai.djl.modality.cv.output.DetectedObjects;
|
||||
import ai.djl.repository.zoo.Criteria;
|
||||
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
|
||||
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
|
||||
import cn.smartjavaai.objectdetection.exception.DetectionException;
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
|
||||
/**
|
||||
* Criteria构建工厂
|
||||
* @author dwj
|
||||
* @date 2025/5/14
|
||||
*/
|
||||
public class CriteriaBuilderFactory {
|
||||
|
||||
public static Criteria<Image, DetectedObjects> createCriteria(DetectorModelConfig config) {
|
||||
//以下模型modelPath不允许为空
|
||||
if(config.getModelEnum() == DetectorModelEnum.YOLOV8_OFFICIAL ||
|
||||
config.getModelEnum() == DetectorModelEnum.YOLOV12_OFFICIAL ||
|
||||
config.getModelEnum() == DetectorModelEnum.YOLOV8_CUSTOM ||
|
||||
config.getModelEnum() == DetectorModelEnum.YOLOV12_CUSTOM){
|
||||
if(StringUtils.isBlank(config.getModelPath())){
|
||||
throw new DetectionException("modelPath is null");
|
||||
}
|
||||
}
|
||||
switch (config.getModelEnum()) {
|
||||
case YOLOV8_OFFICIAL:
|
||||
return new YoloCriteriaBuilder().buildCriteria(config);
|
||||
case YOLOV12_OFFICIAL:
|
||||
return new YoloCriteriaBuilder().buildCriteria(config);
|
||||
case YOLOV8_CUSTOM:
|
||||
return new YoloCriteriaBuilder().buildCriteria(config);
|
||||
case YOLOV12_CUSTOM:
|
||||
return new YoloCriteriaBuilder().buildCriteria(config);
|
||||
// 其他类型
|
||||
default:
|
||||
return new DJLModelCriteriaBuilder().buildCriteria(config);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package cn.smartjavaai.objectdetection.criteria;
|
||||
|
||||
import ai.djl.modality.cv.Image;
|
||||
import ai.djl.modality.cv.output.DetectedObjects;
|
||||
import ai.djl.repository.zoo.Criteria;
|
||||
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
|
||||
|
||||
/**
|
||||
* 模型加载策略接口,用于根据不同模型类型构建对应的 DJL Criteria 实例
|
||||
* @author dwj
|
||||
* @date 2025/5/14
|
||||
*/
|
||||
public interface CriteriaBuilderStrategy {
|
||||
|
||||
/**
|
||||
* 根据模型类型构建对应的 DJL Criteria 实例
|
||||
* @param config
|
||||
* @return
|
||||
*/
|
||||
Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config);
|
||||
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package cn.smartjavaai.objectdetection.criteria;
|
||||
|
||||
import ai.djl.Application;
|
||||
import ai.djl.Device;
|
||||
import ai.djl.modality.cv.Image;
|
||||
import ai.djl.modality.cv.output.DetectedObjects;
|
||||
import ai.djl.repository.zoo.Criteria;
|
||||
import ai.djl.training.util.ProgressBar;
|
||||
import cn.smartjavaai.common.enums.DeviceEnum;
|
||||
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
|
||||
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
|
||||
|
||||
import java.util.Objects;
|
||||
|
||||
/**
|
||||
* DJL提供的Criteria 构建器
|
||||
* @author dwj
|
||||
* @date 2025/5/14
|
||||
*/
|
||||
public class DJLModelCriteriaBuilder implements CriteriaBuilderStrategy {
|
||||
|
||||
private static final String DJL_MODEL_PREFIX = "djl://";
|
||||
|
||||
@Override
|
||||
public Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config) {
|
||||
Device device = null;
|
||||
if(!Objects.isNull(config.getDevice())){
|
||||
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
|
||||
}
|
||||
Criteria<Image, DetectedObjects> criteria = Criteria.builder()
|
||||
.optApplication(Application.CV.OBJECT_DETECTION)
|
||||
.setTypes(Image.class, DetectedObjects.class)
|
||||
.optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : DetectorConstant.DEFAULT_THRESHOLD)
|
||||
.optModelUrls(DJL_MODEL_PREFIX + config.getModelEnum().getModelUri())
|
||||
.optDevice(device)
|
||||
.optProgress(new ProgressBar())
|
||||
.build();
|
||||
return criteria;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package cn.smartjavaai.objectdetection.criteria;
|
||||
|
||||
import ai.djl.modality.cv.Image;
|
||||
import ai.djl.modality.cv.output.DetectedObjects;
|
||||
import ai.djl.modality.cv.translator.YoloV8TranslatorFactory;
|
||||
import ai.djl.repository.zoo.Criteria;
|
||||
import ai.djl.training.util.ProgressBar;
|
||||
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
|
||||
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
|
||||
|
||||
import java.nio.file.Paths;
|
||||
|
||||
/**
|
||||
* YOLO模型Criteria 构建器
|
||||
* @author dwj
|
||||
* @date 2025/5/14
|
||||
*/
|
||||
public class YoloCriteriaBuilder implements CriteriaBuilderStrategy {
|
||||
@Override
|
||||
public Criteria<Image, DetectedObjects> buildCriteria(DetectorModelConfig config) {
|
||||
Criteria.Builder criteriaBuilder = Criteria.builder()
|
||||
.setTypes(Image.class, DetectedObjects.class)
|
||||
//.optModelUrls("/Users/wenjie/Documents/develop/face_model/yolo")
|
||||
.optModelPath(Paths.get(config.getModelPath()))
|
||||
.optEngine("OnnxRuntime")
|
||||
.optArgument("width", 640) //将输入图像的宽度缩放为 640 像素
|
||||
.optArgument("height", 640)
|
||||
.optArgument("resize", true)
|
||||
.optArgument("toTensor", true)
|
||||
.optArgument("applyRatio", true)
|
||||
.optTranslatorFactory(new YoloV8TranslatorFactory())
|
||||
.optProgress(new ProgressBar())
|
||||
.optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : DetectorConstant.DEFAULT_THRESHOLD);
|
||||
if(config.getMaxBox() > 0){
|
||||
criteriaBuilder.optArgument("maxBox", config.getMaxBox());
|
||||
}
|
||||
Criteria<Image, DetectedObjects> criteria = criteriaBuilder.build();
|
||||
return criteria;
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
package cn.smartjavaai.objectdetection;
|
||||
package cn.smartjavaai.objectdetection.enums;
|
||||
|
||||
/**
|
||||
* 目标检测模型枚举
|
||||
@@ -34,7 +34,14 @@ public enum DetectorModelEnum {
|
||||
YOLO3_DARKNET_COCO_608("ai.djl.mxnet/yolo/0.0.1/yolo3_darknet_coco_608"),
|
||||
YOLO3_MOBILENET_COCO_320("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_320"),
|
||||
YOLO3_MOBILENET_COCO_416("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_416"),
|
||||
YOLO3_MOBILENET_COCO_608("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_608");
|
||||
YOLO3_MOBILENET_COCO_608("ai.djl.mxnet/yolo/0.0.1/yolo3_mobilenet_coco_608"),
|
||||
|
||||
YOLOV12_OFFICIAL(""),
|
||||
YOLOV8_OFFICIAL(""),
|
||||
|
||||
YOLOV8_CUSTOM(""),
|
||||
|
||||
YOLOV12_CUSTOM("");
|
||||
|
||||
/**
|
||||
* 根据名称获取枚举 (忽略大小写和下划线变体)
|
||||
@@ -1,41 +1,34 @@
|
||||
package cn.smartjavaai.objectdetection.model;
|
||||
|
||||
import ai.djl.Application;
|
||||
import ai.djl.Device;
|
||||
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.opencv.OpenCVImageFactory;
|
||||
import ai.djl.modality.cv.translator.YoloV8TranslatorFactory;
|
||||
import ai.djl.repository.zoo.Criteria;
|
||||
import ai.djl.repository.zoo.ModelNotFoundException;
|
||||
import ai.djl.repository.zoo.ZooModel;
|
||||
import ai.djl.training.util.ProgressBar;
|
||||
import ai.djl.translate.TranslateException;
|
||||
import cn.smartjavaai.common.entity.DetectionResponse;
|
||||
import cn.smartjavaai.common.enums.DeviceEnum;
|
||||
import cn.smartjavaai.common.pool.ModelPredictorPoolManager;
|
||||
import cn.smartjavaai.common.pool.PredictorFactory;
|
||||
import cn.smartjavaai.common.utils.FileUtils;
|
||||
import cn.smartjavaai.common.utils.ImageUtils;
|
||||
import cn.smartjavaai.common.utils.OpenCVUtils;
|
||||
import cn.smartjavaai.objectdetection.DetectorConfig;
|
||||
import cn.smartjavaai.objectdetection.DetectorModelConfig;
|
||||
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
|
||||
import cn.smartjavaai.objectdetection.constant.DetectorConstant;
|
||||
import cn.smartjavaai.objectdetection.criteria.CriteriaBuilderFactory;
|
||||
import cn.smartjavaai.objectdetection.exception.DetectionException;
|
||||
import cn.smartjavaai.objectdetection.utils.DetectorUtils;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.apache.commons.lang3.Validate;
|
||||
import org.apache.commons.lang3.StringUtils;
|
||||
import org.apache.commons.pool2.ObjectPool;
|
||||
import org.apache.commons.pool2.impl.GenericObjectPool;
|
||||
import org.apache.commons.pool2.impl.GenericObjectPoolConfig;
|
||||
|
||||
import javax.imageio.ImageIO;
|
||||
import java.awt.image.BufferedImage;
|
||||
import java.io.*;
|
||||
import java.nio.file.Path;
|
||||
import java.nio.file.Paths;
|
||||
import java.time.Duration;
|
||||
import java.util.Objects;
|
||||
|
||||
/**
|
||||
@@ -50,23 +43,16 @@ public class DetectorModel implements AutoCloseable{
|
||||
|
||||
//private Predictor<Image, DetectedObjects> predictor;
|
||||
|
||||
private static final String DJL_MODEL_PREFIX = "djl://";
|
||||
|
||||
|
||||
private ObjectPool<Predictor<Image, DetectedObjects>> predictorPool;
|
||||
|
||||
public void loadModel(DetectorModelConfig config){
|
||||
Device device = null;
|
||||
if(!Objects.isNull(config.getDevice())){
|
||||
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu();
|
||||
if(Objects.isNull(config.getModelEnum())){
|
||||
throw new DetectionException("未配置模型枚举");
|
||||
}
|
||||
Criteria<Image, DetectedObjects> criteria = Criteria.builder()
|
||||
.optApplication(Application.CV.OBJECT_DETECTION)
|
||||
.setTypes(Image.class, DetectedObjects.class)
|
||||
.optArgument("threshold", config.getThreshold() > 0 ? config.getThreshold() : DetectorConfig.DEFAULT_THRESHOLD)
|
||||
.optModelUrls(DJL_MODEL_PREFIX + config.getModelEnum().getModelUri())
|
||||
.optDevice(device)
|
||||
.optProgress(new ProgressBar())
|
||||
.build();
|
||||
Criteria<Image, DetectedObjects> criteria = CriteriaBuilderFactory.createCriteria(config);
|
||||
|
||||
try {
|
||||
model = criteria.loadModel();
|
||||
// 创建池子:每个线程独享 Predictor
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
package cn.smartjavaai.objectdetection.model;
|
||||
|
||||
import cn.smartjavaai.common.config.Config;
|
||||
import cn.smartjavaai.objectdetection.DetectorModelConfig;
|
||||
import cn.smartjavaai.objectdetection.DetectorModelEnum;
|
||||
import cn.smartjavaai.objectdetection.config.DetectorModelConfig;
|
||||
import cn.smartjavaai.objectdetection.enums.DetectorModelEnum;
|
||||
import cn.smartjavaai.objectdetection.exception.DetectionException;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
|
||||
import java.util.Map;
|
||||
import java.util.Objects;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user