Сжимаем трансформеры: простые, универсальные и прикладные способы cделать их компактными и быстрыми

transformer_press


Сейчас в сфере ML постоянно слышно про невероятные "успехи" трансформеров в разных областях. Но появляется все больше статей о том, что многие из этих успехов мягко говоря надуманы (из недавнего помню статью про пре-тренировку больших CNN в компьютерном зрении, огромную MLP сетку, статью про деконструкцию достижений в сфере трансформеров).


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


Тем не менее у self-attention модуля есть ряд плюсов: (i) относительная простота при правильной реализации (ii) простота квантизации (iii) относительная эффективность на коротких (до нескольких сотен элементов) последовательностях и (iv) относительная популярность (но большая часть имплементаций имеет код раздутый раз в 5).


Также есть определенный пласт статей про улучшение именно асимптотических свойств self-attention модуля (например Linformer и его аналоги). Но несмотря на это, если например открыть список пре-тренированных языковых моделей на основе self-attention модулей, то окажется, что "эффективных" моделей там буквально пара штук и они были сделаны довольно давно. Да и последовательности длиннее 500 символов нужны не очень часто (если вы не Google).


Попробуем ответить на вопрос — а как существенно снизить размер и ускорить self-attention модуль и при этом еще удовлетворить ряду production-ready требований:



Тут важно еще сделать оговорку, что проседание качества будет тем сильнее, чем сложнее ваша задача. Например на бинарной классификации ужаться можно сколько угодно (да им может проще использовать более простые методы), а вот на sequence-to-sequence задачах будут моменты.


Простейшие оптимизации


Какое-то время назад в сети проскакивала такая презентация. Если абстрагироваться от ее "космической" (когда я читаю такие материалы, мне кажется что авторы строят башню на луну) академической составляющей, то между строк можно найти такую информацию:



В целом, вопрос состоит в том, стоит ли сразу тренировать более компактные модели или все-таки дистиллировать, но это зависит уже от вашего конкретного кейса.


Плюсы:



Минусы:



Квантизация


В PyTorch где-то примерно начиная с версии 1.3 завезли динамическую квантизацию Linear и LSTM модулей. Не считая подготовительного кода она действительно (без преувеличения) делается в одну строчку кода:


quantized_model = torch.quantization.quantize_dynamic(
    model, {torch.nn.Linear}, dtype=torch.qint8
)

Примеры можно посмотреть тут и тут.


Плюсы:



Минусы:



Факторизация


Я кратко описывал этот подход в заметке на канале в телеграме. Идея состоит в том, чтобы взять готовую сетку, применить Singular Value Decomposition (который недавно тоже завезли в PyTorch) к матрице весов, заранее выбрав нужный уровень разреженности.


Чтобы не плодить лишние классы, для начала можно элегантно сделать monkey-patching своей модели при ее загрузке, заменив Linear на FactorizedLinear модуль.


Плюсы:



Минусы:



Прунинг и дистилляция


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


FNet


Недавно появилась вот такая статья. В ней по сути предложили заменить self-attention на разложение Фурье. Получается, что размер модели снижается в два раза, скорость на GPU становится меньше чуть ли не на порядок (чего нельзя сказать про CPU). Якобы на задачах авторов потеря качества в районе 10%.


Плюсы:



Минусы:



На десерт — оптимизации самого PyTorch


Все знают про fusion сверток, бнорма и relu. Но в недавней версии 1.9 добавили еще "заморозку" модели и inference mode. Я протестировал их, и при прочих равных inference mode не добавил скорости, а заморозка модели докинула 14% скорости. Тут важно отметить, что все очевидные оптимизации с моделями уже были проделаны, поэтому приросты уже небольшие. Еще важный момент состоит в том, что если вы возьмете медленную модель и примените какой-либо хак из этой статьи, вы получите условно заявленные x2. Но если вы примените много хаков сразу, то в какой-то момент начнет показывать себя закон убывающей полезности.


Плюсы:



Минусы:



Примерное сравнение


Оптимизация Снижение размера Ускорение Комментарий
Сужение и 2 головы 2x 2-3x Модель де-факто становится меньше
Квантизация 4x 2x на CPU Работает только на CPU
Факторизация 2-4x 2-3x на GPU, нет на CPU Можно и сильнее, но дальше качество проседает
Замена attention на FFT 2x 2x на CPU, 7x на GPU
Заморозка и прочие оптимизации - 15-25% Заморозка, fusion

Качество итоговых моделей


Поскольку все очень зависит от сложности вашей задачи и конкретики, точных цифр приводить не буду, скорее отранжирую способы по качеству и сложности достижения:


Оптимизация Сложность имплементации Сложность получения качества
Заморозка и прочие оптимизации Из коробки По идее ничего не меняется, но приросты небольшие
Квантизация Из коробки, если нормальный код Даже на сложных задачах почти нет просадок
Сужение и 2 головы Из коробки, если нормальный код Возможно будет дольше сходиться
Факторизация 70 строк кода Нужно тюнить
Замена attention на FFT 10 строк кода Как мне показалось, дольше еще дольше тюнить

Краткий итог


Просуммируем все вышесказанное. Если вынести за скобки подкрутку гипер-параметров самой модели, то методы оптимизации можно разделить на два класса (i) работающие почти из коробки, но с меньшим эффектом (ii) и требующие тюнинга, но суммарно дающие больше результата.


