设置num_parallel_calls后TF Map函数并行无性能提升问题排查
问题分析与修正方案
你遇到的并行不生效问题,主要是代码里的几个错误加上tf.py_func的特性导致的,咱们一步步拆解:
1. 数据集处理的低级错误
看你的代码里这段:
if batch: dataset_y = dataset_x.batch(batch_size)
这里明显写错了!你把dataset_y的batch处理写成了用dataset_x,导致dataset_y其实和dataset_x是同一个对象,后续两个迭代器都是从同一个数据集拿数据,自然不存在并行处理两个数据集的情况,调整成:
if batch: dataset_y = dataset_y.batch(batch_size)
2. 迭代器复用错误
同样,你创建Y的迭代器时,用的是dataset_x.make_one_shot_iterator(),这意味着X和Y都是从同一个dataset_x取数据,完全没有用到dataset_y,正确的写法应该是:
Y = dataset_y.make_one_shot_iterator().get_next()
3. tf.py_func的GIL限制
这是很关键的一点:你用tf.py_func调用的是Python自定义函数squarer,而Python的全局解释器锁(GIL)会限制同一时间只有一个Python线程执行代码。也就是说,哪怕你设置了num_parallel_calls,这些并行调用的Python函数其实没法真正同时运行,只能串行执行,这也是你看不到性能提升的核心原因之一。
如果想要真正利用并行,建议:
- 把
squarer这类简单操作改成纯TensorFlow的算子(比如tf.square(x)),避开Python函数的GIL限制; - 如果必须用Python函数,可以考虑用
tf.data.experimental.map_with_optional_args结合多进程,或者升级到TensorFlow 2.x的tf.py_function配合tf.data.AUTOTUNE,同时注意处理GIL的问题。
4. 测试细节的优化
- 用
%timeit测试时,最好先手动运行一次函数做热身,避免把初始化图的时间算进去; - 你的
num_iterations=1000,但原始数据集是range(1000),repeat=1,所以迭代1000次刚好取完数据,可能没法完全体现并行的优势,可以加大repeat的次数,比如repeat=10,让迭代次数足够多。
修正后的测试代码示例
import tensorflow as tf import time def squarer(x): # 模拟耗时操作,方便看出并行效果 time.sleep(0.001) return x * x def test_two_custom_function_parallelism(num_parallel_calls=1, batch=False, batch_size=1, repeat=1, num_iterations=10): tf.reset_default_graph() start = time.time() # 修正dataset_y的batch处理 dataset_x = tf.data.Dataset.range(1000).map(lambda x: tf.py_func(squarer, [x], [tf.int64]), num_parallel_calls=num_parallel_calls).repeat(repeat) if batch: dataset_x = dataset_x.batch(batch_size) dataset_y = tf.data.Dataset.range(1000).map(lambda x: tf.py_func(squarer, [x], [tf.int64]), num_parallel_calls=num_parallel_calls).repeat(repeat) if batch: dataset_y = dataset_y.batch(batch_size) # 这里修正 # 修正Y的迭代器来源 X = dataset_x.make_one_shot_iterator().get_next() Y = dataset_y.make_one_shot_iterator().get_next() with tf.Session() as sess: sess.run(tf.global_variables_initializer()) i = 0 while True: try: res = sess.run([X, Y]) i += 1 if i == num_iterations: break except tf.errors.OutOfRangeError as e: break # 测试时先热身 test_two_custom_function_parallelism(num_iterations=10, num_parallel_calls=2, batch_size=2, batch=True) # 再正式测试 %timeit test_two_custom_function_parallelism(num_iterations=1000, num_parallel_calls=2, batch_size=2, batch=True) %timeit test_two_custom_function_parallelism(num_iterations=1000, num_parallel_calls=10, batch_size=2, batch=True)
额外提示
如果换成纯TensorFlow算子(比如把tf.py_func(squarer, [x], [tf.int64])改成tf.square(tf.cast(x, tf.int64))),你会发现num_parallel_calls的提升效果非常明显,因为TensorFlow的原生算子是避开GIL的,可以真正利用多核并行。
内容的提问来源于stack exchange,提问作者Kracekumar
相关产品推荐
相关产品推荐

