Всем привет, меня зовут Антон, я работаю в Сбере разработчиком Java, в продукте GigaIDE. В этой статье мы перепишем нейронную сеть c Python’а на Java, которая распознаёт рукописные цифры MNIST. Попробуем распознавать свои цифры, рисуя их мышкой, сделаем обратный запрос в сеть и заглянем в ее «мозги», а в конце сделаем выводы.

Я не имею отношения к разработке нейронных сетей, только использую их (GigaChat, GigaCode) для исполнения своих ежедневных профессиональных обязанностей. Однажды  захотелось хорошенько разобраться в нейронках, и для этого я прочитал несколько вводных простых книг, чтобы освежить свои знания и понимание всей «магии». Одной из них была книга «Создаём нейронную сеть» Тарика Рашида — хороший материал для начала.

После прочтения можно получить работающую нейронную сеть, правда, на Python’е. Но мне удобнее экспериментировать и изучать нейросеть на Java, поэтому я и занялся построчным переписыванием кода.

Как работает и как устроенна нейронная сеть?

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

Нейронная сеть представляет собой математическую модель, которая преобразует входной сигнал в выходной. Чаще всего сеть состоит из нескольких слоёв: входного, выходного и нескольких скрытых. Каждый слой представляет собой набор нейронов. Все нейроны одного слоя соединены с каждым нейроном следующего слоя. У каждой связи есть вес. При прохождении сигнала через связь он корректируется (умножается) исходя из веса связи. Значение сигнала в каждой связи складывается, пропускается через функцию активации и подаётся на выход нейрона. Сигнал с выхода передаётся на вход нейрона следующего слоя.

Для наглядности приведу картинку из Википедии. Красным отмечен нейрон.

Пример простой нейронной сети:

Зелёные — входные нейроны, в которые подаётся сигнал; голубые — нейроны скрытого слоя, в которых происходит вся «магия»; жёлтые — нейроны выходного слоя, то есть желаемый результат. Здесь в выходном слое всего один нейрон, это обычно не так, в выходном слое может быть произвольное количество нейронов.

Что делает нейронная сеть?

Нейронная сеть, представленная в книге Тарика Рашида выполняет классическую задачу распознавания картинок, то есть задачу классификации. Есть набор рукописных цифр MNIST, сеть использует его для обучения и проверки. Этот набор состоит из записей в виде картинки 28 на 28 пикселей и цифры, изображённой на картинке. Проще всего работать с MNIST как с CSV-файлом, где каждая строка — запись. Первое число в записи это эталон цифры, а далее 784 (28 на 28) значений от 0 до 255, кодирующих цвет пикселей на картинке.

Например:

Скрытый текст

7,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,84,185,159,151,60,36,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,222,254,254,254,254,241,198,198,198,198,198,198,198,198,170,52,0,0,0,0,0,0,0,0,0,0,0,0,67,114,72,114,163,227,254,225,254,254,254,250,229,254,254,140,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,17,66,14,67,67,67,59,21,236,254,106,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,83,253,209,18,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,22,233,255,83,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,129,254,238,44,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,59,249,254,62,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,133,254,187,5,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,9,205,248,58,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,126,254,182,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,75,251,240,57,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,19,221,254,166,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,3,203,254,219,35,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,38,254,254,77,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,31,224,254,115,1,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,133,254,254,52,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,61,242,254,254,52,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,121,254,254,219,40,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,121,254,207,18,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0

Значения отличные от нуля, это градации серого. CSV, конечно, хорошо, но хотелось бы увидеть всё-таки картинки вместо чисел

Всего в наборе 60 000 картинок для обучения и 10 000 картинок для тестирования и проверки сети. Получается, что на каждую цифру в наборе приходится 6 000 картинок. Мне стало интересно, как можно 6 000 раз по разному написать цифру «ноль», или цифру «один», а также взглянуть на эти картинки «вживую». Для этого я написал небольшое приложение для просмотра наборов MNIST — MnistCsvViewer. Его интерфейс:

Взглянув на цифры, я увидел, что они действительно различаются и написаны в американской манере. Обычно цифру «один» мы пишем двумя чертами: короткой и длинной; короткая под некоторым углом к длинной. В датасете есть и такие варианты, но чаще всего цифра «один» представляет собой просто слегка наклонённую черту.

Другие цифры тоже имеют свои региональные особенности, например, цифра 9 чаще всего без завитка внизу и имеет примерно такой вид:


