Всем привет, меня зовут Антон, я работаю в Сбере разработчиком на Java, в продукте GigaIDE. В этой статье я расскажу, как оптимизировал расчёты в простой нейронной сети. Сначала попробую перемножать матрицы в многопоточном режиме, потом перейду к многопоточному обучению сети. Буду использовать сторонние библиотеки EJML и ND4J, последняя к тому же позволит обучить нейронную сеть на видеокарте (GPU).

В прошлой статье я написал простую нейронную сеть, которая распознаёт рукописные цифры из наборе MNIST. На тестовом множестве сеть распознаёт цифры с точностью в 97%, а цифры, написанные моей рукой, — с точностью 50%. Что-то тут не так, нужны эксперименты, а для экспериментов нужна высокая скорость расчётов. Утомительно ждать несколько минут, чтобы подтвердить или опровергнуть небольшую гипотезу, например, просто увеличив коэффициент обучения. Хочется, чтобы всё было побыстрее, поэтому я взялся за оптимизацию расчётов. 

Самый длительный процесс в нейронной сети — обучение. Я обучал на множестве из 60 000 картинок. На одну эпоху изначально уходило около 50 секунд, замерял с помощью @RepeatTest(5) из JUnit. JMH для таких оценок не нужен:

В консоль вывел точность сети, которая после одной эпохи составила примерно 96%.

Напомню, что моя сеть — это переписанная строчка в строчку питновоская сеть из книги Тарика Рашида. Я могу замерить производительность и на Python’е. Правда, совсем не разбираюсь в этом языке, поэтому просто записал время в начале скрипта, а в конце вычел текущее и получил разницу. Стандартная методика, может быть, для Python’а надо сделать как-то по-другому, я не знаю. В любом случае, в репу положил скрипт, желающие могут проверить. Итак, в итоге:

На Pyhon’e получалось примерно 75 секунд, то есть в 1,5 раза медленнее, чем на Java. Хотя там используется numpy, который обещает хорошую производительность. Почему numpy не помогает, предположу и расскажу ниже.

Профилируем

В вопросах оптимизации нужны точные инструменты и числа, поэтому воспользуюсь JProfiler, и поищу самые «горячие» методы. Общая картина при обучении сети:

«Горячие» методы:

Конечно, не сложно было предположить, что самым «горячим» методом будет перемножения матриц, но всё равно необходимо получить числа и общую картину происходящего. 

Умножаем в многопотоке

Как ускорить умножение матриц? Взглянув на код метода перемножения сразу понимаешь, что оптимизировать там нечего: всё настолько просто, насколько возможно. Тогда пробуем многопоток. Каждую строку матрицы с результатом перемножения можно получать независимо от других строк, следовательно, можно легко распараллелить этот процесс. В пределе можно распараллелить независимый расчёт каждого элемента матрицы, но это уж слишком мелко. 

Чтобы рассчитать каждый ряд матрицы в многопотоке, удобно воспользоваться методом Arrays.parallelSetAll(). Он вторым параметром принимает лямбду, которая, в свою очередь, считается в отдельном потоке:

public static double[][] multiplyInParallel(double[][] ma, double[][] mb) {
   double[][] result = new double[ma.length][mb[0].length];
   Arrays.parallelSetAll(result, arrayRowIndex -> multiplyRow(ma, mb, arrayRowIndex));
   return result;
}


private static double[] multiplyRow(double[][] ma, double[][] mb, int row) {
   double[] resultRow = new double[mb[0].length];


   for (int j = 0; j < mb[0].length; j++) {
       for (int k = 0; k < ma[0].length; k++) {
           resultRow[j] += ma[row][k] * mb[k][j];
       }
   }
   return resultRow;
}

Запускаем и смотрим на результат:

Стало только хуже. Примерно на 15 секунд медленнее, то есть хуже примерно на 30%. Для интереса можно также спрофилировать и посмотреть на потоки:

Во всю трудится ForkJoinPool, но большую часть времени потоки просто ничего не делают, а в «горячие» методы попадает Thread.run()

Наверное, тут слишком большая гранулярность. На каждую строку массива свой поток. Можно попробовать гранулярность повысить, например, считать в отдельном потоке не по одной, а по 100 строк матрицы. Для этого написал уже свою реализацию RecursiveAction, которая исполняется на FJP (привожу только метод compute): 

@Override
protected void compute() {
   int rowCount = endRow — startRow;
   if (rowCount <= THRESHOLD) {
       for (int i = startRow; i < endRow; i++) {
           result[i] = multiplyRow(ma, mb, i);
       }
   } else {
       int mid = startRow + rowCount / 2;
       MatrixMultiplyTask t1 = new MatrixMultiplyTask(ma, mb, result, startRow, mid);
       MatrixMultiplyTask t2 = new MatrixMultiplyTask(ma, mb, result, mid, endRow);
       invokeAll(t1, t2);
   }
}

