Tensorflow JS堆叠有状态LSTM在线训练时出现stack参数报错是什么原因?
报错含义
这个报错表示TensorFlow JS中的tf.stack()算子接收到的输入参数不符合要求,该算子要求输入必须是张量数组或者类张量数组,当前触发错误时传入的参数不是符合要求的数组结构。
问题原因
- 堆叠stateful LSTM时,所有LSTM层都必须开启
stateful: true配置。你只给第一层LSTM设置了stateful模式,第二层LSTM默认是stateless模式,两层的状态管理逻辑不兼容,在处理单时间步、batch size=1的输入时,第二层LSTM内部计算需要拼接的张量生成失败,导致传入stack算子的参数非法。 - stateful模式下的LSTM层要求输入的batch维度、时间步维度固定,你没有给第二层LSTM显式指定输入shape,TensorFlow JS对上层传入的shape为
[1,1,256]的输出做自动shape推导时出现逻辑异常,也会触发该报错。
修复方法
给第二层LSTM补充stateful: true配置即可,修复后的模型代码如下:
this.model = tf.sequential(); this.model.add(tf.layers.lstm({ units: 256, returnSequences: true, batchInputShape: [1, 1, input_dim], stateful: true })); this.model.add(tf.layers.lstm({ units: 128, returnSequences: false, stateful: true })); this.model.add(tf.layers.dense({ units: output_dim})); this.model.compile({ optimizer: 'adam', loss: 'meanSquaredError' });
内容的提问来源于stack exchange,提问作者Zain Syed
相关产品推荐
相关产品推荐

