- 【人脸检测】新增6个模型(MTCNN、YOLOV5、RetinaFace小尺寸版),大幅提升性能

- 【人脸识别】新增Seetaface6轻量模型
- 【目标检测】支持视频流目标检测(rtsp、视频文件等)
- 【目标检测】支持tensorflow2目标检测模型
- 【目标检测】新增行人检测模型(yolo-person)
- 【通用视觉】新增4个动作识别模型
- 【通用视觉】新增语义分割模型
- 【通用视觉】新增5个实例分割模型(含yolov8-seg、yolov11-seg)
- 【通用视觉】新增yolo-obb11旋转框检测(含yolov11-obb)
- 【通用视觉】新增5个姿态估计模型(含yolov8-pose、yolov11-pose)
This commit is contained in:
dengwenjie
2025-09-07 17:19:19 +08:00
parent 2b044fda29
commit a8e7ce6c4e
102 changed files with 4473 additions and 1368 deletions

View File

@@ -36,7 +36,7 @@ public class InstanceSegModelConfig extends ModelConfig {
/**
* 置信度阈值
*/
private float threshold = 0.3f;
private float threshold = 0.25f;
public InstanceSegModelConfig() {

View File

@@ -3,10 +3,12 @@ package cn.smartjavaai.instanceseg.criteria;
import ai.djl.Device;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.modality.cv.translator.InstanceSegmentationTranslatorFactory;
import ai.djl.repository.zoo.Criteria;
import ai.djl.training.util.ProgressBar;
import cn.smartjavaai.common.enums.DeviceEnum;
import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig;
import cn.smartjavaai.instanceseg.enums.InstanceSegModelEnum;
import cn.smartjavaai.instanceseg.translator.YoloSegmentationTranslatorFactory2;
import org.apache.commons.lang3.StringUtils;
@@ -27,21 +29,42 @@ public class InstanceSegCriteriaFactory {
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
}
Criteria<Image, DetectedObjects> criteria = null;
ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
params.putAll(config.getCustomParams());
// ConcurrentHashMap params = new ConcurrentHashMap<String, String>();
// params.putAll(config.getCustomParams());
// YoloV5Translator.Builder builder = new YoloV5Translator.Builder()
// .optSynsetArtifactName("synset.txt").setPipeline()
criteria =
Criteria.builder()
.setTypes(Image.class, DetectedObjects.class)
.optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null :
config.getModelEnum().getModelUri())
.optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null)
.optDevice(device)
.optEngine("PyTorch")
.optTranslatorFactory(new YoloSegmentationTranslatorFactory2())
.optProgress(new ProgressBar())
.build();
if(config.getModelEnum() == InstanceSegModelEnum.SEG_MASK_RCNN){
criteria =
Criteria.builder()
.setTypes(Image.class, DetectedObjects.class)
.optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null :
config.getModelEnum().getModelUri())
.optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null)
.optDevice(device)
.optEngine(config.getModelEnum().getEngine())
.optArgument("normalize","true")
.optArgument("synsetFileName","classes.txt")
.optTranslatorFactory(new InstanceSegmentationTranslatorFactory())
.optProgress(new ProgressBar())
.build();
}else{
criteria =
Criteria.builder()
.setTypes(Image.class, DetectedObjects.class)
.optModelUrls(StringUtils.isNotBlank(config.getModelPath()) ? null :
config.getModelEnum().getModelUri())
.optModelPath(StringUtils.isNotBlank(config.getModelPath()) ? Paths.get(config.getModelPath()) : null)
.optDevice(device)
.optArgument("width", config.getModelEnum().getInputWidth())
.optArgument("height", config.getModelEnum().getInputHeight())
.optArgument("resize", "true")
.optArgument("threshold", config.getThreshold())
.optEngine(config.getModelEnum().getEngine())
.optTranslatorFactory(new YoloSegmentationTranslatorFactory2())
.optProgress(new ProgressBar())
.build();
}
return criteria;
}
}

View File