Вообще, хочется более пристально проанализировать изображения цифр и попробовать их отсортировать от самых простых к сложным (хотя бы по количеству значимых пикселей), описав эволюцию их написания. Возможно это сделать не получится, задача на будущее.

Разобравшись с картинками, я приступил к переписыванию нейросети на Java.

Init

Нейронная сеть имеет три слоя. Входной слой из 784 нейронов, скрытый — из 200 нейронов и выходной — из 10 нейронов. Во входной слой подаём значения из CSV, нормализовав их. В скрытом слое происходят вычисления. В выходном слое каждый нейрон представляет собой цифру от 0 до 9. После прохождения сигнала через сеть на каждом выходном нейроне появляются значение от 0 до 1, и чем ближе к единице, тем выше «вероятность», что цифра распознана.

Нейронка на Python’e представляет собой класс с тремя методами: init (конструктор), query и train. Прямой перевод названий раскрывает их смысл. Я перепишу всё строчка в строчку, чтобы можно было бы воспользоваться комментариями из исходника и минимизировать свои ошибки.

В конструкторе задаём количество входных, скрытых и выходных узлов (нейронов), а также коэффициент обучения. Затем у скрытого и выходного слоя создаём две матрицы весов и заполняем их начальными значениями, которые очень важны. Можно задать веса, близкие к нулю, что, в целом, будет работать. В книге советуют поступить более хитрым способом: назначить начальные значения весов в соответствии с нормальным распределением с центром в нуле и со стандартным отклонением, величина которого обратно пропорциональна корню из количества узлов матрицы. На слух сложновато звучит, на языке Java выглядит так:

random.nextGaussian(0, Math.pow(matrix.length, -0.5)); 

Оказывается, в Random есть для этого специальный метод.

Код конструктора прост:

public NeuralNetwork(int inputNodesNumber,
                     int hiddenNodesNumber,
                     int outputNodesNumber,
                     double learningRate) {

    Checker.checkNodesNumbers(inputNodesNumber, hiddenNodesNumber, outputNodesNumber);

    this.inputNodesNumber = inputNodesNumber;
    this.hiddenNodesNumber = hiddenNodesNumber;
    this.outputNodesNumber = outputNodesNumber;
    this.learningRate = learningRate;

    initWeights();
}

Для детерминированного результата необходимо иметь возможность задать веса точно. Для этого добавил ещё две стратегии, когда все веса нули и когда все веса единицы. Такие стратегии заведомо обречены на провал, но зато я смогу сравнить свои результаты запросов к сети с питоновскими, и если они совпадут, то это можно будет считать успехом.

Query

Думаю, ни для кого не секрет, что почти все расчёты в нейронной сети сводятся к перемножению матриц. В Python’e для этого есть специальный пакет numpy, который в сложных случаях использует Си (скомпилированный код). В Java тоже есть библиотеки для расчёта матриц, и не одна. Есть даже работы, в которых сравнивается их производительность. Но я не хотел сразу обрастать чужими зависимостями, которые всё делают за меня. Для самообучения решил написать собственный прикладной класс, который будет заниматься математикой, причём в лоб по определению. После этого реализовать запрос к сети уже было не сложно, переписав строчка в строчку с Python’а:

public double[][] query(double[] inputs) {
    if (inputs.length != inputNodesNumber) {
        throw new IllegalArgumentException("Wrong count of inputs");
    }

    double[][] inputMatrix = MatrixUtils.transformToMatrix(inputs);
    double[][] hiddenInputs = MatrixUtils.multiply(inputToHiddenWeights, inputMatrix);
    double[][] hiddenOutputs = MatrixUtils.applyFunction(hiddenInputs, activationFunction);
    double[][] finalInputs = MatrixUtils.multiply(hiddenToOutputsWeights, hiddenOutputs);

    return MatrixUtils.applyFunction(finalInputs, activationFunction);
}

После этого я создал одинаковые сети на Python’e и Java с изначальными весами, равными нулю и единице. Затем в каждой сделал запрос с одинаковым input’ом. Убедившись, что результаты совпадают, перешёл к реализации обучения сети.

Train

С обучением немного сложнее в плане математики, для этого нужно сделать запрос в нейросеть, получить результат, затем вычислить ошибку и на её основе скорректировать веса методом обратного распространения. Формула корректировки весов в книге выглядит так:

На Python’e она записывается так:

self.who += self.lr * numpy.dot((output_errors * final_outputs * 
                 (1.0 - final_outputs)), numpy.transpose(hidden_outputs))

