【发布时间】:2020-08-21 18:12:02
【问题描述】:
我正在使用大型数据集训练神经网络,因此我需要使用多个工作人员/多处理来加快训练速度。
之前我使用的是 keras 生成器,并使用 fit 生成器,其中 multiprocessing 设置为 false,worker 设置为 16,但是最近我不得不使用自己的生成器,所以我创建了自己的 flow_from_directory 生成器,如下所示:
train_generator = train_datagen.flow_from_directory(
train_data_dir,
target_size=(image_size, image_size),
batch_size=training_batch_size,
class_mode='categorical') # set as training data
bal_gen = balanced_flow_from_directory(train_generator)
def balanced_flow_from_directory(flow_from_directory):
for x, y in flow_from_directory:
yield custom_balance(x, y)
但是,在 fit 生成器中,当我将 worker > 1 和 MultiProcessing 设置为 False 时,它会告诉我我的生成器不是线程安全的,因此不能与 worker > 1 和 Multiprocessing 设置为 False 一起使用。当我将工人保持 > 1 并将 MultiProcessing 设置为 True 时,代码会运行,但它会给我如下警告:
警告:tensorflow:使用带有
use_multiprocessing=True的生成器和多个工作人员可能会复制您的数据。请考虑使用tf.data.Dataset
此外,验证会给出非常奇怪的输出,例如:
1661/1661 [===============================] - ETA:0s - 损失:0.1420 - 准确度:0.9662警告:tensorflow:使用带有
use_multiprocessing=True的生成器和多个工作人员可能会复制您的数据。请考虑使用tf.data.Dataset。 1661/1661 [==============================] - 475s 286ms/步 - 损失:0.1420 - 准确度:0.9662 - val_loss : 6.2723 - val_accuracy: 0.0108elines tf.data 推荐。
验证准确率总是很低,val_loss 总是很高。我可以做些什么来解决这个问题吗?
更新:我找到了使生成器函数线程安全的代码,如下所示:
import threading
class threadsafe_iter:
"""
Takes an iterator/generator and makes it thread-safe by
serializing call to the `next` method of given iterator/generator.
"""
def __init__(self, it):
self.it = it
self.lock = threading.Lock()
def __iter__(self):
return self
def __next__(self):
with self.lock:
return self.it.__next__()
def threadsafe_generator(f):
def g(*a, **kw):
return threadsafe_iter(f(*a, **kw))
return g
@threadsafe_generator
def balanced_flow_from_directory(flow_from_directory):
for x, y in flow_from_directory:
yield custom_balance(x, y)
现在我可以使用 workers = 16 并将 Multiprocessing 设置为 False,就像我在制作自定义生成器之前使用的那样。但是,当我这样做时,每个 epoch 需要 30 分钟,而以前需要 7 分钟。
当我使用 workers=16 并将 multiprocessing 设置为 true 时,它给我的问题与我在上面将 multiprocessing 设置为 true 时遇到的问题相同 - 即验证准确性破坏。
【问题讨论】:
-
删除了
imbalanced data标签,因为这与处理类/目标大小的不平衡有关
标签: python multithreading tensorflow keras multiprocessing