Python训练的TensorFlow ECG CNN-LSTM模型在Java TensorFlow推理时出现变量不存在错误
从你提供的错误信息来看,Java端加载SavedModel后执行推理时找不到模型变量sequential/dense_1/kernel,这通常和模型导出的正确性、Java端的调用逻辑,或者TensorFlow版本兼容性有关。下面是针对性的解决方案:
1. 确保导出的是训练完成且权重正确的模型
你的训练代码中使用了EarlyStopping(restore_best_weights=True),理论上训练后的model已经恢复了最优权重,但为了避免潜在的权重未同步问题,建议显式加载ModelCheckpoint保存的最优模型再导出为SavedModel格式:
修改Python代码中的导出部分:
# ====================================== # 6. Salvar o Modelo para TensorFlow Java # ====================================== # Certifique-se de que o modelo foi treinado antes de exportar export_dir = "ecg_model_to_javav2" # 加载ModelCheckpoint保存的最优模型 best_model = tf.keras.models.load_model('best_model_cnn_lstm.keras') # 导出为SavedModel(兼容Java TensorFlow的格式) tf.saved_model.save(best_model, export_dir)
2. 修正Java端的推理调用逻辑(传递命名输入张量)
Keras模型导出为SavedModel后,serving_default签名的输入是命名张量的字典(对应Keras模型的输入层名称),而不是单个张量。你当前直接传递单个inputTensor会导致模型无法正确映射输入,进而引发变量查找错误。
步骤2.1:确认输入名称
先通过saved_model_cli命令查看模型的输入签名详情:
saved_model_cli show --dir ./ecg_model_to_javav2 --tag_set serve --signature_def serving_default
输出会类似这样(注意输入名称,比如conv1d_input):
The given SavedModel SignatureDef contains the following input(s): inputs['conv1d_input'] tensor_info: dtype: DT_FLOAT shape: (-1, 100, 1) name: serving_default_conv1d_input:0
步骤2.2:修改Java代码的推理部分
将单个张量传递改为传递命名输入字典:
// 2. Criar tensor de entrada corretamente try (TFloat32 inputTensor = TFloat32.tensorOf( Shape.of(1, 100, 1), data -> { for (int i = 0; i < 100; i++) { data.setFloat(ecgSample[0][i][0], 0, i, 0); } } )) { // 3. 准备命名输入字典(键为saved_model_cli显示的输入名称) Map<String, Tensor> inputs = new HashMap<>(); inputs.put("conv1d_input", inputTensor); // 替换为你的实际输入名称 // 4. Executar inferência usando o dicionário de entrada try (Tensor outputTensor = model.function("serving_default").call(inputs)) { if (outputTensor instanceof TFloat32) { processOutput((TFloat32) outputTensor); } else { System.err.println("Erro: O modelo não retornou um TFloat32."); } } }
3. 验证TensorFlow版本兼容性
确保Python端使用的TensorFlow版本与Java端的TensorFlow库版本完全一致(主版本和次版本都要匹配,比如Python用2.15.0,Java也要用2.15.0)。版本不兼容会导致模型结构解析错误,进而出现变量找不到的问题。
4. 检查Java端张量填充的正确性
可以在填充张量后添加日志,验证输入数据是否符合模型预期:
// 在填充张量的循环后添加日志 System.out.println("Amostra de entrada do ECG:"); for (int i = 0; i < 10; i++) { System.out.printf("Posição %d: %.4f%n", i, ecgSample[0][i][0]); }
确认数据范围和形状是否与Python训练时的输入一致。
备注:内容来源于stack exchange,提问作者Tatiana Filgueiras

