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

Python训练的TensorFlow ECG CNN-LSTM模型在Java TensorFlow推理时出现变量不存在错误

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 10:04:36