【问题标题】:Keras Multiprocessing breaks validation accuracyKeras 多处理破坏了验证的准确性
【发布时间】: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


【解决方案1】:

也许您应该将相同的数据平衡功能应用于您的验证数据生成器?

【讨论】:

  • 我试过了,但没有帮助。尽管如此,还是感谢您的回答。
猜你喜欢
  • 2018-10-20
  • 2020-03-23
  • 2020-06-15
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2011-08-22
  • 2020-02-12
  • 2020-08-25
相关资源
最近更新 更多