Train.java 2.8 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283
  1. package com.example.demo.Util;
  2. import org.opencv.core.Mat;
  3. import org.opencv.core.MatOfInt;
  4. import org.opencv.face.FaceRecognizer;
  5. import org.opencv.face.LBPHFaceRecognizer;
  6. import org.opencv.imgcodecs.Imgcodecs;
  7. import org.opencv.imgproc.Imgproc;
  8. import org.opencv.objdetect.CascadeClassifier;
  9. import java.io.File;
  10. import java.io.IOException;
  11. import java.util.ArrayList;
  12. import java.util.HashMap;
  13. import java.util.List;
  14. import java.util.Map;
  15. public class Train {
  16. /*public static void main(String[] args) throws IOException {
  17. trainPath();
  18. }*/
  19. public static void trainPath() throws IOException {
  20. train("D:\\demos\\src\\main\\resources\\static\\imagedb",
  21. "D:\\demos\\src\\main\\resources\\static\\model");
  22. }
  23. /**
  24. * 训练模型的方法,传入人脸图片所在的文件夹路径,和模型输出的路径
  25. * 训练结束后模型文件会在模型输出路径里边
  26. **/
  27. public static void train(String imageFolder, String saveFolder)
  28. throws IOException {
  29. System.loadLibrary("opencv_java410");
  30. FaceRecognizer faceRecognizer = LBPHFaceRecognizer.create();
  31. CascadeClassifier faceCascade = new CascadeClassifier();
  32. // opencv的模型
  33. faceCascade.load("D:/opencv/build/etc/haarcascades/haarcascade_frontalface_alt.xml");
  34. // 读取文件于数组中
  35. File[] files = new File(imageFolder).listFiles();
  36. Map<String, Integer> nameMapId = new HashMap<String, Integer>(10);
  37. // 图片集合
  38. List<Mat> images = new ArrayList<Mat>(files.length);
  39. // 名称集合
  40. List<String> names = new ArrayList<String>(files.length);
  41. List<Integer> ids = new ArrayList<Integer>(files.length);
  42. for (int index = 0; index < files.length; index++ ) {
  43. // 解析文件名 获取名称
  44. File file = files[index];
  45. String name = file.getName().split("\\.")[1];
  46. Integer id = nameMapId.get(name);
  47. if (id == null) {
  48. id = names.size();
  49. names.add(name);
  50. nameMapId.put(name, id);
  51. faceRecognizer.setLabelInfo(id, name);
  52. }
  53. Mat mat = Imgcodecs.imread(file.getCanonicalPath());
  54. Mat gray = new Mat();
  55. // 图片预处理
  56. Imgproc.cvtColor(mat, gray, Imgproc.COLOR_BGR2GRAY);
  57. images.add(gray);
  58. System.out.println("add total " + images.size());
  59. ids.add(id);
  60. }
  61. int[] idsInt = new int[ids.size()];
  62. for (int i = 0; i < idsInt.length; i++) {
  63. idsInt[i] = ids.get(i).intValue();
  64. }
  65. // 显示标签
  66. MatOfInt labels = new MatOfInt(idsInt);
  67. // 调用训练方法
  68. faceRecognizer.train(images, labels);
  69. // 输出持久化模型文件 训练一次后就可以一直调用
  70. faceRecognizer.save(saveFolder + "/face_model.yml");
  71. }
  72. }