Как работает gridsearchcv
Перейти к содержимому

Как работает gridsearchcv

Сеточный поиск лучшей архитектуры нейронной сети с помощью Keras и Sklearn

Добрый день, уважаемые читатели. Темой нашей сегодняшней статьи станет объединение библиотек Keras и Sklearn.

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

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

Импорт модулей и загрузка данных

Начнём наш сегодняшний эксперимент с импортирования всех необходимых функций и классов.

Не буду здесь объяснять назначение каждого импорта, я буду использовать все функции постепенно.

В качестве набора данных я предпочёл Breast Cancer. Этот датасет можно считать классическим в сфере классификации. Набор содержит 569 наблюдений, каждое из которых содержит 30 признаков.

Разделим набор на обучающие и тестовые данные. Для этого используем такой инструмент, как функцию train_test_split из модуля model_selection библиотеки sklearn. В моём случае тестовый набор будет составлять 20% от общего.

Пользовательский класс для построения модели

Немного забегая вперёд, скажу, что GridSearchCV использует подстановку гиперпараметров в классы моделей для дальнейшего их обучения. Однако использование класса Sequential() не очень подходит под эту задачу. Поэтому мы реализуем свой класс, в котором построим модель и приведём её к программному интерфейсу sklearn.

Я объявляю класс и также описываю, что его конструктор не принимает никаких аргументов, а также ничего не выполняет.

Класс GridSearchCV использует метод set_params(**params) для подстановки параметров.

Сюда я буду передавать список слоёв сети, тип оптимизатора, меру измерения ошибки, метрику оценивая, кол-во эпох и размер мини-пакета для обучения.

С помощью Keras класса KerasClassifier() я преобразовываю нашу сеть, реализованную в классе Sequential() , так, чтобы я мог использовать модель в различных операциях sklearn . Конструктор класса принимает функцию, которая должна возвращать саму нейронную сеть, кол-во эпох и размер батча.

Также мы переопределяем методы fit() , predict() и score() , т.к. их использует алгоритм сеточного поиска.

Примечание. Полное описание тела класса содержится в документе, ссылку на который вы сможете найти в конце статьи.

Сеточный поиск

Так выглядит словарь со всеми значениями параметров нашего класса:

Позже я рассмотрю инструмент, с помощью которого вы сможете визуализировать все три предложенные мною архитектуры.

Т.к. нейронные сети «требуют» масштабирования признаков, я создал конвейер с транформером, который будет проводить операцию стандартизации.

Далее я создал экземпляр класса, отвечающего за сеточный поиск.

И сразу запустил поиск на тренировочных данных:

Результат

Этот самый класс сохраняет лучшие параметры, а также лучшую модель, основываясь на тренировочных данных.

Теперь, основываясь на этих параметрах, я создад экземпляр класса Sequential() с такой архитектурой и параметра компилирования.

С помощью функции plot_model() (прежде чем использовать её, установите PyDot — pip/conda install pydot ) я сохраню изображение, которое будет содержать отображение архитектуры нашей сети.

С помощью Pillow.Image выводится изображение на экран:

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

Заключение

Сегодня я продемонстрировал способ автоматизировать поиск нужного кол-ва слоёв в нейронной сети, связав sklearn и Keras. Надеюсь, эта статья была для вас полезна. Желаю вам низких ошибок и надёжных опорных векторов.

Документ с расширением .ipynb, который вы сможете просмотреть и изменить на своё усмотрение, находится по ссылке.

Кроме того, рекомендую прочитать статью Обработка естественного языка с Python — NLP. А также подписывайтесь на группу ВКонтакте, Telegram и YouTube-канал. Там еще больше полезного и интересного для программистов.

Добавить комментарий

Ваш адрес email не будет опубликован. Обязательные поля помечены *