You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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
}

关键修正点

  1. 模型格式转换:将仅含参数的state_dict转为完整的TorchScript模型,解决加载错误
  2. 输入预处理:将DJL默认的HWC图像格式转为PyTorch要求的CHW格式,并完成归一化
  3. 人脸检测:添加MTCNN人脸裁剪逻辑,与Python流程保持一致,确保模型输入为标准160x160人脸图像
  4. 资源管理:添加模型、predictor的关闭逻辑,避免内存泄漏

内容的提问来源于stack exchange,提问作者Zaryab Ali

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.20 01:22:32