1、目标检测:支持自己训练的模型推理

2、目标检测:支持yolo12模型
3、支持JDK8使用
4、引入离线依赖库,支持完全离线使用
5、优化FaceNet人脸比对速度
6、支持4通道图片检测
This commit is contained in:
dengwenjie
2025-05-17 11:19:46 +08:00
parent ab669d58e3
commit a7a118c5aa
27 changed files with 401 additions and 237 deletions

View File

@@ -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() {
}

View File

@@ -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;
}

View File

@@ -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);
}
}
}

View File

@@ -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);
}

View File

@@ -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;
}
}

View File

@@ -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;
}
}

View File

@@ -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("");
/**
* 根据名称获取枚举 (忽略大小写和下划线变体)

View File

@@ -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

View File

@@ -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;