如何在Java中使用Python训练的TensorFlow Keras图像识别模型
问题解决方案
你遇到的两个报错均来自Deeplearning4j对新版Keras特性的兼容不足,可按以下两种路径解决:
方案1:调整Python端代码适配Deeplearning4j导入规则
- 移除模型内的
Rescaling预处理层,将归一化逻辑外置:训练阶段手动对数据集做像素值/255的归一化处理,推理阶段在Java端对输入的BufferedImage做同样的归一化操作即可。 - 调整损失函数声明方式:模型编译时用字符串格式声明损失,同时建议在最后一层增加Softmax激活适配常规推理逻辑,修改后的核心代码如下:
model = Sequential([ # 移除Rescaling层,input_shape直接写(30,30,3) layers.Conv2D(10, 3, padding='same', activation='relu', input_shape=(30, 30, 3)), layers.MaxPooling2D(), layers.Conv2D(15, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Conv2D(20, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Dropout(0.2), layers.Flatten(), layers.Dense(64, activation='relu'), layers.Dense(num_classes, activation='softmax') ]) model.compile( optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'] )
- 导出模型时优先选择HDF5格式,确保没有自定义层和自定义损失。
方案2:更换Java端调用方案,兼容性更强
推荐直接使用TensorFlow官方Java API加载模型,无需修改Python训练代码,所有Keras层和损失函数都能完美兼容,调用示例如下:
import org.tensorflow.SavedModelBundle; import org.tensorflow.Tensor; import org.tensorflow.ndarray.Shape; import org.tensorflow.ndarray.FloatNdArray; import org.tensorflow.ndarray.NdArrays; import java.awt.image.BufferedImage; import javax.imageio.ImageIO; import java.io.File; public class TfModelTest { public static void main(String[] args) throws Exception { // Python端先导出为SavedModel格式:model.save("saved_model") try (SavedModelBundle model = SavedModelBundle.load("saved_model", "serve")) { BufferedImage img = ImageIO.read(new File("testImage.png")); // 构造输入张量,尺寸为[1, 30, 30, 3],同时做归一化处理 FloatNdArray input = NdArrays.ofFloats(Shape.of(1, 30, 30, 3)); for (int h = 0; h < 30; h++) { for (int w = 0; w < 30; w++) { int rgb = img.getRGB(w, h); input.setFloat((rgb >> 16 & 0xff)/255f, 0, h, w, 0); input.setFloat((rgb >> 8 & 0xff)/255f, 0, h, w, 1); input.setFloat((rgb & 0xff)/255f, 0, h, w, 2); } } // 执行推理 try (Tensor<Float> output = model.session().runner() .feed("serving_default_input_1", Tensor.of(input)) .fetch("StatefulPartitionedCall") .run() .get(0) .expect(Float.class)) { FloatNdArray result = output.asNdArray(); // 读取分类结果 System.out.println("预测结果: " + result); } } } }
两种方案相比更推荐方案2,不需要调整已训练好的模型结构和训练逻辑,稳定性和兼容性更高。
内容的提问来源于stack exchange,提问作者Andreas Radauer
相关产品推荐
相关产品推荐

