TFJS搭建深度LSTM模型训练时stack参数张量类型报错如何解决
问题排查方案
- 第一步:先注释所有
batchNormalization层调用,重新运行测试代码
tfjs的BatchNormalization层处理LSTM输出的三维时序张量时,内部维度计算逻辑和Python TensorFlow存在差异,会触发张量堆叠错误,这是最可能导致你当前报错的原因。你可以先把所有BN层删除跑通基础模型,再调整层的顺序(比如将BN移到Dropout之前,或者仅在全连接层部分添加BN)。 - 第二步:替换LSTM的激活函数为默认值
tanh测试
你当前所有LSTM层都指定了relu作为激活,而tfjs的LSTM内部默认循环激活为hardSigmoid,两种激活配合在小时间步(你当前输入的时间步为1)场景下会出现内部张量维度丢失的问题。你可以先去掉LSTM层的activation参数,用默认的tanh跑通后再根据需求调整激活函数。 - 第三步:修正测试用的标签张量
你当前用tf.ones([1,2])作为分类任务的标签不符合categoricalCrossentropy的输入要求,正确的one-hot标签可参考如下写法:
const label = tf.oneHot(tf.tensor1d([0], 'int32'), 2);
虽然这个问题不是你当前报错的直接原因,但会导致后续损失计算和模型训练异常。
内容的提问来源于stack exchange,提问作者jeefo
相关产品推荐
相关产品推荐

