如何在Scala中加载.pt人脸识别模型?DJL加载报错解决
如何在Scala(JVM环境)中使用facenet-pytorch的预训练VGG-Face2模型?
问题背景
使用facenet-pytorch仓库的20180402-114759-vggface2.pt预训练模型在Python中实现人脸识别运行正常,但尝试通过DJL(Deep Java Library)在Scala中加载该模型时,出现错误:
ai.djl.engine.EngineException: PytorchStreamReader failed reading zip archive: failed finding central directory
Python中的工作代码
from facenet_pytorch import MTCNN, InceptionResnetV1 from PIL import Image import torch mtcnn = MTCNN(image_size=160, margin=0) resnet = InceptionResnetV1(pretrained='vggface2').eval() resnet.load_state_dict(torch.load('../20180402-114759-vggface2.pt'), strict=False) img1 = Image.open('../img1') img2 = Image.open('../img2') img1_cropped = mtcnn(img1) img2_cropped = mtcnn(img2) if img1_cropped is not None and img2_cropped is not None: img1_embedding = resnet(img1_cropped.unsqueeze(0)) img2_embedding = resnet(img2_cropped.unsqueeze(0)) cos = torch.nn.CosineSimilarity(dim=1, eps=1e-6) similarity = cos(img1_embedding, img2_embedding) print(f"Cosine Similarity: {similarity.item()}") threshold = 0.6 if similarity > threshold: print("The faces are similar!") else: print("The faces are different!") else: print("Face not detected in one or both images.")
错误原因
你尝试加载的20180402-114759-vggface2.pt文件是PyTorch的state_dict(仅包含模型参数),而非完整的TorchScript模型。DJL无法直接加载state_dict,必须加载包含模型结构和参数的完整序列化模型(如TorchScript或ONNX格式)。
解决方案
步骤1:将state_dict转换为TorchScript模型
在Python中运行以下代码,将参数文件转换为DJL可加载的TorchScript模型:
from facenet_pytorch import InceptionResnetV1 import torch # 加载模型结构并导入参数 resnet = InceptionResnetV1(pretrained='vggface2').eval() resnet.load_state_dict(torch.load('../20180402-114759-vggface2.pt'), strict=False) # 创建示例输入(匹配模型输入格式:batch_size=1, channels=3, height=160, width=160) example_input = torch.randn(1, 3, 160, 160) # 追踪模型并导出为TorchScript traced_model = torch.jit.trace(resnet, example_input) # 保存为可被DJL加载的模型文件 traced_model.save("../facenet_vggface2_scripted.pt")
步骤2:修正Scala依赖与模型加载
更新build.sbt,添加人脸检测所需的CV库依赖:
libraryDependencies ++= Seq( "ai.djl" % "api" % "0.29.0", "ai.djl.pytorch" % "pytorch-engine" % "0.29.0" % "runtime", "ai.djl.pytorch" % "pytorch-model-zoo" % "0.29.0", "ai.djl.pytorch" % "pytorch-native-cpu" % "2.3.1" % "runtime" classifier "linux-x86_64", "ai.djl.pytorch" % "pytorch-jni" % "2.3.1-0.29.0" % "runtime", "ai.djl.opencv" % "opencv" % "0.29.0" % "runtime", "ai.djl.modality" % "cv" % "0.29.0" )
步骤3:修正Scala代码(包含人脸检测与输入预处理)
以下是完整的可运行代码,包含人脸检测(对应Python中的MTCNN)、特征提取与相似度计算:
import ai.djl.Model import ai.djl.modality.cv.Image import ai.djl.modality.cv.ImageFactory import ai.djl.modality.cv.translator.FaceDetectionTranslator import ai.djl.ndarray.{NDArray, NDList, NDManager} import ai.djl.ndarray.types.DataType import ai.djl.translate.{Batchifier, Translator, TranslatorContext} import java.nio.file.Paths object FaceRecognitionDJL { def main(args: Array[String]): Unit = { val image1Path = Paths.get("../img_1.png") val image2Path = Paths.get("../img_2.png") val image1 = ImageFactory.getInstance().fromFile(image1Path) val image2 = ImageFactory.getInstance().fromFile(image2Path) // 初始化NDManager val manager = NDManager.newBaseManager() // 检测并裁剪人脸 val face1 = detectFace(manager, image1) val face2 = detectFace(manager, image2) // 加载转换后的TorchScript模型 val model = Model.newInstance("face_recognition_model") model.load(Paths.get("../facenet_vggface2_scripted.pt")) // 提取人脸特征 val embeddings1 = getEmbeddings(model, face1) val embeddings2 = getEmbeddings(model, face2) // 计算相似度 val similarity = compareEmbeddings(embeddings1, embeddings2) println(s"Cosine Similarity: $similarity") val threshold = 0.6 if (similarity > threshold) { println("The faces are similar!") } else { println("The faces are different!") } // 释放资源 model.close() manager.close() } /** * 使用DJL的MTCNN模型检测并裁剪人脸 */ def detectFace(manager: NDManager, image: Image): Image = { val detectorModel = Model.newInstance("mtcnn") // 加载DJL预训练的MTCNN模型 detectorModel.load(Paths.get("https://resources.djl.ai/test-models/pytorch/mtcnn.zip")) val detectorTranslator = FaceDetectionTranslator.builder() .setOutputWidth(160) .setOutputHeight(160) .build() val detector = detectorModel.newPredictor(detectorTranslator) val detectedFaces = detector.predict(image) detectorModel.close() if (detectedFaces.isEmpty) { throw new RuntimeException("Face not detected in image") } detectedFaces.head.getImage } /** * 提取人脸特征向量 */ def getEmbeddings(model: Model, image: Image): Array[Float] = { val predictor = model.newPredictor(new FaceFeatureTranslator) try { predictor.predict(image) } finally { predictor.close() } } /** * 计算两个特征向量的余弦相似度 */ def compareEmbeddings(embedding1: Array[Float], embedding2: Array[Float]): Double = { val dotProduct = embedding1.zip(embedding2).map { case (a, b) => a * b }.sum val norm1 = Math.sqrt(embedding1.map(x => x * x).sum) val norm2 = Math.sqrt(embedding2.map(x => x * x).sum) if (norm1 == 0 || norm2 == 0) 0.0 else dotProduct / (norm1 * norm2) } } /** * 自定义Translator:处理图像输入,适配模型要求 */ class FaceFeatureTranslator extends Translator[Image, Array[Float]] { override def processInput(ctx: TranslatorContext, input: Image): NDList = { val manager = ctx.getNDManager // 将HWC格式(DJL默认)转换为CHW格式(PyTorch要求),并归一化到0-1 val imgArray: NDArray = input.toNDArray(manager, Image.Flag.COLOR) .transpose(2, 0, 1) // HWC -> CHW .toType(DataType.FLOAT32, false) .div(255.0f) .expandDims(0) // 添加batch维度 new NDList(imgArray) } override def processOutput(ctx: TranslatorContext, list: NDList): Array[Float] = { // 移除batch维度,返回特征向量 list.get(0).squeeze(0).toFloatArray } override def getBatchifier: Batchifier = Batchifier.STACK }
关键修正点
- 模型格式转换:将仅含参数的state_dict转为完整的TorchScript模型,解决加载错误
- 输入预处理:将DJL默认的HWC图像格式转为PyTorch要求的CHW格式,并完成归一化
- 人脸检测:添加MTCNN人脸裁剪逻辑,与Python流程保持一致,确保模型输入为标准160x160人脸图像
- 资源管理:添加模型、predictor的关闭逻辑,避免内存泄漏
内容的提问来源于stack exchange,提问作者Zaryab Ali
相关产品推荐
相关产品推荐