@@ -6,15 +6,15 @@ package cn.smartjavaai.instanceseg.enums;
*/
public enum InstanceSegModelEnum {
SEG_YOLO11N_PYTORCH("djl://ai.djl.pytorch/yolo11n-seg"),
SEG_YOLO11N_PYTORCH("PyTorch", 640, 640, "djl://ai.djl.pytorch/yolo11n-seg"),
SEG_YOLOV8N_PYTORCH("djl://ai.djl.pytorch/yolo11n-seg"),
SEG_YOLOV8N_PYTORCH("PyTorch", 640, 640, "djl://ai.djl.pytorch/yolov8n-seg"),
SEG_YOLO11N_ONNX("djl://ai.djl.onnxruntime/yolo11n-seg"),
SEG_YOLO11N_ONNX("OnnxRuntime", 640, 640, "djl://ai.djl.onnxruntime/yolo11n-seg"),
SEG_YOLOV8N_ONNX("djl://ai.djl.onnxruntime/yolov8n-seg"),
SEG_YOLOV8N_ONNX("OnnxRuntime", 640, 640, "djl://ai.djl.onnxruntime/yolov8n-seg"),
SEG_MASK_RCNN("djl://ai.djl.mxnet/mask_rcnn");
SEG_MASK_RCNN("MXNet", 0,0, "djl://ai.djl.mxnet/mask_rcnn");
/**
@@ -30,14 +30,43 @@ public enum InstanceSegModelEnum {
throw new IllegalArgumentException("未知模型名称: " + name);
}
/**
* 模型输入尺寸:宽
*/
private final int inputWidth;
/**
* 模型输入尺寸:高
*/
private final int inputHeight;
private final String modelUri;
InstanceSegModelEnum(String modelUri) {
/**
* 模型引擎
*/
private final String engine;
InstanceSegModelEnum(String engine, int inputWidth, int inputHeight, String modelUri) {
this.inputWidth = inputWidth;
this.inputHeight = inputHeight;
this.modelUri = modelUri;
this.engine = engine;
}
public String getModelUri() {
return modelUri;
}
public String getEngine() {
return engine;
}
public int getInputWidth() {
return inputWidth;
}
public int getInputHeight() {
return inputHeight;
}
}

View File

@@ -113,6 +113,9 @@ public class CommonInstanceSegModel implements InstanceSegModel {
@Override
public R<DetectionResponse> detectAndDraw(Image image) {
DetectedObjects detectedObjects = detectCore(image);
if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){
return R.fail(R.Status.NO_OBJECT_DETECTED);
}
image.drawBoundingBoxes(detectedObjects);
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
detectionResponse.setDrawnImage(image);
@@ -124,6 +127,9 @@ public class CommonInstanceSegModel implements InstanceSegModel {
try {
Image img = SmartImageFactory.getInstance().fromFile(Paths.get(imagePath));
DetectedObjects detectedObjects = detectCore(img);
if(Objects.isNull(detectedObjects) || detectedObjects.getNumberOfObjects() == 0){
return R.fail(R.Status.NO_OBJECT_DETECTED);
}
img.drawBoundingBoxes(detectedObjects);
img.save(Files.newOutputStream(Paths.get(outputPath)), "png");
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, img);

View File

@@ -0,0 +1,111 @@
package cn.smartjavaai.instanceseg.model;
import cn.smartjavaai.common.config.Config;
import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig;
import cn.smartjavaai.instanceseg.enums.InstanceSegModelEnum;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import lombok.extern.slf4j.Slf4j;
import java.util.Map;
import java.util.Objects;
import java.util.concurrent.ConcurrentHashMap;
/**
* 实例分割 模型工厂
* @author dwj
*/
@Slf4j
public class InstanceSegModelFactory {
// 使用 volatile 和双重检查锁定来确保线程安全的单例模式
private static volatile InstanceSegModelFactory instance;
private static final ConcurrentHashMap<InstanceSegModelEnum, InstanceSegModel> modelMap = new ConcurrentHashMap<>();
/**
* 模型注册表
*/
private static final Map<InstanceSegModelEnum, Class<? extends InstanceSegModel>> registry =
new ConcurrentHashMap<>();
// 私有构造函数,防止外部创建实例
private InstanceSegModelFactory() {}
// 双重检查锁定的单例方法
public static InstanceSegModelFactory getInstance() {
if (instance == null) {
synchronized (InstanceSegModelFactory.class) {
if (instance == null) {
instance = new InstanceSegModelFactory();
}
}
}
return instance;
}
/**
* 获取模型(通过配置)
* @param config
* @return
*/
public InstanceSegModel getModel(InstanceSegModelConfig config) {
if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型");
}
return modelMap.computeIfAbsent(config.getModelEnum(), k -> {
return createFaceDetModel(config);
});
}
/**
* 使用ModelConfig创建模型
* @param config
* @return
*/
private InstanceSegModel createFaceDetModel(InstanceSegModelConfig config) {
Class<?> clazz = registry.get(config.getModelEnum());
if(clazz == null){
throw new DetectionException("Unsupported model");
}
InstanceSegModel model = null;
try {
model = (InstanceSegModel) clazz.newInstance();
} catch (InstantiationException | IllegalAccessException e) {
throw new DetectionException(e);
}
model.loadModel(config);
return model;
}
/**
* 注册模型
* @param modelEnum
* @param clazz
*/
private static void registerAlgorithm(InstanceSegModelEnum modelEnum, Class<? extends InstanceSegModel> clazz) {
registry.put(modelEnum, clazz);
}
/**
* 移除缓存的模型
* @param modelEnum
*/
public static void removeFromCache(InstanceSegModelEnum modelEnum) {
modelMap.remove(modelEnum);
}
// 初始化默认算法
static {
registerAlgorithm(InstanceSegModelEnum.SEG_YOLOV8N_ONNX, CommonInstanceSegModel.class);
registerAlgorithm(InstanceSegModelEnum.SEG_YOLOV8N_PYTORCH, CommonInstanceSegModel.class);
registerAlgorithm(InstanceSegModelEnum.SEG_YOLO11N_PYTORCH, CommonInstanceSegModel.class);
registerAlgorithm(InstanceSegModelEnum.SEG_YOLO11N_ONNX, CommonInstanceSegModel.class);
registerAlgorithm(InstanceSegModelEnum.SEG_MASK_RCNN, CommonInstanceSegModel.class);
log.debug("缓存目录:{}", Config.getCachePath());
}
}