Запускаем и проверяем:

По скорости — нет эффекта, всё равно на 25 секунд медленнее (примерно на 50%), чем простое последовательное перемножение матриц. 

Делаю выводы, что перемножать в многопотоке, по крайней мере небольшие матрицы, — плохая идея. Это важно. На протяжении всей статьи я буду делать акцент на то, что матрицы у меня небольшие. В моём случае это 728 на 200 элементов. С виду вроде и не маленькая, но по меркам текущих современных нейронных сетей —  просто «наноматрицы». То, что распараллеливать умножение матриц — плохая идея, было видно ещё после самого первого профилирования. Метод multiply исполняется 88 микросекунд, такие методы не параллелятся. Но всё-таки попробовать надо было, к тому же, подход можно использовать для больших матриц, где могут быть совсем другие эффекты. 

Можно было ещё попробовать Vector API, но немного изучив вопрос, я выяснил, что в случае небольших матриц он не поможет. К тому же современные JVM в определённых ситуациях сами могут векторизовать инструкции. 

Асинхронный процесс обучения

Если нельзя ускорить самый «горячий» метод, то можно попробовать ускорить метод, который лежит выше по стеку. Из профиля видно, что это, конечно, метод train(). Обучать сеть можно независимо и параллельно каждой картинкой. Разбиваем 60 000 картинок на батчи, допустим, по 1 000. Далее, для каждого такого набора запускаем обучение в отдельном потоке.

Здесь батчи — это не мини батчи из мира нейронных сетей, а пока что лишь просто наборы данных меньшего размера. Корректировка весов сети выполняется после каждой картинки, поэтому нужна синхронизация на запись и чтение весов. Для этого в сеть добавляю ReentrantLock и читаю веса под замком:

try {
   weightsLock.lock();
   currentInputToHiddenWeights = inputToHiddenWeights;
   currentHiddenToOutputsWeights = hiddenToOutputsWeights;
} finally {
   weightsLock.unlock();
}

Записываю веса под замком:

private void adjustWeights(double[][] deltaInputsToHidden,
                          double[][] deltaHiddenToOutputs) {
   try {
       weightsLock.lock();
       inputToHiddenWeights = MatrixUtils.add(inputToHiddenWeights, deltaInputsToHidden);
       hiddenToOutputsWeights = MatrixUtils.add(hiddenToOutputsWeights, deltaHiddenToOutputs);
   } finally {
       weightsLock.unlock();
   }
}

Запись весов приводит к пересозданию матрицы и перезаписыванию ссылки на неё, таким образом веса обновляются парой, а при чтении весов получаю их актуальную версию. 

Позаботившись о синхронизации, запускаю и проверяю:

Ура! Стало в 2,5 раза быстрее. Отлично, я сдвинулся с мёртвой точки. Причём возросла и точность, не сильно, но всё же. 

Профилируем и смотрим на потоки:

Здесь потоки уже не просто ждут, а иногда блокируются на замках. Что можно с этим сделать? Можно феерически забить на синхронизацию! Неужели? Убираю замки и запускаю:

Получаю прирост по производительности около 15% и немного худшую точность, но, в целом, это работает!

Профилируем:

Тут, ожидаемо, наступает «потоковый рай».

Так нужны ли замки или нет? По-хорошему, конечно, нужны, но можно первую эпоху прогнать без синхронизации, проверить какие-то гипотезы, а далее, если это необходимо, дообучить сеть более строго на замках. Возможен вариант записи под замками, а чтение без них. В этом случае я получил первоначальную высокую точность и прирост скорости примерно в 10%. Тут нужны отдельные эксперименты.

Замечу, что при простом последовательном (в один поток) обучении пиковое потребление памяти примерно 750 МБ, а при асинхронном — 2,5 Гб, то есть в три раза больше. На небольших сетях это не критично, на огромных — может быть непреодолимым барьером.

На этом этапе я прекращаю оптимизировать свои методы и попробую использовать готовые библиотеки. Ведь всё уже написано до нас, надо лишь найти.  

ND4J

ND4J — библиотека для работы с матрицами. Является частью более крупного проекта Deeplearning4j. ND4J предоставляет оптимизацию расчётов на оборудовании почти для всех платформ, причём позволяет считать матрицы и на GPU. Почти все вычисления уходят в натив на С++, а на Java только API. Замечу важную особенность: ND4J не хранит свои объекты в Java Heap, а использует память напрямую от ОС (off-heap memory). 

Порадовала очень подробная документация. Видно, что библиотека заточена под расчёты в нейронной сети. Удобные методы инициализации весов, встроенные функции активации, функции ошибок и прочее. Использовать удобно. 

В общем, немного переписываю исходную нейронную сеть, чтобы можно было использовать в ней сущности библиотеки (INDArray), а не массивы double. Заменяю математику и запускаю в один поток:

