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

导入Keras模型至DL4J后预测时出现形状不匹配错误求助

Keras序列模型导入DL4J预测时形状不匹配错误

Keras训练模型代码

model = Sequential()
model.add(Embedding(50000, 128, input_length=10))
model.add(Conv1D(48, 5, activation='relu', padding='valid'))
model.add(GlobalMaxPooling1D())
model.add(Dropout(0.5))
model.add(Flatten())
model.add(Dropout(0.5))

model.add(Dense(7, activation='softmax'))
model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy'])

DL4J预测Scala代码

val modelPath = File("C:/Users/Ayodele/Desktop/Development/classification.h5")
val model: MultiLayerNetwork = 
KerasModelImport.importKerasSequentialModelAndWeights(modelPath.absolutePath)
val list = arrayOf(1,2,3,4,5,6,7,8,9,10)
val inputs = 10
val features = Nd4j.create(1,inputs)
 for (i in 0 until inputs) {
        features.putScalar(intArrayOf(i), list.get(i))
    }
    System.out.println(features)
    val pred = model.predict(features) 

报错信息

Exception in thread "main" org.nd4j.linalg.exception.ND4JIllegalStateException: New 
shape length doesn't match original length: [288] vs [48]. Original shape: [1, 48] 
New Shape: [1, 288]
at org.nd4j.linalg.api.ndarray.BaseNDArray.reshape(BaseNDArray.java:3804)
at org.nd4j.linalg.api.ndarray.BaseNDArray.reshape(BaseNDArray.java:3749)
at org.nd4j.linalg.api.ndarray.BaseNDArray.reshape(BaseNDArray.java:3872)
at org.nd4j.linalg.api.ndarray.BaseNDArray.reshape(BaseNDArray.java:4099)
at 
org.deeplearning4j.preprocessors.KerasFlattenRnnPreprocessor.preProcess(KerasFlattenRnnPreprocessor.java:49)
at org.deeplearning4j.nn.multilayer.MultiLayerNetwork.outputOfLayerDetached(MultiLayerNetwork.java:1299)
at org.deeplearning4j.nn.multilayer.MultiLayerNetwork.output(MultiLayerNetwork.java:2467)
at org.deeplearning4j.nn.multilayer.MultiLayerNetwork.output(MultiLayerNetwork.java:2430)
at org.deeplearning4j.nn.multilayer.MultiLayerNetwork.output(MultiLayerNetwork.java:2421)
at org.deeplearning4j.nn.multilayer.MultiLayerNetwork.output(MultiLayerNetwork.java:2408)
at org.deeplearning4j.nn.multilayer.MultiLayerNetwork.predict(MultiLayerNetwork.java:2270)

解决方案

  • 问题根源:Keras中GlobalMaxPooling1D()输出的是(batch_size, filters)的2D张量(此处为(1,48)),后续添加的Flatten()属于冗余操作,不会改变张量形状,但DL4J的Keras模型导入器错误地将该Flatten()层识别为RNN类展平处理器,导致形状计算出现偏差,引发不匹配错误。

  • 修复步骤:

    1. 修改Keras模型,移除冗余的Flatten()层,修改后的模型代码如下:
      model = Sequential()
      model.add(Embedding(50000, 128, input_length=10))
      model.add(Conv1D(48, 5, activation='relu', padding='valid'))
      model.add(GlobalMaxPooling1D())
      model.add(Dropout(0.5))
      # 移除冗余的Flatten层
      model.add(Dropout(0.5))
      
      model.add(Dense(7, activation='softmax'))
      model.compile(loss='categorical_crossentropy', optimizer='adam', metrics=['accuracy'])
      
    2. 重新训练模型并导出为h5文件,再导入DL4J进行预测。
    3. 确认输入张量:DL4J中创建的features形状为(1,10),与Keras模型的input_length=10匹配,无需调整;注意输入数据应为整数类型(Embedding层要求输入为词索引)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 18:20:58