mirror of
https://github.com/geekwenjie/SmartJavaAI.git
synced 2026-09-10 03:28:49 +00:00
- 【人脸检测】新增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:
@@ -1,5 +1,6 @@
|
||||
package cn.smartjavaai.face.model.facedect;
|
||||
|
||||
import ai.djl.Device;
|
||||
import ai.djl.MalformedModelException;
|
||||
import ai.djl.engine.Engine;
|
||||
import ai.djl.inference.Predictor;
|
||||
@@ -18,6 +19,7 @@ import ai.djl.translate.NoopTranslator;
|
||||
import cn.smartjavaai.common.entity.*;
|
||||
import cn.smartjavaai.common.entity.Point;
|
||||
import cn.smartjavaai.common.entity.face.FaceInfo;
|
||||
import cn.smartjavaai.common.enums.DeviceEnum;
|
||||
import cn.smartjavaai.common.pool.PredictorFactory;
|
||||
import cn.smartjavaai.common.utils.*;
|
||||
import cn.smartjavaai.face.config.FaceDetConfig;
|
||||
@@ -78,12 +80,16 @@ public class MtcnnFaceDetModel implements FaceDetModel{
|
||||
throw new FaceException("MTCNN 模型需要指定存放模型文件的目录路径");
|
||||
}
|
||||
try {
|
||||
Device device = null;
|
||||
if(!Objects.isNull(config.getDevice())){
|
||||
device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(config.getGpuId());
|
||||
}
|
||||
Path pnetPath = modelPath.resolve("pnet_script.pt");
|
||||
Path rnetPath = modelPath.resolve("rnet_script.pt");
|
||||
Path onetPath = modelPath.resolve("onet_script.pt");
|
||||
pNetModel = getModel(pnetPath);
|
||||
rNetModel = getModel(rnetPath);
|
||||
oNetModel = getModel(onetPath);
|
||||
pNetModel = getModel(pnetPath, device);
|
||||
rNetModel = getModel(rnetPath, device);
|
||||
oNetModel = getModel(onetPath, device);
|
||||
|
||||
this.pnetPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(pNetModel));
|
||||
this.rnetPredictorPool = new GenericObjectPool<>(new PredictorFactory<>(rNetModel));
|
||||
@@ -110,7 +116,7 @@ public class MtcnnFaceDetModel implements FaceDetModel{
|
||||
* @throws MalformedModelException
|
||||
* @throws IOException
|
||||
*/
|
||||
public ZooModel<NDList, NDList> getModel(Path modelPath) throws ModelNotFoundException, MalformedModelException, IOException {
|
||||
public ZooModel<NDList, NDList> getModel(Path modelPath, Device device) throws ModelNotFoundException, MalformedModelException, IOException {
|
||||
Criteria<NDList, NDList> criteria =
|
||||
Criteria.builder()
|
||||
.setTypes(NDList.class, NDList.class)
|
||||
@@ -118,6 +124,7 @@ public class MtcnnFaceDetModel implements FaceDetModel{
|
||||
.optEngine("PyTorch")
|
||||
.optModelPath(modelPath)
|
||||
.optProgress(new ProgressBar())
|
||||
.optDevice(device)
|
||||
.build();
|
||||
return criteria.loadModel();
|
||||
}
|
||||
|
||||
@@ -31,9 +31,6 @@ 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.bytedeco.javacv.FFmpegFrameGrabber;
|
||||
import org.bytedeco.javacv.Frame;
|
||||
import org.bytedeco.javacv.Java2DFrameUtils;
|
||||
import org.opencv.core.Mat;
|
||||
|
||||
import javax.imageio.ImageIO;
|
||||
|
||||
Reference in New Issue
Block a user