TensorFlow 2.17中添加Masking层后出现OperatorNotAllowedInGraphError错误求助
TensorFlow 2.17中添加Masking层后出现OperatorNotAllowedInGraphError错误求助
Hey Abdul, 我太懂你卡了三天的崩溃感——这种和mask相关的图模式错误真的特别磨人,咱们来快速搞定它!
问题根源
从错误信息能看出来,问题出在TimeDistributed层处理上游传递来的mask张量时:在TensorFlow的图模式下,代码里把符号化的tf.Tensor(这里就是那个形状为(None,13)的bool型mask)当成了Python原生的bool来用,而图模式完全不允许这种操作。
快速解决方案
你其实已经接近答案了!下面两种方法都能解决这个问题,选一个就行:
方法1:给TimeDistributed层显式禁用mask传递
直接在TimeDistributed层的参数里加上mask=None,阻止上游的mask传递到这一层:
model.add(TimeDistributed(Dense(y_train.shape[2], kernel_regularizer=l2(0.001)), mask=None))
解释:你的模型里,Masking层已经在输入阶段过滤了无效的0值,LSTM层也正确利用mask完成了序列计算,后续的Dense层不需要再依赖mask做处理,所以直接关掉mask传递是完全安全的。
方法2:用Lambda层提前清除mask
取消你注释掉的那行Lambda层代码就行,它会把上游的mask清除,让TimeDistributed层收不到mask:
model.add(Lambda(lambda x: x, mask=None))
这和方法1的本质是一样的,都是切断mask的传递链,避免触发图模式下的不兼容操作。
测试建议
修改完代码后重新运行训练,应该就能绕过这个错误了。如果还有问题,可以检查下输入数据的mask值是否正确(比如确实是用0.0作为无效值),或者尝试把整个build_lstm函数用@tf.function装饰一下,不过前两种方法基本就能解决问题啦。
备注:内容来源于stack exchange,提问作者Abdul Basit
相关产品推荐
相关产品推荐

