Генераторы

Генераторы

Категория статьи: Нейронные сети

Проект: Shazam своими руками

Обвертка для генераторов. Генератор порции данных.

Доброго времени суток!

Для загрузки файлов с диска и подачу данных для обучения модели буду использовать генераторы.

Как я уже говорил раньше, нарезанные файлы mp3, т.е. наша обучающая выборка, достаточно тяжелая и мы не можем её загрузить в память сразу всю, для этого будем использовать генератор порции (batch generator). Для того, чтобы мы могли бесконечно формировать порции данных хорошо бы зациклить этот процесс. К счастью в пакете itertools есть итератор cycle который позволяет это делать. Однако, при обучении модели, я заметил, что после определенного момента свободной оперативной памяти почти не осталось, а еще, спустя несколько эпох, ноутбук свалился по памяти. Что же происходит? А происходит следующее: итератор cycle создает список и все элементы, которые он перебирает, заносит в него. После того как закончились данные, начинает их брать из списка, который создан на памяти. Это равносильно тому, что загрузить сразу все данные в память и оттуда подавать их на обучение. Но сделать этого мы не можем, т.к. данные весят очень много!

* Итератор cycle можно применять для последовательностей, которые не занимают много памяти

Принято было решение написать обвёртку для генератора, которая отлавливает исключение StopIteration (это исключение возникает когда генератор доходит до конца итерируемой последовательности) и пересоздаёт генератор. Таким образом в памяти будет только наш batch.

class GenWrapper():
  def __init__(self, gen_fn, *args, **kwargs):
    self.gen_fn = gen_fn
    self.args = args
    self.kwargs = kwargs
    self.gen = self.__create_gen()

  def __create_gen(self):
    return self.gen_fn(*self.args, **self.kwargs)

  def __iter__(self):
    return self

  def __next__(self):
    try:
      return next(self.gen)
    except StopIteration:
      self.gen = self.__create_gen()
      return next(self.gen)

Дальше создаём сам генератор batch'a, в котором есть список генераторов для каждого класса композиции.

import random
import itertools
import numpy as np
from tensorflow.keras import utils

def gen_mel(src_dir):
  files = os.listdir(src_dir)

  # Фильтруем файлы от системных
  files = list(filter(lambda x: x[0] != '.', files))

  random.shuffle(files)
  for f in files:
    d_ = unpickle(os.path.join(src_dir, f))
    yield (d_['X'], d_['y'])

def gen_batch(src_dir, batch_size):
  files = os.listdir(src_dir)
  files = sorted(files)

  # Фильтруем файлы от системных
  files = list(filter(lambda x: x[0] != '.', files))

  cl_count = int(files[-1]) + 1

  random.shuffle(files)

  l_gen_mel = []
  for f in files:
    l_gen_mel.append(GenWrapper(gen_mel, os.path.join(src_dir, f)))
  gen_mel_cycle = itertools.cycle(l_gen_mel)

  x_shape = (batch_size, lhp.n_mels, lhp.n_timeframe)
  while True:
    X = np.zeros(x_shape)
    Y = np.zeros((batch_size, ), dtype=int)
    for i in range(batch_size):
      X[i], Y[i] = next(next(gen_mel_cycle))

    X = X.reshape(x_shape + (1,))
    Y = utils.to_categorical(Y, cl_count)
    yield (X, Y)

Генератор gen_mel при создании получает директорию в которой лежат файлы-нарезки одной музыкальной композиции (соответствует директории), перемешивает список файлов и при каждом обращении методом next к генератору будет отдавать (X и y) входные и выходные данные для обучения модели. Когда дойдет до последнего файла в последовательности, генератор перезапускается, т.к. используется обвертка GenWrapper и снова перемешивает список файлов.

Генератор gen_batch при создании получает директорию train или val и размер батча. Далее читает список директорий музыкальных композиций, перемешивает их и для каждой директории создает генератор gen_mel. При обращении к генератору методом next в цикле формирует batch, делает преобразования размерностей и возвращает (X, Y) - данные, которые использует модель для обучения.