使用@tf.function装饰器赋值类属性时触发报错如何解决?
报错原因
@tf.function 装饰器的作用是将函数内的张量运算转换为静态计算图实现执行加速,你遇到的报错是因为被装饰的方法包含两类无法被静态图支持的操作:
- 类实例属性赋值操作:
self.input_ae = inputs、self.autoencoder = Model(...)这类属于Python层面的副作用操作,tf.function在构建静态图的跟踪阶段只会执行一次Python逻辑,后续调用不会再执行这些Python代码,且静态图本身不支持维护Python对象的属性状态。 - Keras模型结构定义操作:
Model(input, output)是Python层面的模型结构定义逻辑,不属于可转换的张量运算。你传入的Input生成的是仅用于搭结构的Keras符号占位张量,没有实际数值,tf.function尝试将其转换为静态图支持的张量类型时就会触发类型错误。
你提供的最小测试用例触发报错的原因完全一致:在@tf.function装饰的方法内执行类属性赋值,且传入了Keras符号张量。
解决方案
直接移除build_ae和compile_ae方法上的@tf.function装饰器即可。这两类方法都是模型结构定义、编译的配置逻辑,本身不需要计算图加速,只有训练步、前向推理、损失计算这类张量密集型的运算才需要用@tf.function装饰优化。
内容的提问来源于stack exchange,提问作者Daniel Kelshaw
相关产品推荐
相关产品推荐

