在向TensorFlow的 tf.data.Dataset.map() 添加一个Python函数后,速度变慢
我正在使用tf.data.Dataset来进行数据预处理。当我在map()中使用TensorFlow操作时,流水线运行得很快。但是,在map()中添加一个简单的自定义Python函数后,训练速度变慢了。
import tensorflow as tf
def custom_fn(x):
return x * 2
dataset = tf.data.Dataset.range(100000)
dataset = dataset.map(lambda x: custom_fn(x))
为什么添加一个简单的Python函数会降低tf.data流水线的性能?保持预处理更高效的推荐做法是什么?
解决方案
在Python中,张量运算的行为与TensorFlow的运算相同。如果不使用tf.py_function,就不会有上下文切换的额外开销。为了准确定位变慢的真正原因,能否分享你自定义函数和流水线的完整代码?在此期间,为了解决常见瓶颈,请确保对数据进行批处理,并使用AUTOTUNE对操作进行并行映射。
import tensorflow as tf
def custom_fn(x):
return x * 2
dataset = tf.data.Dataset.range(100000)
# Batch first for vectorization, then map in parallel using AUTOTUNE
dataset = dataset.batch(32)
dataset = dataset.map(custom_fn, num_parallel_calls=tf.data.AUTOTUNE)
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。