А на Java с учётом моего прикладного математического класса принимает такой вид:

double[][] deltaHiddenToOutputs = MatrixUtils.multiply(
        MatrixUtils.multiply(
                MatrixUtils.multiplyByElements(
                        outputErrors,
                        MatrixUtils.multiplyByElements(
                                finalOutputs,
                                MatrixUtils.subtract(1, finalOutputs))),
                MatrixUtils.transpose(hiddenOutputs)),
        learningRate);

hiddenToOutputsWeights = MatrixUtils.add(hiddenToOutputsWeights, deltaHiddenToOutputs);

Немножко монструозно, но это потом исправим.

Переписав всё строчка в строчку, я получил рабочую сеть на Java (NeuralNetwork.class) и перешел к обучению сети.

Обучение и проверка

Обучал на 60 000 картинок и в 5 эпох. Казалось бы, сеть небольшая, задача не сверхсложная, но всё равно обучение занимает ощутимое время, примерно по минуте на каждую эпоху. После обучения и проверки нейросети на тестовом множестве в 10 000 картинок получил точность распознавания 0,975, то есть ошибка всего 2,5%.

В Python’е процесс обучения и проверки нейросети происходит прямо в том же скрипте, где она создавалась. В своих же экспериментах я сделал отдельный класс NetworkTrainer, который занимается обучением (подает картинки в метод train) и проверкой (подаёт в метод query тестовую картинку и сравнивает результат с эталоном). При проверке нейросети решил сохранить те картинки, что не получилось распознать, положив их в CSV-файл, а подобные CSV-файлы я умею открывать и просматривать своим MnistCsvViewer:

Таким образом я увидел те цифры, которые не смогла распознать сеть. Да, есть сложные варианты, но для человека распознать большую часть из этого набора не составит труда.

Распознай теперь меня

Теперь, когда у меня есть рабочая и обученная нейросеть, я хочу, чтобы она распознавала мои «каракули». Писать на бумажке цифры, потом фотографировать или сканировать их, вырезать по одной для меня показалось слишком утомительным. Я решил, что буду рисовать с помощью «мышки», в целом почерк это передаёт. Для этого я реализовал сохранение нейросети в файл и чтение её из файла. Задача несложная, нужно сохранить количество узлов в каждом слое, коэффициент обучения и веса. Зная структуру файла, можно прочитать ее и передать в конструктор. 

Для визуализации и рисования своих цифр я написал второе приложение, которое позволяет открыть сохранённую нейросеть и посмотреть её структуру — NeworkViewer.

На второй вкладке приложения можно порисовать и посмотреть, как нейросеть распознаёт мои цифры:

И оказалось, что очень плохо… точность там примерно 50-60%, про 3% ошибки речи и не идёт. Для экспериментов я добавил возможность изменять размер кисти, а также добавил размытие (blur), чтобы рисунок был более похож на рукописные картинки (края линий не такие чёткие).

Ничего не помогало. Несколько раз перепроверил — всё верно. Я вижу ответ нейросети и уровень сигнала для каждой цифры. Зачастую, когда сеть угадывает и показывает 0,95, я рисую почти такую же цифру и сигнал может стать 0,95 совсем на другой цифре. Мне не понятно, как можно кружок в середине картинки принять за что-либо другое, кроме нуля, однако нейросеть это прекрасно делает:

Back Query

Интересно, но можно развернуть направление прохождение сигнала в сети: на выход подать желаемую цифру, а на входе получить картинку, то есть заглянуть в «мозги» сети. Это сделано в книги, это же повторил и я. На третьей вкладке приложения можно увидеть эти образы:

Рассматривая их, кроме ноля и семи, я ничего не увидел, может быть ещё немного девятки. Думал, что должны были получится сильно размытые, но узнаваемые силуэты цифр, но нет.

В книге Тарика Рашида приведён такой результат:

Тут чёткий ноль. Может быть, узнаётся двойка и пятёрка, остальные цифры, особенно, например, 8 — мазня.

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

Выводы

Я сделал первые шаги в изучении нейронных сетей, и пока что они меня не впечатлили. Обучение слишком долгое, количество обучающих множеств — огромно, а практический результат слабый: мои цифры нейросеть распознаёт очень плохо. Да, она обучалась не на моих цифрах, но как исправить ситуацию? Самому нарисовать 6000 нулей и переобучить сеть? И так для каждого человека?