К первому можно отнести заморозку и квантизацию. В сумме они дают приятное уменьшение размера модели (4x) и ускорение в районе 2-3x на CPU.


Ко второму можно отнести факторизацию и FFT. Их стоит рассматривать как некую дополнительную оптимизацию, причем они скорее всего исключают друг друга. В сумме с первым типов методов можно получить суммарное снижение размера модели почти на порядок и ускорение тоже почти на порядок. Если при этом еще подкрутить гипер-параметры модели, то "порядок" в принципе не кажется недостижимым.


Как сделать ускорение на два порядка, если честно я не знаю. Возможно вы знаете?

@snakers4
21.06.2021 12:25 UTC
Первоисточник

Комментарии

@VPryadchenko
21.06.2021 16:04 UTC
0

Прунинг и дистилляция

Я пробовал дистиллировать...

А что насчёт прунинга?

@snakers4
21.06.2021 16:23 UTC
0

Не знаю как делать прунинг, чтобы он в продакшене давал приросты по скорости. Все эти "научные" статьи, которые сначала хвастаются 95% степенью спарсификации (на матрице fp32 множителей размером со всю модель), а потом стыдливо в одно предложение говорят "ой, ну алгоритмов перемножения sparse матриц еще нет", как-то не впечатляют.


С другой стороны именно в случае self-attention модуля обрезание голов можно по идее назвать structured прунингом.

21.06.2021 16:31 UTC
+1

То то и оно, что пользы от unstructured прунинга действительно немного. А вот structured вполне себе заходит в прод. По специфике своей работы с трансформерами не имею дела, но очень интересно посмотреть, как они прунятся.

21.06.2021 16:34 UTC
0

А какой подход к structured прунингу реально работает? Есть рецепты? Укладывается в 50 строк?)

21.06.2021 16:39 UTC
0

Реально (в моих CV задачах) работает модификация прунинга, основного на Taylor expansion, рецепты которого есть, но (пока) делиться которыми не могу. В 50 строк не укладывается, но уложится в меньшее их число, если алгоритмы запихнуть в либу)

21.06.2021 16:59 UTC
0

А в чем суть? Ряды Тейлора? Есть какие-то примеры / статьи / референсные имплементации?

22.06.2021 08:28 UTC
0

На текущий момент, самый точный метод, известный мне по прунингу — Optimal Brain Surgeon. В этом подходе необходимо считать матрицу Гессе модели (берется приближение функцией второго порядка в окрестности минимума), и оптимальный сдвиг берется такой, что он убирает часть весов и рост ошибки наименьший. Оригинальная работа очень старая, поэтому там совсем игрушечные модели, в своей практике я работаю с моделями порядка несколько тысяч весов (ТЗ требует сверхкомпактный моделей для обработки сигналов) — и на текущий момент OBS — дает самое лучшее качество при том же количество весов. Про наличие открытых имплементаций увы не знаю. Самый простой, но грубоватый подход — magnitude based. Он добавлен в стандартные Фреймворк, есть в Tensorflow (Tutorial). Но без серьезной просадки в качестве, полагаю, более чем раза в 2 модель ужать не получится в большинстве задач.

22.06.2021 08:28 UTC
0

А как эти методы потом работают в продакшене? А сетками размера реальных сеток из CV это все работает?

22.06.2021 08:39 UTC
0

В Signal processing это работает на реальных архитектурах для продакшена. Детали обычно являются корпоративной тайной — поэтому в открытом доступе не появляются. В компьютерном зрении считать матрицу Гессе честно слишком накладно (ее размер для типичной архитектуры типа Resnet, EfficientNet будет порядка 10^7 x 10^7 — 10^14), что никуда не влезет. Но методы основанные на приближениях как-то работают https://arxiv.org/pdf/2101.08940.pdf. В приведенной работе где-то в 3-4 раза ужали размер модели (но правда все на CIFAR-10), а интересно было в  real-life applications увидеть пример.

22.06.2021 13:57 UTC
0
Ну, собственно, для структурного прунинга (например, вырезания фильтров в сверточных сетях) достаточно маскировать карты признаков набором масок, считать для них Гессианы (притом, в окрестности локального минимума они факторизуются на матрицы градиентов) послойно, и для каждого слоя находить, соответственно, оптимальные подмножества ненулевых маскирующих значений.
22.06.2021 08:44 UTC
0

Еще интересная работа, на мой взгляд, обучение разреженной модели с перешиванием весов — RigL. Веса убираются по magnitude-based прунингу, а новые выращиваются на основе градиента лосс-функции. В примерах есть сжатие до 4-5 раз почти без просадки качества и есть открытая реализация на Tensorflow

22.06.2021 08:48 UTC
0

А в проде такие модели реально быстрее?

22.06.2021 09:13 UTC
0

Нашел пример по ускорению Берта благодаря прунингу (https://blog.rasa.com/pruning-bert-to-accelerate-inference/). Но начинает уже просаживаться качество при этом. В CV — в основном работы смотрят на (количество параметров и флопс)[https://blog.rasa.com/pruning-bert-to-accelerate-inference/], но в экспериментах просто кладется маска на веса, что не ускоряет модель на практике. В обработке сигналов обычно смотрят на скорость работы — она фиксированная, а чтобы число умножений было не более определенного порога