Получаю кучу предупреждений от JVM и почти в 1,5 раза худшую производительность. Напомню, моя нейронная сеть на математике «в лоб» обучается за 50 секунд. 

Профилирую. Здесь интересна картина по памяти:

Видно, что память куда-то утекает потихоньку. Основная сущность (INDArray) билилотеки ND4J реализует интерфейс AutoCloseable, поэтому объекты надо вовремя закрывать. И вроде бы все закрываю… Замечу, что здесь на картинке мы видим память Java Heap. Если глянуть в ОС, то потребление памяти растёт с течение времени и в конце достигает примерно 10 ГБ (а это очень много), то есть в ОС память тоже куда то утекает. Надо разбираться.

Нагрузка на процессор примерно 55%.

Если глянуть на «горячие» методы, то там будет всё знакомо:

На первом месте метод перемножения матриц. Ожидаемо. Только тут он как раз исполняется почти в 1,5 раза медленнее, чем математика «в лоб».

Можно попробовать просчитать всё на GPU, раз библиотека это позволяет. Для этого в pom.xml меняю зависимость для бэкенда, и все дела. У меня Geforce 1660 Super. Запускаю:

Получилось медленнее всех, не для моих «наноматриц» всё это. Слишком большой overhead на поддержание этих расчётов. На GPU их интересно профилировать:

Оперативной памяти потребляется в среднем всего 100 Мб, сколько GPU-памяти — мне неизвестно. Нагрузка на ЦП — 10% (с бэкендом на CPU было около 55%). С точки зрения Java, тут тишь да гладь. 

А если асинхронно? Страшно. Выше я упоминал, что ND4J хранит свои примитивы в off-heap memory и совсем непонятно, как на них синхронизироваться, JMM тут будет сходить с ума. Но ссылки под моим контролем, причём, как я писал выше, можно попробовать даже без синхронизации и замков. Я несколько раз пробовал запустить многопоточную версию, но эта реализация очень быстро «пожирала» все 32 ГБ моей ОП и подвешивала систему, либо Fedora просто «убивала» приложение как слишком прожорливое.

Библиотека ND4J взята из проекта для глубокого обучения (deep learning), где считаются поистине огромные матрицы, и в моём случае это как из пушки по воробьям. Необходимо найти грань, после которой ND4J будет действительно эффективна, а также разобраться с утечками памяти. Дело будущего.

По тем же причинам в Python’e библиотека numpy не даёт хороших результатов. Для небольших матриц уходить в натив и возвращаться обратно — дорого.

EJML 

EJML расшифровывается как Efficient Java Matrix Library, дословно — «Эффективная Java-библиотека для матриц». Здесь всё попроще, в натив и в память ОС не уходят, всё считаем в пределах Java. В библиотеке применяют три разных подхода к расчётам и программированию: процедурный, объектно-ориентированный и расчёты через выражения. Процедурный — набор статических методов, которые манипулирует матрицами (как и в моей реализации). Объектно-ориентированный подход использует объекты SImpleMatrix и flow-операции. Подход через выражения использует формулы, записанные в виде строки, например: 

eq.process("K = P*H'*inv( H*P*H' + R )"). Библиотека сама всё распарасит и посчитает. Самым быстрым является процедурный подход, поэтому буду использовать его.

Опять немного переписал нейронную сеть под новые сущности и математику. Причём в этот раз сразу оптимизирую все выражения. В EJML есть методы вроде multTransB(double, ma, mb), который за раз перемножает матрицы, транспонирует вторую и умножает результат на число. Тут тоже подумали над производительностью. Также уберу лишнее создание объектов, буду переиспользовать существующие, то есть максимально всё оптимизирую. Запускаю:

Отлично, отыграл секунд 10, а это 10%. Теперь асинхронная версия и без замков (помним, что это можно сделать):

Отлично. Я добился более чем пятикратного увеличения скорости, причём даже без потери точности. 

EJML в целом мне понравилась. Если посмотреть в «кишки», то там обёрнутый массив double[] и более продвинутая математика, а что-то совсем простое, как у меня. В некоторых методах оставлены комментарии, что операцию надо векторизовать. Библиотека обновляется часто, последняя версия выпущена месяц назад, поэтому предположу, что векторизацию некоторых операций сделают достаточно скоро, поэтому я и отказался от собственной реализации Vector API. Лучше просто подожду. 

Не понравилась документация, скромная она, спасает Java Doc, там получше. 

Выводы

Для расчётов с небольшими матрицами лучше всего использовать собственный прикладной класс, либо EJML. Обучать нейронную сеть в многопоточном режиме можно и нужно, причём в некоторых случаях даже без синхронизации. Для больших матриц и нейронных сетей необходимы дальнейшие эксперименты, чтобы определить порог, после которого становится эффективно использовать ND4J и расчёты на GPU.