Files
SmartJavaAI/vision/src/main/java/cn/smartjavaai/instanceseg/model/CommonInstanceSegModel.java
dengwenjie a8e7ce6c4e - 【人脸检测】新增6个模型(MTCNN、YOLOV5、RetinaFace小尺寸版),大幅提升性能
- 【人脸识别】新增Seetaface6轻量模型
- 【目标检测】支持视频流目标检测(rtsp、视频文件等)
- 【目标检测】支持tensorflow2目标检测模型
- 【目标检测】新增行人检测模型(yolo-person)
- 【通用视觉】新增4个动作识别模型
- 【通用视觉】新增语义分割模型
- 【通用视觉】新增5个实例分割模型(含yolov8-seg、yolov11-seg)
- 【通用视觉】新增yolo-obb11旋转框检测(含yolov11-obb)
- 【通用视觉】新增5个姿态估计模型(含yolov8-pose、yolov11-pose)
2025-09-07 17:19:19 +08:00

160 lines
6.2 KiB
Java

package cn.smartjavaai.instanceseg.model;
import ai.djl.MalformedModelException;
import ai.djl.engine.Engine;
import ai.djl.inference.Predictor;
import ai.djl.modality.cv.Image;
import ai.djl.modality.cv.output.DetectedObjects;
import ai.djl.repository.zoo.Criteria;
import ai.djl.repository.zoo.ModelNotFoundException;
import ai.djl.repository.zoo.ZooModel;
import cn.smartjavaai.common.cv.SmartImageFactory;
import cn.smartjavaai.common.entity.DetectionResponse;
import cn.smartjavaai.common.entity.R;
import cn.smartjavaai.common.pool.PredictorFactory;
import cn.smartjavaai.instanceseg.config.InstanceSegModelConfig;
import cn.smartjavaai.instanceseg.criteria.InstanceSegCriteriaFactory;
import cn.smartjavaai.instanceseg.exception.InstanceSegException;
import cn.smartjavaai.objectdetection.exception.DetectionException;
import cn.smartjavaai.vision.utils.DetectedObjectsFilter;
import cn.smartjavaai.vision.utils.DetectorUtils;
import cn.smartjavaai.vision.utils.CategoryMaskFilter;
import lombok.extern.slf4j.Slf4j;
import org.apache.commons.collections.CollectionUtils;
import org.apache.commons.pool2.impl.GenericObjectPool;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.Objects;
/**
* 实例分割模型
* @author dwj
*/
@Slf4j
public class CommonInstanceSegModel implements InstanceSegModel {
private InstanceSegModelConfig config;
private ZooModel<Image, DetectedObjects> model;
private GenericObjectPool<Predictor<Image, DetectedObjects>> predictorPool;
@Override
public void loadModel(InstanceSegModelConfig config) {
if(Objects.isNull(config.getModelEnum())){
throw new DetectionException("未配置模型枚举");
}
Criteria<Image, DetectedObjects> criteria = InstanceSegCriteriaFactory.createCriteria(config);
this.config = config;
try {
model = criteria.loadModel();
// 创建池子:每个线程独享 Predictor
this.predictorPool = new GenericObjectPool<>(new PredictorFactory<>(model));
int predictorPoolSize = config.getPredictorPoolSize();
if(config.getPredictorPoolSize() <= 0){
predictorPoolSize = Runtime.getRuntime().availableProcessors(); // 默认等于CPU核心数
}
predictorPool.setMaxTotal(predictorPoolSize);
log.debug("当前设备: " + model.getNDManager().getDevice());
log.debug("当前引擎: " + Engine.getInstance().getEngineName());
log.debug("模型推理器线程池最大数量: " + predictorPoolSize);
} catch (IOException | ModelNotFoundException | MalformedModelException e) {
throw new DetectionException("模型加载失败", e);
}
}
@Override
public R<DetectionResponse> detect(Image image) {
DetectedObjects detectedObjects = detectCore(image);
DetectionResponse detectionResponse = DetectorUtils.convertToDetectionResponse(detectedObjects, image);
return R.ok(detectionResponse);
}
/**
* 模型核心推理方法
* @param image
* @return
*/
@Override
public DetectedObjects detectCore(Image image) {
Predictor<Image, DetectedObjects> predictor = null;
try {
predictor = predictorPool.borrowObject();
DetectedObjects detectedObjects = predictor.predict(image);
//过滤
if(Objects.nonNull(detectedObjects) && detectedObjects.getNumberOfObjects() > 0){
DetectedObjectsFilter detectedObjectsFilter = new DetectedObjectsFilter(config.getAllowedClasses(), config.getThreshold());
detectedObjects = detectedObjectsFilter.filter(detectedObjects);
}
return detectedObjects;
} catch (Exception e) {
throw new DetectionException("实例分割错误", e);
}finally {
if (predictor != null) {
try {
predictorPool.returnObject(predictor); //归还
log.debug("释放资源");
} catch (Exception e) {
log.warn("归还Predictor失败", e);
try {
predictor.close(); // 归还失败才销毁
} catch (Exception ex) {
log.error("关闭Predictor失败", ex);
}
}
}
}
}
@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);
return R.ok(detectionResponse);
}
@Override
public R<DetectionResponse> detectAndDraw(String imagePath, String outputPath) {
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);
return R.ok(detectionResponse);
} catch (IOException e) {
throw new InstanceSegException(e);
}
}
@Override
public void close() throws Exception {
try {
if (predictorPool != null) {
predictorPool.close();
}
} catch (Exception e) {
log.warn("关闭 predictorPool 失败", e);
}
try {
if (model != null) {
model.close();
}
} catch (Exception e) {
log.warn("关闭 model 失败", e);
}
}
}