Всем привет, меня зовут Антон, я работаю в Сбере разработчиком 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 нулей и переобучить сеть? И так для каждого человека?
