mirror of
https://github.com/geekwenjie/SmartJavaAI.git
synced 2026-09-13 13:18:58 +00:00
优化人脸检测速度
This commit is contained in:
@@ -33,6 +33,10 @@ public class FeatureExtractionAlgo extends AbstractFaceAlgorithm {
|
||||
|
||||
private Criteria<Image, float[]> faceFeatureCriteria;
|
||||
|
||||
private Predictor<Image, float[]> predictor;
|
||||
|
||||
private ZooModel<Image, float[]> model;
|
||||
|
||||
public static final List<Float> mean =
|
||||
Arrays.asList(
|
||||
127.5f / 255.0f,
|
||||
@@ -62,6 +66,8 @@ public class FeatureExtractionAlgo extends AbstractFaceAlgorithm {
|
||||
.optProgress(new ProgressBar())
|
||||
.optEngine("PyTorch") // Use PyTorch engine
|
||||
.build();
|
||||
model = faceFeatureCriteria.loadModel();
|
||||
predictor = model.newPredictor();
|
||||
}
|
||||
|
||||
|
||||
@@ -76,10 +82,7 @@ public class FeatureExtractionAlgo extends AbstractFaceAlgorithm {
|
||||
Path imageFile = Paths.get(imagePath);
|
||||
Image img = ImageFactory.getInstance().fromFile(imageFile);
|
||||
img.getWrappedImage();
|
||||
try (ZooModel<Image, float[]> model = faceFeatureCriteria.loadModel()) {
|
||||
Predictor<Image, float[]> predictor = model.newPredictor();
|
||||
return predictor.predict(img);
|
||||
}
|
||||
return predictor.predict(img);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -92,10 +95,7 @@ public class FeatureExtractionAlgo extends AbstractFaceAlgorithm {
|
||||
public float[] featureExtraction(InputStream inputStream) throws Exception {
|
||||
Image img = ImageFactory.getInstance().fromInputStream(inputStream);
|
||||
img.getWrappedImage();
|
||||
try (ZooModel<Image, float[]> model = faceFeatureCriteria.loadModel()) {
|
||||
Predictor<Image, float[]> predictor = model.newPredictor();
|
||||
return predictor.predict(img);
|
||||
}
|
||||
return predictor.predict(img);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -38,6 +38,10 @@ public class RetinaFace extends AbstractFaceAlgorithm {
|
||||
|
||||
private Criteria<Image, float[]> faceFeatureCriteria;
|
||||
|
||||
private Predictor<Image, DetectedObjects> predictor;
|
||||
|
||||
private ZooModel<Image, DetectedObjects> model;
|
||||
|
||||
/**
|
||||
* 特征图层的基础缩放比例
|
||||
*/
|
||||
@@ -57,7 +61,7 @@ public class RetinaFace extends AbstractFaceAlgorithm {
|
||||
* @param config
|
||||
*/
|
||||
@Override
|
||||
public void loadModel(ModelConfig config) {
|
||||
public void loadModel(ModelConfig config) throws ModelNotFoundException, MalformedModelException, IOException {
|
||||
FaceDetectionTranslator translator =
|
||||
new FaceDetectionTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), variance, config.getMaxFaceCount(), scales, steps);
|
||||
criteria =
|
||||
@@ -71,6 +75,8 @@ public class RetinaFace extends AbstractFaceAlgorithm {
|
||||
.optProgress(new ProgressBar())
|
||||
.optEngine("PyTorch") // Use PyTorch engine
|
||||
.build();
|
||||
model = criteria.loadModel();
|
||||
predictor = model.newPredictor();
|
||||
}
|
||||
|
||||
|
||||
@@ -85,11 +91,8 @@ public class RetinaFace extends AbstractFaceAlgorithm {
|
||||
public FaceDetectedResult detect(String imagePath) throws Exception{
|
||||
Path facePath = Paths.get(imagePath);
|
||||
Image img = ImageFactory.getInstance().fromFile(facePath);
|
||||
try (ZooModel<Image, DetectedObjects> model = criteria.loadModel();
|
||||
Predictor<Image, DetectedObjects> predictor = model.newPredictor()) {
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
}
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -101,13 +104,8 @@ public class RetinaFace extends AbstractFaceAlgorithm {
|
||||
@Override
|
||||
public FaceDetectedResult detect(InputStream imageInputStream) throws Exception {
|
||||
Image img = ImageFactory.getInstance().fromInputStream(imageInputStream);
|
||||
try (ZooModel<Image, DetectedObjects> model = criteria.loadModel();
|
||||
Predictor<Image, DetectedObjects> predictor = model.newPredictor()) {
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
/*saveBoundingBoxImage(img, detection);
|
||||
return detection;*/
|
||||
}
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -1,17 +1,20 @@
|
||||
package cn.smartjavaai.face.algo;
|
||||
|
||||
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.modality.cv.translator.ImageFeatureExtractorFactory;
|
||||
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 cn.smartjavaai.common.entity.Point;
|
||||
import cn.smartjavaai.common.entity.Rectangle;
|
||||
import cn.smartjavaai.face.*;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.InputStream;
|
||||
import java.nio.file.Path;
|
||||
import java.nio.file.Paths;
|
||||
@@ -41,6 +44,10 @@ public class UltraLightFastGenericFace extends AbstractFaceAlgorithm {
|
||||
*/
|
||||
private static final double[] variance = {0.1f, 0.2f};
|
||||
|
||||
private Predictor<Image, DetectedObjects> predictor;
|
||||
|
||||
private ZooModel<Image, DetectedObjects> model;
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -49,7 +56,7 @@ public class UltraLightFastGenericFace extends AbstractFaceAlgorithm {
|
||||
* @param config
|
||||
*/
|
||||
@Override
|
||||
public void loadModel(ModelConfig config) {
|
||||
public void loadModel(ModelConfig config) throws ModelNotFoundException, MalformedModelException, IOException {
|
||||
FaceDetectionTranslator translator =
|
||||
new FaceDetectionTranslator(config.getConfidenceThreshold(), config.getNmsThresh(), variance, config.getMaxFaceCount(), scales, steps);
|
||||
criteria =
|
||||
@@ -60,6 +67,8 @@ public class UltraLightFastGenericFace extends AbstractFaceAlgorithm {
|
||||
.optProgress(new ProgressBar())
|
||||
.optEngine("PyTorch") // Use PyTorch engine
|
||||
.build();
|
||||
model = criteria.loadModel();
|
||||
predictor = model.newPredictor();
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -72,11 +81,8 @@ public class UltraLightFastGenericFace extends AbstractFaceAlgorithm {
|
||||
public FaceDetectedResult detect(String imagePath) throws Exception{
|
||||
Path facePath = Paths.get(imagePath);
|
||||
Image img = ImageFactory.getInstance().fromFile(facePath);
|
||||
try (ZooModel<Image, DetectedObjects> model = criteria.loadModel();
|
||||
Predictor<Image, DetectedObjects> predictor = model.newPredictor()) {
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
}
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -88,13 +94,8 @@ public class UltraLightFastGenericFace extends AbstractFaceAlgorithm {
|
||||
@Override
|
||||
public FaceDetectedResult detect(InputStream imageInputStream) throws Exception {
|
||||
Image img = ImageFactory.getInstance().fromInputStream(imageInputStream);
|
||||
try (ZooModel<Image, DetectedObjects> model = criteria.loadModel();
|
||||
Predictor<Image, DetectedObjects> predictor = model.newPredictor()) {
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
/*saveBoundingBoxImage(img, detection);
|
||||
return detection;*/
|
||||
}
|
||||
DetectedObjects detection = predictor.predict(img);
|
||||
return convertToFaceDetectedResult(detection,img);
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
Reference in New Issue
Block a user