将Keras模型导入TensorFlow Java后预测阶段报错求助
问题解决:
In[0] is not a matrix 错误 错误含义
这个错误直白来说就是:模型的第一个输入不是二维矩阵,而是一维张量。你模型里的dense_1全连接层要做矩阵乘法运算,而矩阵乘法要求输入至少是二维结构(比如(批量数, 特征数)的格式),一维张量没法满足这个运算要求,所以触发报错。
和输入维度的关系
完全相关。从saved_model_cli的输出能看到,模型要求的输入形状是(-1,6):
-1代表任意批量大小(比如1个样本、10个样本都可以)6是每个样本的特征数量
所以输入必须是二维张量,但你用TFloat32.vectorOf(x)生成的是一维张量(形状为(6)),和模型要求的输入结构不匹配,这就是问题根源。
修正后的Java代码
需要把一维数组转换成形状为(1,6)的二维张量,同时修正输出值的错误获取方式:
public static void importKerasModel() { try (SavedModelBundle model = SavedModelBundle.load("PATH\\kerasModel", "serve")) { float[] x = {0.48f, 0.48f, 0.48f, 0.48f, 0.48f, 0.48f}; // 生成形状为(1,6)的二维张量,对应单样本批量输入 try (Tensor input = TFloat32.tensorOf(Shape.of(1, 6), x); Tensor output = model.session() .runner() .feed("serve_keras_tensor:0", input) // 使用完整张量名称避免匹配问题 .fetch("StatefulPartitionedCall:0") .run() .get(0)) { // 输出形状是(1,1)的二维张量,用二维数组接收后提取结果 float[][] outputArray = new float[1][1]; output.copyTo(outputArray); float prediction = outputArray[0][0]; System.out.println("prediction = " + prediction); } } catch (Exception e) { e.printStackTrace(); } }
额外提醒
- 喂入张量时尽量使用完整名称(比如
serve_keras_tensor:0),避免潜在的名称匹配歧义 - 模型输出形状是
(-1,1),属于二维张量,必须用对应维度的数组来接收结果
内容的提问来源于stack exchange,提问作者Cyrano
相关产品推荐
相关产品推荐

