如何关闭TensorFlow Autograph?tf.data.Dataset.map函数转换问题
搞定tf.data.map的Autograph问题:原因+解决方案
为啥Autograph非要碰你的process_fn?
其实是tf.data的默认行为——map()方法会自动调用Autograph,试图把你传入的Python函数转换成高效的图模式代码,哪怕你写的已经全是TF原生操作。但你的函数里有个Python原生的assert语句,Autograph没法把这种纯Python语法转成图能识别的操作,所以它只能抛出警告,然后退而求其次按原函数执行。
怎么让Autograph别瞎折腾,直接用你的代码?
有两个简单的办法:
办法1:调用map时直接关Autograph
在map里加个autograph=False参数,明确告诉TF不用转换:
dataset = dataset.map(process_fn, autograph=False)
办法2:给函数加个“禁止转换”的装饰器
用TF提供的装饰器标记你的函数,让Autograph直接跳过它:
import tensorflow as tf @tf.autograph.experimental.do_not_convert def process_fn(item): assert 'image' in item # 现在这个断言会立即执行啦 image = tf.image.convert_image_dtype(item.pop('image'), tf.float32) image = tf.multiply(tf.subtract(image, 0.5), 2) return image
怎么让断言立即执行,不是等图跑起来才检查?
你写的assert 'image' in item是纯Python断言,只有在Eager模式下执行才会立刻触发。之前Autograph要转它,就把它塞进图里了,导致要等图运行才会检查,还报错。
只要用上面的任意一种方法关掉Autograph转换,这个Python断言就会在处理每个item的时候立刻执行(这时候item还是Python字典,不是图里的TensorDict),一旦发现缺了'image'键,马上就会报错,不会等到图运行阶段。
要是你的数据源是生成TensorDict的(比如从TFRecord加载的),那Python的in操作没法直接查TensorDict的键,这时候换成TF原生的断言就可以兼容所有情况了:
def process_fn(item): tf.debugging.assert_in('image', item.keys(), message="没找到'image'键!") image = tf.image.convert_image_dtype(item.pop('image'), tf.float32) image = tf.multiply(tf.subtract(image, 0.5), 2) return image
这种方式会在图执行时检查,但不管你用什么数据源都能正常工作。
内容的提问来源于stack exchange,提问作者user3970726
相关产品推荐
相关产品推荐

