diff --git a/examples/face-example/src/main/java/smartai/examples/face/headpose/HeadPoseDetDemo.java b/examples/face-example/src/main/java/smartai/examples/face/headpose/HeadPoseDetDemo.java new file mode 100644 index 0000000..48d7fc1 --- /dev/null +++ b/examples/face-example/src/main/java/smartai/examples/face/headpose/HeadPoseDetDemo.java @@ -0,0 +1,208 @@ +package smartai.examples.face.headpose; + +import cn.smartjavaai.common.cv.SmartImageFactory; +import cn.smartjavaai.common.entity.DetectionInfo; +import cn.smartjavaai.common.entity.DetectionResponse; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.entity.face.HeadPose; +import cn.smartjavaai.face.config.FaceDetConfig; +import cn.smartjavaai.face.config.HeadPoseConfig; +import cn.smartjavaai.face.enums.FaceDetModelEnum; +import cn.smartjavaai.face.enums.HeadPoseModelEnum; +import cn.smartjavaai.face.factory.FaceDetModelFactory; +import cn.smartjavaai.face.factory.HeadPoseModelFactory; +import cn.smartjavaai.face.model.facedect.FaceDetModel; +import cn.smartjavaai.face.model.headpose.HeadPoseModel; +import ai.djl.modality.cv.Image; +import com.alibaba.fastjson.JSONObject; +import lombok.extern.slf4j.Slf4j; +import org.junit.BeforeClass; +import org.junit.Test; + +import java.io.IOException; +import java.util.List; + +/** + * 人脸姿态检测 demo + *
+ * 演示两种后端的使用方式: + * 1. SeetaFace6 PoseEstimator + * 2. SixDRepNet ONNX 模型 + *
+ * + * @author hyw + */ +@Slf4j +public class HeadPoseDetDemo { + + @BeforeClass + public static void beforeAll() throws IOException { + SmartImageFactory.setEngine(SmartImageFactory.Engine.OPENCV); + } + + /** + * 使用 SeetaFace6 进行人脸姿态检测(结合人脸检测) + */ + @Test + public void testSeetaFace6HeadPose() { + try { + // 需替换为实际模型存储路径 + String seetaModelPath = "C:/Users/DengWenJie/Downloads/sf3.0_models/sf3.0_models"; + + // 1. 创建人脸姿态检测模型(SeetaFace6) + HeadPoseConfig headPoseConfig = new HeadPoseConfig(); + headPoseConfig.setModelEnum(HeadPoseModelEnum.SEETA_FACE6_MODEL); + headPoseConfig.setModelPath(seetaModelPath); + HeadPoseModel headPoseModel = HeadPoseModelFactory.getInstance().getModel(headPoseConfig); + + // 2. 创建人脸检测模型 + FaceDetConfig faceDetConfig = new FaceDetConfig(); + faceDetConfig.setModelEnum(FaceDetModelEnum.SEETA_FACE6_MODEL); + faceDetConfig.setModelPath(seetaModelPath); + FaceDetModel faceDetModel = FaceDetModelFactory.getInstance().getModel(faceDetConfig); + + // 3. 检测人脸 + Image image = SmartImageFactory.getInstance().fromFile("src/main/resources/iu_1.jpg"); + R+ * 支持 SeetaFace6 和 SixDRepNet 两种后端实现,通过配置切换。 + *
+ * + * @author hyw + */ +@Slf4j +public class HeadPoseModelFactory { + + // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 + private static volatile HeadPoseModelFactory instance; + + private static final ConcurrentHashMap+ * 支持多种后端实现(SeetaFace6、SixDRepNet等),检测人脸的 pitch/yaw/roll 三个欧拉角。 + *
+ * + * @author hyw + */ +public interface HeadPoseModel extends AutoCloseable { + + /** + * 加载模型 + * @param config 模型配置 + */ + void loadModel(HeadPoseConfig config); + + /** + * 单人脸姿态检测(基于人脸框) + * @param image 原始图片 BufferedImage + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return 姿态结果(pitch/yaw/roll,单位:度) + */ + default HeadPose predict(BufferedImage image, DetectionRectangle faceDetectionRectangle) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 单人脸姿态检测(基于人脸框) + * @param imagePath 图片路径 + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return 姿态结果 + */ + default HeadPose predict(String imagePath, DetectionRectangle faceDetectionRectangle) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 单人脸姿态检测(基于人脸框) + * @param imageData 图片字节流 + * @param faceDetectionRectangle 人脸检测结果-人脸框 + * @return 姿态结果 + */ + default HeadPose predict(byte[] imageData, DetectionRectangle faceDetectionRectangle) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 单人脸姿态检测(裁剪后的人脸) + * @param croppedFace 已裁剪的人脸图片 BufferedImage + * @return 姿态结果 + */ + default HeadPose predictCropedFace(BufferedImage croppedFace) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 单人脸姿态检测(裁剪后的人脸) + * @param imagePath 已裁剪的人脸图片路径 + * @return 姿态结果 + */ + default HeadPose predictCropedFace(String imagePath) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 单人脸姿态检测(裁剪后的人脸) + * @param imageData 已裁剪的人脸图片字节流 + * @return 姿态结果 + */ + default HeadPose predictCropedFace(byte[] imageData) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + /** + * 多人脸姿态检测(基于已检测结果) + * @param image 原始图片 BufferedImage + * @param faceDetectionResponse 人脸检测结果 + * @return 每张人脸的姿态结果列表 + */ + default List+ * 基于 SeetaFace6 的 PoseEstimator 实现人脸姿态(pitch/yaw/roll)检测。 + *
+ * + * @author hyw + * @date 2026/8/25 + */ +@Slf4j +public class Seetaface6HeadPoseModel implements HeadPoseModel { + + private PoseEstimatorPool poseEstimatorPool; + + private HeadPoseConfig config; + + private boolean fromFactory = false; + + @Override + public void loadModel(HeadPoseConfig config) { + if (StringUtils.isBlank(config.getModelPath())) { + throw new FaceException("modelPath is null"); + } + this.config = config; + // 加载 SeetaFace6 依赖库 + NativeLoader.loadNativeLibraries(config.getDevice()); + log.debug("Loading seetaFace6 library successfully."); + + String[] poseEstimatorModelPath = {config.getModelPath() + File.separator + "pose_estimation.csta"}; + SeetaDevice device = SeetaDevice.SEETA_DEVICE_AUTO; + int gpuId = 0; + if (Objects.nonNull(config.getDevice())) { + device = config.getDevice() == DeviceEnum.CPU ? SeetaDevice.SEETA_DEVICE_CPU : SeetaDevice.SEETA_DEVICE_GPU; + if (config.getGpuId() >= 0 && device == SeetaDevice.SEETA_DEVICE_GPU) { + gpuId = config.getGpuId(); + } + } + + try { + SeetaModelSetting poseEstimatorPoolSetting = new SeetaModelSetting(gpuId, poseEstimatorModelPath, device); + SeetaConfSetting poseEstimatorPoolConfSetting = new SeetaConfSetting(poseEstimatorPoolSetting); + this.poseEstimatorPool = new PoseEstimatorPool(poseEstimatorPoolConfSetting); + + int predictorPoolSize = config.getPredictorPoolSize(); + if (predictorPoolSize <= 0) { + predictorPoolSize = Runtime.getRuntime().availableProcessors(); + } + poseEstimatorPool.setMaxTotal(predictorPoolSize); + log.debug("SeetaFace6 HeadPose 模型推理器线程池最大数量: {}", predictorPoolSize); + } catch (FileNotFoundException e) { + throw new FaceException(e); + } + } + + @Override + public HeadPose predict(BufferedImage image, DetectionRectangle faceDetectionRectangle) { + if (!BufferedImageUtils.isImageValid(image)) { + throw new FaceException("图像无效"); + } + if (Objects.isNull(faceDetectionRectangle)) { + throw new FaceException("无人脸数据"); + } + PoseEstimator poseEstimator = null; + try { + poseEstimator = poseEstimatorPool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(image.getWidth(), image.getHeight(), 3); + imageData.data = BufferedImageUtils.getMatrixBGR(image); + SeetaRect seetaRect = Seetaface6Utils.convertToSeetaRect(faceDetectionRectangle); + return estimatePose(poseEstimator, imageData, seetaRect); + } catch (Exception e) { + throw new FaceException("人脸姿态检测错误", e); + } finally { + PoolUtils.returnToPool(poseEstimatorPool, poseEstimator); + } + } + + @Override + public HeadPose predict(String imagePath, DetectionRectangle faceDetectionRectangle) { + if (!FileUtils.isFileExists(imagePath)) { + throw new FaceException("图像文件不存在"); + } + BufferedImage image; + try { + image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + } catch (IOException e) { + throw new FaceException("无效图片路径", e); + } + return predict(image, faceDetectionRectangle); + } + + @Override + public HeadPose predict(byte[] imageData, DetectionRectangle faceDetectionRectangle) { + if (Objects.isNull(imageData)) { + throw new FaceException("图像无效"); + } + try { + return predict(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionRectangle); + } catch (IOException e) { + throw new FaceException("错误的图像", e); + } + } + + @Override + public HeadPose predictCropedFace(BufferedImage croppedFace) { + if (!BufferedImageUtils.isImageValid(croppedFace)) { + throw new FaceException("图像无效"); + } + PoseEstimator poseEstimator = null; + try { + poseEstimator = poseEstimatorPool.borrowObject(); + SeetaImageData imageData = new SeetaImageData(croppedFace.getWidth(), croppedFace.getHeight(), 3); + imageData.data = BufferedImageUtils.getMatrixBGR(croppedFace); + // 裁剪后的人脸使用全图区域作为人脸框 + SeetaRect seetaRect = new SeetaRect(); + seetaRect.x = 0; + seetaRect.y = 0; + seetaRect.width = croppedFace.getWidth(); + seetaRect.height = croppedFace.getHeight(); + return estimatePose(poseEstimator, imageData, seetaRect); + } catch (Exception e) { + throw new FaceException("人脸姿态检测错误", e); + } finally { + PoolUtils.returnToPool(poseEstimatorPool, poseEstimator); + } + } + + @Override + public HeadPose predictCropedFace(String imagePath) { + if (!FileUtils.isFileExists(imagePath)) { + throw new FaceException("图像文件不存在"); + } + BufferedImage image; + try { + image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + } catch (IOException e) { + throw new FaceException("无效图片路径", e); + } + return predictCropedFace(image); + } + + @Override + public HeadPose predictCropedFace(byte[] imageData) { + if (Objects.isNull(imageData)) { + throw new FaceException("图像无效"); + } + try { + return predictCropedFace(ImageIO.read(new ByteArrayInputStream(imageData))); + } catch (IOException e) { + throw new FaceException("错误的图像", e); + } + } + + @Override + public List+ * 基于 6DRepNet 的 ONNX 模型实现人脸姿态(pitch/yaw/roll)检测。 + * 模型输出 3x3 旋转矩阵,通过后处理转换为欧拉角。 + *
+ *+ * 预处理流程对应 Python 端的 torchvision transforms: + * Resize(224) → CenterCrop(224) → ToTensor → Normalize(ImageNet) + *
+ * + * @author hyw + * @date 2026/8/25 + */ +@Slf4j +public class SixDrepNetHeadPoseModel implements HeadPoseModel { + + // ImageNet 归一化参数 + private static final float[] MEAN = {0.485f, 0.456f, 0.406f}; + private static final float[] STD = {0.229f, 0.224f, 0.225f}; + private static final int INPUT_SIZE = 224; + + private OrtSession session; + private OrtEnvironment env; + + private HeadPoseConfig config; + private boolean fromFactory = false; + + @Override + public void loadModel(HeadPoseConfig config) { + if (StringUtils.isBlank(config.getModelPath())) { + throw new FaceException("modelPath is null"); + } + this.config = config; + + try { + env = OrtEnvironment.getEnvironment(); + OrtSession.SessionOptions opts = new OrtSession.SessionOptions(); + + // GPU 支持需要 onnxruntime-gpu 依赖 + if (Objects.nonNull(config.getDevice()) && config.getDevice() == cn.smartjavaai.common.enums.DeviceEnum.GPU) { + opts.addCUDA(config.getGpuId() >= 0 ? config.getGpuId() : 0); + log.debug("SixDRepNet 使用 GPU 模式, gpuId={}", config.getGpuId()); + } else { + log.debug("SixDRepNet 使用 CPU 模式"); + } + + session = env.createSession(config.getModelPath(), opts); + log.info("SixDRepNet ONNX 模型已加载: {}", config.getModelPath()); + } catch (OrtException e) { + throw new FaceException("加载 SixDRepNet ONNX 模型失败: " + config.getModelPath(), e); + } + } + + @Override + public HeadPose predict(BufferedImage image, DetectionRectangle faceDetectionRectangle) { + if (!BufferedImageUtils.isImageValid(image)) { + throw new FaceException("图像无效"); + } + if (Objects.isNull(faceDetectionRectangle)) { + throw new FaceException("无人脸数据"); + } + try { + // 根据人脸框裁剪人脸区域 + BufferedImage croppedFace = cropFace(image, faceDetectionRectangle); + return predictInternal(croppedFace); + } catch (FaceException e) { + throw e; + } catch (Exception e) { + throw new FaceException("人脸姿态检测错误", e); + } + } + + @Override + public HeadPose predict(String imagePath, DetectionRectangle faceDetectionRectangle) { + if (!FileUtils.isFileExists(imagePath)) { + throw new FaceException("图像文件不存在"); + } + BufferedImage image; + try { + image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + } catch (IOException e) { + throw new FaceException("无效图片路径", e); + } + return predict(image, faceDetectionRectangle); + } + + @Override + public HeadPose predict(byte[] imageData, DetectionRectangle faceDetectionRectangle) { + if (Objects.isNull(imageData)) { + throw new FaceException("图像无效"); + } + try { + return predict(ImageIO.read(new ByteArrayInputStream(imageData)), faceDetectionRectangle); + } catch (IOException e) { + throw new FaceException("错误的图像", e); + } + } + + @Override + public HeadPose predictCropedFace(BufferedImage croppedFace) { + if (!BufferedImageUtils.isImageValid(croppedFace)) { + throw new FaceException("图像无效"); + } + try { + return predictInternal(croppedFace); + } catch (FaceException e) { + throw e; + } catch (Exception e) { + throw new FaceException("人脸姿态检测错误", e); + } + } + + @Override + public HeadPose predictCropedFace(String imagePath) { + if (!FileUtils.isFileExists(imagePath)) { + throw new FaceException("图像文件不存在"); + } + BufferedImage image; + try { + image = ImageIO.read(new File(Paths.get(imagePath).toAbsolutePath().toString())); + } catch (IOException e) { + throw new FaceException("无效图片路径", e); + } + return predictCropedFace(image); + } + + @Override + public HeadPose predictCropedFace(byte[] imageData) { + if (Objects.isNull(imageData)) { + throw new FaceException("图像无效"); + } + try { + return predictCropedFace(ImageIO.read(new ByteArrayInputStream(imageData))); + } catch (IOException e) { + throw new FaceException("错误的图像", e); + } + } + + @Override + public List+ * 旋转顺序: X(pitch) → Y(yaw) → Z(roll) + * + * @param R 3x3 旋转矩阵 + * @return float[3],分别为 pitch, yaw, roll(弧度) + */ + private float[] rotationMatrixToEuler(float[][] R) { + float sy = (float) Math.sqrt(R[0][0] * R[0][0] + R[1][0] * R[1][0]); + boolean singular = sy < 1e-6f; + + float x, y, z; + if (!singular) { + x = (float) Math.atan2(R[2][1], R[2][2]); // pitch + y = (float) Math.atan2(-R[2][0], sy); // yaw + z = (float) Math.atan2(R[1][0], R[0][0]); // roll + } else { + x = (float) Math.atan2(-R[1][2], R[1][1]); // pitch + y = (float) Math.atan2(-R[2][0], sy); // yaw + z = 0; // roll + } + + return new float[]{x, y, z}; + } + + /** + * 根据人脸框从原图裁剪人脸区域 + * + * @param image 原始图片 + * @param rect 人脸检测框 + * @return 裁剪后的人脸图片 + */ + private BufferedImage cropFace(BufferedImage image, DetectionRectangle rect) { + int x = Math.max(0, rect.getX()); + int y = Math.max(0, rect.getY()); + int w = Math.min(rect.getWidth(), image.getWidth() - x); + int h = Math.min(rect.getHeight(), image.getHeight() - y); + if (w <= 0 || h <= 0) { + throw new FaceException("人脸框区域无效"); + } + return image.getSubimage(x, y, w, h); + } + + @Override + public void setFromFactory(boolean fromFactory) { + this.fromFactory = fromFactory; + } + + public boolean isFromFactory() { + return fromFactory; + } + + @Override + public void close() throws Exception { + if (fromFactory) { + HeadPoseModelFactory.removeFromCache(config.getModelEnum()); + } + if (Objects.nonNull(session)) { + session.close(); + } + if (Objects.nonNull(env)) { + env.close(); + } + } + +}