/
Tags: искусственный интеллект
ISBN: 978-5-93700-192-4
Text
Григорий Сапунов
Глубокое обучение с JAX
Deep Learning with JAX
G R I G O RY S A P U N OV
Глубокое обучение с JAX
ГРИГОРИЙ САПУНОВ
Москва, 2025
УДК 004.85JAX
ББК 16.6
С19
С19
Сапунов Г.
Глубокое обучение с JAX / пер. с англ. А. В. Снастина. – М.: ДМК Пресс, 2025. –
526 с.: ил.
ISBN 978-5-93700-192-4
Книга обучает созданию эффективных нейронных сетей с применением биб
лиотеки JAX от Google. Специальные инструменты JAX помогут вам справиться
с проблемами нехватки производительности, характерными для глубокого обучения
на больших данных. Примеры реалистичных проектов и листинги исходного кода
с подробными комментариями показывают, как парадигма функционального про
граммирования JAX улучшает сочетаемость различных программных компонентов
и возможности распараллеливания.
Для читателей, имеющих опыт программирования на языке Python и знакомых
с принципами глубокого обучения.
УДК 004.85JAX
ББК 16.6
Copyright © DMK Press 2025. Authorized translation of the English edition.
© 2024 Manning Publications. This translation is published and sold by permission of Manning
Publications, the owner of all rights to publish and sell the same.
Все права защищены. Любая часть этой книги не может быть воспроизведена в какой
бы то ни было форме и какими бы то ни было средствами без письменного разрешения вла
дельцев авторских прав.
ISBN 978-1-6334-3888-0 (англ.)
ISBN 978-5-93700-192-4 (рус.)
© Manning Publications, 2024
© Перевод, оформление, издание, ДМК Пресс, 2025
Моим родителям, которые поощряли все мои увлечения,
буквально окружали меня великолепными книгами
и подарили мне первый компьютер в те времена,
когда он все еще оставался предметом роскоши
(и использовался в основном для развлечений)
Оглавление
Часть I
ПЕРВЫЕ ШАГИ .......................................................................................... 23
1
Когда и зачем используется JAX......................................................................... 25
2
Первая программа в JAX....................................................................................... 46
Часть II ЯДРО JAX ...................................................................................................... 79
3
Работа с массивами. .............................................................................................. 81
4
Вычисление градиентов......................................................................................132
5
Компиляция кода..................................................................................................178
6
Векторизация кода................................................................................................223
7
Распараллеливание вычислений.......................................................................254
8
Использование сегментирования тензоров...................................................305
9
Случайные числа в JAX. .......................................................................................333
10
Работа с pytree........................................................................................................368
Часть III ЭКОСИСТЕМА ...........................................................................................393
11
Высокоуровневые библиотеки поддержки нейронных сетей....................395
12
Другие члены экосистемы JAX...........................................................................445
Содержание
Оглавление............................................................................................................ 6
Предисловие........................................................................................................ 12
Благодарности................................................................................................... 14
О книге. ............................................................................................................... 16
Об авторе........................................................................................................... 21
Об иллюстрации на обложке............................................................................ 22
Часть I
1
2
ПЕРВЫЕ ШАГИ ........................................................................... 23
Когда и зачем используется JAX. .................................................. 25
1.1
Почему нужно использовать JAX........................................................ 30
Производительность вычислений............................................... 30
Функциональный подход............................................................... 34
Экосистема JAX........................................................................... 35
1.2
Чем JAX отличается от NumPy............................................................. 37
1.2.1 JAX как NumPy............................................................................. 38
1.2.2 Компонуемые трансформации.................................................... 39
1.3
Чем JAX отличается от TensorFlow и PyTorch.................................... 42
Резюме................................................................................................................ 45
1.1.1
1.1.2
1.1.3
Первая программа в JAX.................................................................... 46
2.1
2.2
2.3
2.4
Учебная задача машинного обучения: классификация
рукописных цифр.................................................................................. 47
Общий обзор проекта глубокого обучения с использованием
JAX............................................................................................................ 49
Загрузка и подготовка набора данных............................................... 51
Простая нейронная сеть с использованием JAX............................... 54
2.4.1
2.4.2
Инициализация нейронной сети.................................................. 56
Нейронная сеть с прямой связью................................................. 58
Содержание
8
2.5
2.6
2.7
2.8
2.9
vmap: автоматически векторизованные вычисления
для обработки пакетов. ........................................................................ 61
Autodiff: как вычислять градиенты, ничего не зная
о производных....................................................................................... 65
Функция потерь........................................................................... 67
Определение градиентов............................................................. 68
Шаг обновления градиента.......................................................... 68
Цикл тренировки......................................................................... 70
JIT: компиляция кода для ускорения его выполнения.................... 72
Сохранение и развертывание модели................................................ 74
2.6.1
2.6.2
2.6.3
2.6.4
Чистые функции и компонуемые трансформации: почему
они так важны........................................................................................ 76
Упражнение 2.1................................................................................................. 77
Резюме................................................................................................................ 78
Часть II
3
ЯДРО JAX. ........................................................................................ 79
Работа с массивами.............................................................................. 81
3.1
3.2
3.3
3.4
Обработка изображений с использованием массивов NumPy....... 82
Загрузка изображения в массив NumPy....................................... 84
Выполнение простых операций предварительной
обработки с изображением......................................................... 88
3.1.3 Добавление шума в изображение................................................. 90
3.1.4 Реализация фильтрации изображения........................................ 92
3.1.5 Сохранение тензора как файла изображения............................. 98
Массивы в JAX........................................................................................ 99
3.2.1 Переход на NumPy-подобный API JAX.........................................100
3.2.2 Что такое Array?.......................................................................102
3.2.3 Операции, связанные с устройствами.......................................104
3.2.4 Асинхронная диспетчеризация...................................................111
3.2.5 Выполнение вычислений на TPU..................................................112
Отличия от NumPy................................................................................117
3.3.1 Неизменяемость.........................................................................117
3.3.2 Типы............................................................................................121
3.1.1
3.1.2
Интерфейсы высокого и низкого уровней: jax.numpy
и jax.lax.................................................................................................126
Базисные элементы управления потоком выполнения..............127
Повышение (расширение) типа..................................................129
Упражнение 3.1................................................................................................130
Резюме...............................................................................................................130
3.4.1
3.4.2
4
Вычисление градиентов.....................................................................132
4.1
4.2
Различные способы вычисления производных...............................134
4.1.1
4.1.2
4.1.3
4.1.4
Дифференцирование вручную.....................................................136
Символьное дифференцирование................................................137
Численное дифференцирование...................................................139
Автоматическое дифференцирование.......................................141
Вычисление градиентов с использованием autodiff.......................144
4.2.1 Работа с градиентами в TensorFlow..........................................146
Содержание
9
4.2.2
4.2.3
4.2.4
4.2.5
Работа с градиентами в PyTorch...............................................148
Работа с градиентами в JAX......................................................149
Производные более высоких порядков.........................................158
Вариант со многими переменными............................................160
4.3
Прямой и обратный режимы autodiff. ..............................................164
4.3.1 Трассировка оценок.....................................................................165
4.3.2 Прямой режим и jvp()................................................................166
4.3.3 Обратный режим и vjp()...........................................................171
4.3.4 Материалы для более глубокого изучения..................................175
Резюме...............................................................................................................176
5
6
Компиляция кода....................................................................................178
5.1
Использование компиляции. .............................................................180
Использование JIT-компиляции...................................................181
Чистые функции и процесс компиляции.....................................190
5.2
Внутренний механизм JIT...................................................................193
5.2.1 Jaxpr – промежуточное представление для программ JAX........193
5.2.2 XLA..............................................................................................204
5.2.3 Использование AOT-компиляции................................................211
5.3
Ограничения JIT-компиляции............................................................216
5.3.1 Чистые функции и функции, не являющиеся чистыми...............216
5.3.2 Точные числовые данные.............................................................216
5.3.3 Условные выражения, использующие значения входных
параметров.............................................................................................216
5.3.4 Медленная компиляция...............................................................217
5.3.5 Методы класса............................................................................219
5.3.6 Простые функции.......................................................................221
Упражнение 5.1................................................................................................221
Резюме...............................................................................................................221
5.1.1
5.1.2
Векторизация кода...............................................................................223
6.1
Различные способы векторизации функции...................................224
6.1.1
6.1.2
6.1.3
6.1.4
Простейшие методики...............................................................226
Векторизация вручную...............................................................229
Автоматическая векторизация.................................................230
Сравнение скорости выполнения................................................231
6.2
Управление поведением vmap().........................................................233
6.2.1 Управление осями массива для выполнения
преобразования...........................................................................233
6.2.2 Управление осями выходного массива........................................237
6.2.3 Использование именованных аргументов..................................238
6.2.4 Использование стиля декоратора..............................................241
6.2.5 Использование коллективных операций.....................................241
6.3
Варианты использования vmap() из реальной практики...............243
6.3.1 Обработка пакетов данных.......................................................244
6.3.2 Пакетная обработка моделей нейронных сетей.......................246
6.3.3 Поэлементные градиенты..........................................................247
6.3.4 Векторизация циклов.................................................................249
Резюме...............................................................................................................253
Содержание
10
7
Распараллеливание вычислений...................................................254
7.1
7.2
7.3
Распараллеливание вычислений с помощью pmap()......................256
7.1.1
7.1.2
Установка начальных условий задачи........................................257
Использование pmap (почти) так же, как vmap...........................260
Управление поведением pmap().........................................................268
7.2.1 Управление отображением осей входных и выходных
данных.........................................................................................268
7.2.2 Использование именованных осей и коллективных
операций......................................................................................276
Пример программы тренировки нейронной сети
с распараллеливаниемпо данным....................................................285
Подготовка данных и структуры нейронной сети....................286
Реализация процедуры тренировки с распараллеливанием
по данным....................................................................................290
7.4
Использование конфигураций с несколькими хостами. ...............296
Резюме...............................................................................................................303
7.3.1
7.3.2
8
Использование сегментирования тензоров.........................305
8.1
8.2
Основы сегментирования тензоров..................................................307
8.1.1
8.1.2
8.1.3
8.1.4
8.1.5
8.1.6
8.1.7
Сетка устройств........................................................................310
Позиционное сегментирование..................................................310
Пример с применением двумерной сетки...................................311
Использование репликации.........................................................316
Ограничения сегментирования..................................................319
Сегментирование с именованием...............................................320
Стратегия размещения устройств и ошибки...........................322
Многослойный перцептрон с применением
сегментирования тензоров.................................................................325
8.2.1
8.2.2
Восьмиканальное распараллеливание данных............................325
Четырехканальное распараллеливание по данным,
двухканальное распараллеливание тензора...............................327
Резюме...............................................................................................................332
9
Случайные числа в JAX........................................................................333
9.1
Генерация случайных данных............................................................335
9.1.1
9.1.2
9.1.3
Загрузка набора данных..............................................................337
Генерация случайного шума........................................................340
Выполнение случайного расширения данных..............................344
9.2
Отличия от NumPy................................................................................347
9.2.1 Как работает NumPy.................................................................347
9.2.2 Начальное число и состояние в NumPy.......................................349
9.2.3 PRNG в JAX..................................................................................353
9.2.4 Более подробная конфигурация JAX PRNG..................................361
9.3
Генерация случайных чисел в реальных приложениях..................362
9.3.1 Создание полного конвейера расширения данных......................363
9.3.2 Генерация случайных инициализаций для нейронной сети........364
Резюме...............................................................................................................366
Содержание
10
11
Работа с pytree.........................................................................................368
10.1
10.2
Представление сложных структур данных в форме pytree............370
Функции для работы с pytree..............................................................376
10.2.1 Использование tree_map()..........................................................376
10.2.2 Преобразование pytree в плоскую структуру
и восстановление древовидной формы.......................................380
10.2.3 Использование tree_reduce().....................................................383
10.2.4 Транспонирование pytree............................................................384
10.3 Создание специализированных узлов pytree. .................................387
Резюме...............................................................................................................391
Часть III
11
ЭКОСИСТЕМА ............................................................................393
Высокоуровневые библиотеки поддержки
нейронных сетей.....................................................................................395
11.1
Классификация изображений MNISTс использованием
многослойного перцептрона..............................................................397
11.1.1 Многослойный перцептрон в Flax...............................................397
11.1.2 Библиотека трансформации градиентов Optax.......................405
11.1.3 Тренировка нейронной сети с применением Flax.......................408
11.2 Классификация изображений с использованием ResNet...............413
11.2.1 Управление состоянием в Flax....................................................415
11.2.2 Сохранение и загрузка модели с использованием Orbax.............422
11.3 Использование экосистемы Hugging Face........................................425
11.3.1 Использование предварительно натренированной модели
из хранилища Hugging Face Model Hub.......................................427
11.3.2 Более подробное изучение процессов точной настройки
и предварительной тренировки.................................................435
11.3.3 Использование библиотеки диффузоров....................................438
Резюме...............................................................................................................443
12
Другие члены экосистемы JAX......................................................445
12.1
Экосистема глубокого обучения. .......................................................446
12.1.1 Высокоуровневые библиотеки поддержки нейронных сетей......446
12.1.2 Большие языковые модели в JAX.................................................448
12.1.3 Библиотеки утилит...................................................................451
12.2 Модули машинного обучения. ...........................................................454
12.2.1 Обучение с подкреплением..........................................................454
12.2.2 Прочие библиотеки машинного обучения...................................455
12.3 Модули JAX для других сфер деятельности......................................457
Резюме...............................................................................................................459
Приложение A. Установка JAX........................................................................460
Приложение B. Использование Google Colab...................................................464
Приложение C. Использование Google Cloud TPU...........................................467
Приложение D. Экспериментальные средства распараллеливания............472
Предметный указатель. ..................................................................................513
Предисловие
JAX – это мощная библиотека на языке Python, созданная компанией
Google для глубокого обучения и высокопроизводительных вычис
лений. Библиотека широко используется в научных исследованиях
в области машинного обучения и позиционируется как третий по
известности и распространенности фреймворк глубокого обучения,
уступая только TensorFlow и PyTorch. В частности, это практически
«штатный» фреймворк в таких компаниях, как DeepMind, а исследо
вания Google все в большей степени основываются на JAX.
Что мне действительно нравится в JAX, так это его акцент на функ
циональном программировании в глубоком обучении. Этот фрейм
ворк предоставляет надежные функциональные преобразования,
включая вычисление градиента, JIT-компиляцию посредством XLA,
автоматическую векторизацию и возможности распараллеливания.
JAX поддерживает графические (GPU) и тензорные (TPU) процессо
ры, обеспечивая впечатляющую производительность.
Сейчас самое подходящее время для подробного и глубокого
изучения JAX, так как его экосистема стремительно расширяется.
Несмотря на то что фреймворк существует уже несколько лет, все же
наблюдается очевидная недостаточность всеобъемлющих ресурсов
для новичков. Хотя веб-сайт JAX предлагает обширную документа
цию и поддерживающее сообщество, объединение всех источников
информации, особенно при интеграции других библиотек, может
показаться обескураживающим.
Эта книга создана для тех, кто стремится освоить JAX. Ее цель –
объединить самую важную информацию в одном месте и помочь
читателям понять концепции JAX, улучшить их навыки и способ
ность применять JAX в проектах и исследованиях.
Предполагается знание базовых принципов глубокого обучения
и наличие опыта практической работы с языком Python. В книге не
рассматриваются основы глубокого обучения, поскольку по этой
Предисловие
13
теме существует множество информационных ресурсов. Вместо это
го все внимание сосредоточено исключительно на JAX, хотя при не
обходимости главные концепции глубокого обучения все же кратко
излагаются. Такой подход должен оказаться полезным для тех, кто
не имеет опыта работы в сфере глубокого обучения, например фи
зикам.
JAX – это нечто большее, нежели просто фреймворк глубокого
обучения. Диапазон его модулей постоянно расширяется и выходит
за рамки глубокого обучения, создавая и укрепляя потенциальные
возможности в области дифференцируемого программирования,
крупномасштабных физических имитаций и многих других. Наде
юсь, что эта книга также будет полезна тем, кто интересуется подоб
ными приложениями.
JAX продолжает развиваться, поэтому несколько глав книги суще
ственно обновились. Но не стоит беспокоиться о возможных изме
нениях в будущем, потому что основополагающие знания, которые
вы получите при чтении, останутся применимыми в следующих
версиях JAX.
Благодарности
Написание этой книги заняло больше времени, чем я предполагал.
В процессе работы над ней я сменил несколько стран, да и версии
JAX менялись. Некоторые главы пришлось переписывать. Но теперь
все готово!
Прежде всего хочу поблагодарить мою семью – жену Милу, сыно
вей Даню и Федю. Вы так долго скучали без моего внимания! Но все
время поддерживали меня.
Хочу поблагодарить народ Армении, где мы жили некоторое вре
мя, за его доброту и гостеприимство. Особая благодарность Ереван
скому стартап-сообществу за помощь и поддержку. Грант Хачатрян
(Hrant Khachatrian), Завен Навоян (Zaven Navoyan), Арсен Егиазарян
(Arsen Yeghiazaryan), Андраник Хачатрян (Andranik Khachatryan),
Ашот Арзуманян (Ashot Arzumanyan), Аш Варданян (Ash Vardanian),
Адам Биттлингмайер (Adam Bittlingmayer), Артур Алексанян (Ar
tur Aleksanyan), Эрик Аракелян (Erik Arakelyan), Карен Гюльбудагян
(Karén Gyulbudaghyan) – огромное вам спасибо!
Спасибо организации Enterprise Armenia, национальному агент
ству по привлечению инвестиций Армении (National Investment
Promotion Agency of Armenia). Вы делаете огромную работу, и ваша
помощь бесценна.
Благодарю редакторов издательства Manning Патрика Барба (Pat
rick Barb), Бекки Уитни (Becky Whitney) и Франсис Лефковитц (Fran
ces Lefkowitz). Даже несмотря на то, что при работе над книгой сме
нились три редактора и в нее было внесено множество изменений,
каждый из вас внес свой вклад в ее ценность. Также благодарю Май
ка Стивенса (Mike Stephens) и Марьяна Баце (Marjan Bace), которые
с самого начала, с моих первых замыслов верили в то, что книга бу
дет написана.
Спасибо моему техническому редактору Нику МакГрейви (Nick
McGreivy), который, помимо того что является докторантом Прин
Благодарности
15
стонского университета, где он изучает физику плазмы, использует
JAX в своих исследованиях для оптимизации научных эксперимен
тов, а также для интеграции методов глубокого обучения в численное
моделирование. Также благодарю технического корректора Костаса
Пассадиса (Kostas Passadis) и рецензентов Арслана Габдулхакова (Ar
slan Gabdulkhakov), Чан Сун Пака (Chansung Park), Филлипа Дорнеля
(Fillipe Dornelas), Джеймса Блэка (James Black), Джеймса Ванга (James
Wang), Цзюнь Цзян (Jun Jiang), Кейт Ким (Keith Kim), Люсиана-Поля
Торье (Lucian-Paul Torje), Максима Волгина (Maxim Volgin), Наджиба
Арифа (Najeeb Arif), Ору Голану (Or Golan), Ритобрате Гошу (Ritobra
ta Ghosh), Сен Юн Ли (Seunghyun Lee), Симоне Де Бони (Simone De
Bonis), Стивену Оутсу (Stephen Oates), Тони Холдройду (Tony Hold
royd), Видья Винаю (Vidhya Vinay) и Войте Тума (Vojta Tuma). Вы все
предоставили множество полезных комментариев и предложений,
которые помогли улучшить эту книгу. С учетом вышесказанного, все
оставшиеся в книге ошибки целиком и полностью лежат на моей со
вести.
И наконец, благодарю моих друзей, экспертов GDE (Google De
veloper Experts) и компанию Google за поддержку столь крупномас
штабной инициативы. Сообщество GDE достойно наивысшей похва
лы! Многие эксперты GDE просматривали ранние версии и давали
полезную обратную связь. Особая благодарность Дэвиду Кардозо
(David Cardozo) за его весьма ценные отзывы и замечания.
О книге
Книга «Глубокое обучение с JAX» написана для того, чтобы помочь
читателям понять JAX и начать практически применять этот фрейм
ворк в проектах и исследованиях. В книге собрана вся наиболее важ
ная информация, которая позволит понять концепции JAX. Много
численные простые для понимания примеры облегчают восприятие
этой темы.
Для кого предназначена эта книга
Книга «Глубокое обучение с JAX» ориентирована на специалистовпрактиков и исследователей в области глубокого обучения, знако
мых с такими фреймворками, как PyTorch и TensorFlow, и желающих
начать использовать JAX. Исследователи в других областях (напри
мер, в физике или в сфере оптимизации) или аспиранты, специали
зирующиеся в глубоком обучении, численных методах оптимизации
или в распределенных вычислениях, также найдут эту книгу полез
ной для обучения и практической деятельности.
Как организована эта книга: общая схема
Книга состоит из трех частей, содержащих 12 глав.
В части I представлено введение и демонстрация возможностей
JAX:
глава 1 отвечает на самый важный вопрос: «Почему именно
JAX?» Здесь объясняется, что такое JAX, описываются его силь
ные и слабые стороны по сравнению с другими фреймворками,
такими как TensorFlow и PyTorch, а также особо отмечается, ког
да JAX может стать самым лучшим инструментом для вашего
проекта;
О книге
17
в главе 2 вы получите первый практический опыт использо
вания JAX. Мы создадим простую нейронную сеть для класси
фикации изображений, а кроме того, будут представлены ос
новные концепции: преобразования JAX для автоматической
векторизации, вычисления градиентов и JIT-компиляции. Вы
также узнаете, как сохранять и загружать модели, и поймете
различие между чистыми функциями и функциями с побоч
ным эффектом в JAX.
В части II рассматриваются основные функциональные средства
и возможности JAX:
в главе 3 используется «рабочая лошадка» глубокого обучения:
тензоры, или многомерные массивы. Сравниваются массивы
NumPy и JAX, обсуждается работа с ними на разнообразных
аппаратных устройствах, таких как центральные процессоры
(CPU), графические процессоры (GPU) и тензорные процессоры
(TPU), объясняются нюансы адаптации исходного кода с рас
смотрением различий при использовании NumPy и JAX;
в главе 4 рассматривается чрезвычайно важная задача вычис
ления градиентов, решение которой является необходимым
условием для тренировки нейронных сетей. Сравниваются
разнообразные методы дифференцирования, очень подробно
рассматриваются возможности автоматического дифференци
рования в JAX, а также использование режимов прямого и об
ратного автоматического дифференцирования;
в главе 5 показано, как оптимизировать код для повышения
производительности, используя JIT-компиляцию. Рассматри
ваются внутреннее устройство и работа механизма JIT и его
взаимодействие с компилятором XLA, а также способы устра
нения потенциальных ограничений;
в главе 6 представлена автоматическая векторизация, мощная
методика эффективной обработки пакетов данных. Рассматри
ваются разнообразные методы векторизации, объясняется, как
управлять JAX-преобразованием vmap(), а также анализируются
сценарии из реальной практики, где автоматическая реализа
ция показывает себя во всем блеске;
в главе 7 основное внимание уделено распараллеливанию, по
зволяющему одновременно выполнять вычисления на не
скольких устройствах. Объясняется, как использовать преобра
зование pmap() для параллельного выполнения, как управлять
его поведением. Рассматривается распараллеливание данных
для тренировки нейронной сети. Также используется код для
реальной работы на конфигурациях с несколькими хостами для
выполнения крупномасштабных задач;
в главе 8 представлено сегментирование тензоров, инноваци
онная эффективная методика распараллеливания в JAX. По
О книге
18
казано, как применять XLA для автоматического распаралле
ливания. Рассматривается реализация параллельного режима
работы с данными и тензорами для тренировки нейронной
сети, а также преимущества этой методики;
в главе 9 разбирается важная тема – генерация случайных чи
сел в JAX. Описываются различия между JAX и NumPy в этом
аспекте, обсуждается роль ключей в представлении состояния
генераторов случайных чисел, объясняется, как применить эти
концепции в реальных приложениях;
глава 10 знакомит читателей с pytrees, мощным инструмен
тальным средством для представления сложных структур дан
ных в JAX. Рассматриваются методы эффективной работы с py
trees, использование функций для обработки деревьев и даже
создание специализированных узлов pytree для особых потреб
ностей.
Часть III представляет богатую функциями и характеризующуюся
большим разнообразием экосистему библиотек, созданных на осно
ве JAX:
в главе 11 представлены библиотеки поддержки нейронных се
тей более высокого уровня, такие как Flax и Optax, предоставля
ющие удобные абстракции для создания и тренировки сложных
моделей. Мы будем использовать Flax для создания простого
MLP (многослойного перцептрона) и более продвинутой оста
точной нейросети для классификации изображений, а также
узнаем, как применять библиотеки Hugging Face для работы
с трансформерами и диффузионными моделями;
глава 12 дает более широкий обзор экосистемы JAX с описани
ем библиотек для решения разнообразных задач машинного
обучения, в том числе для тренировки больших языковых моде
лей (LLM), обучения с подкреплением и эволюционных вычис
лений. Также рассматриваются модули JAX для использования
в других областях науки, таких как физика, химия и т. д.
Если вы административный работник, то рекомендуется прочи
тать первые две главы, чтобы узнать о сильных сторонах JAX, о его
отличиях от PyTorch и TensorFlow и о том, как выглядит типичный
проект машинного обучения с применением JAX. Глава 12 также не
содержит технической информации и может дать представление
о том, где JAX проявляет себя с самой лучшей стороны.
Для разработчиков, желающих как можно быстрее приступить
к созданию нейронных сетей с помощью JAX, рекомендуется сосре
доточить основное внимание на главе 2, чтобы рассмотреть простой
пример глубокого обучения, на главах 3–6 для изучения основопо
лагающих концепций JAX и на главе 11 для обзора библиотек вы
сокого уровня в экосистеме. Остальную часть книги можно читать
О книге
19
в любом порядке в зависимости от ваших конкретных интересов.
Можно пропустить главы 7 и 8, если распараллеливание пока не вхо
дит в ваши планы, – вы можете вернуться к ним в любое время. Если
вас интересует генерация случайных чисел и использование pytrees,
то обратите особое внимание на главы 9 и 10, хотя предыдущие гла
вы предоставляют вполне достаточную базовую информацию, для
того чтобы начать практическое использование фреймворка JAX.
Исходный код примеров
Книга содержит множество примеров исходного кода – в пронуме
рованных листингах и в строках обычного текста. В обоих случа
ях исходный код выделяется моноширинным шрифтом, для того чтобы
можно было отличить его от обычного текста. Иногда фрагменты
или строки исходного кода дополнительно выделяются полужирным
шрифтом, чтобы выделить код, изменившийся по сравнению с пре
дыдущими примерами текущей главы, например при добавлении
новой функциональности в существующую строку кода.
Во многих случаях оригинальный исходный код был переформа
тирован – добавлялись разрывы строк и корректировалось выравни
вание для размещения в доступном пространстве страницы книги.
В некоторых случаях даже такие меры оказывались недостаточны
ми, поэтому в листинги включались маркеры продолжения текущей
строки (➥). Кроме того, комментарии часто удалялись из исходного
кода, когда код был подробно описан в тексте. Многие листинги со
провождаются примечаниями, особо выделяющими важные кон
цепции.
Готовые для выполнения фрагменты кода можно получить из (он
лайновой) версии liveBook этой книги: https://livebook.manning.com/
book/deep-learning-with-jax. Полные коды примеров из этой книги
доступны для загрузки с веб-сайта издательства Manning www.man
ning.com, а также из репозитория GitHub: https://github.com/che-shrcat/JAX-in-Action.
Почти для каждой главы существует соответствующий блокнот
Colab (или несколько блокнотов). Код протестирован для JAX вер
сии 0.4.14.
Отзывы и пожелания
Мы всегда рады отзывам наших читателей. Расскажите нам, что вы
думаете об этой книге – что понравилось или, может быть, не по
нравилось. Отзывы важны для нас, чтобы выпускать книги, которые
будут для вас максимально полезны.
Вы можете написать отзыв на нашем сайте www.dmkpress.com,
зайдяна страницу книги и оставив комментарий в разделе «Отзы
О книге
20
вы и рецензии». Также можно послать письмо главному редактору
по адресу dmkpress@gmail.com; при этом укажите название книги
в теме письма.
Если вы являетесь экспертом в какой-либо области и заинтересо
ваны в написании новой книги, заполните форму на нашем сайте
по адресу http://dmkpress.com/authors/publish_book/ или напишите
в издательство по адресу dmkpress@gmail.com.
Список опечаток
Хотя мы приняли все возможные меры для того, чтобы обеспечить
высокое качество наших текстов, ошибки все равно случаются.
Если вы найдете ошибку в одной из наших книг, мы будем очень
благодарны, если вы сообщите о ней главному редактору по адресу
dmkpress@gmail.com. Сделав это, вы избавите других читателей от
недопонимания и поможете нам улучшить последующие издания
этой книги.
Нарушение авторских прав
Пиратство в интернете по-прежнему остается насущной проблемой.
Издательства «ДМК Пресс» и Manning Publications очень серьезно от
носятся к вопросам защиты авторских прав и лицензирования. Если
вы столкнетесь в интернете с незаконной публикацией какой-либо
из наших книг, пожалуйста, пришлите нам ссылку на интернет-ре
сурс, чтобы мы могли применить санкции.
Ссылку на подозрительные материалы можно прислать по адресу
электронной почты dmkpress@gmail.com.
Мы высоко ценим любую помощь по защите наших авторов, бла
годаря которой мы можем предоставлять вам качественные мате
риалы.
Об авторе
Григорий Сапунов (Grigory Sapunov) – сооснователь
и технический директор (CTO) компании Intento. Григо
рий является инженером по разработке программного
обеспечения с более чем 20-летним опытом, имеет сте
пень к. т. н. (PhD) в области искусственного интеллекта,
а также входит в состав группы экспертов GDE (Google
Developer Expert) по машинному обучению.
Об иллюстрации
на обложке
На обложке книги «Глубокое обучение с JAX» изображена фигура
женщины, называемая «La Bearnaise» или «The Bearnese», что озна
чает человека из конкретной области Французских Пиренеев. Изо
бражение взято из коллекции Жака Грассе де Сэнт-Совёра (Jacques
Grasset de Saint-Sauveur), опубликованной в 1797 г. Иллюстрация де
тально проработана и раскрашена вручную.
В те дни можно было легко определить, где живут люди, какова
их профессия или положение в обществе, просто по их одежде. Из
дательство Manning поощряет изобретательность и инициативу
компьютерного бизнеса с помощью обложек книг, основанных на
богатом разнообразии региональной культуры многовековой дав
ности, возвращенных к жизни копиями иллюстраций из коллекций,
подобных этой.
Часть I
Первые шаги
О
тправляемся в путешествие в мир JAX, библиотеки, основан
ной на новейших технологиях, которая произвела революцию в глу
боком обучении и высокопроизводительных вычислениях. В этой
вводной части книги мы закладываем основу для понимания того,
почему JAX является одним из самых главных инструментов в по
стоянно развивающейся экосистеме фреймворков машинного
обучения. Прочитав две начальные главы, являющиеся основой для
дальнейшего обучения, вы получите представление об особенных
преимуществах JAX по сравнению с другими популярными библио
теками, такими как TensorFlow, PyTorch и NumPy, а также узнаете,
как использовать его мощь для конкретных проектов глубокого
обучения.
В главе 1 представлены основополагающие концепции и силь
ные стороны JAX. Она закладывает прочный фундамент, углубля
ясь в историю создания фреймворков глубокого обучения и особо
выделяя путь развития, который привел к рождению JAX. Эта гла
ва освещает нишу JAX в экосистеме, демонстрируя ее способность
компилировать и распараллеливать код для широкого спектра при
ложений, от моделирования океана до крупномасштабных нейрон
ных сетей. Сравнивая JAX с TensorFlow, PyTorch и NumPy, вы поймете
стратегические сценарии, требующие использования JAX, заклады
вая надежную основу для чтения остальной части книги.
Глава 2 представляет собой практическое обучение с элементами
исследования, которое проведет вас через процесс разработки прос
того приложения, использующего нейронную сеть. Эта глава – ваш
«ключ» к дальнейшему непосредственному изучению возможностей
JAX. Вы узнаете о высокоуровневой структуре проекта JAX – от за
24
Первые шаги
грузки наборов данных до создания нейронных сетей и откроете для
себя мощь функциональных преобразований JAX. Глава завершает
ся подробным описанием операций сохранения и загрузки моделей,
а также различий между чистыми функциями и функциями с побоч
ными эффектами. Эта информация вооружит вас знаниями для соз
дания более сложных проектов.
В совокупности эти главы представляют собой своеобразный
подготовительный этап для вступления во впечатляющий мир JAX,
обеспечивая хорошую подготовку к решению более сложных задач
в области глубокого обучения и за ее пределами. Независимо от того,
являетесь ли вы опытным разработчиком или новичком в этой об
ласти, часть I этой книги предлагает всеобъемлющее руководство,
позволяющее продвигаться дальше с уверенностью и ясным пони
манием.
1
Когда и зачем
используется JAX
Темы главы:
введение в JAX;
когда и где следует использовать JAX;
сравнение JAX с TensorFlow, PyTorch и NumPy.
Еще одна библиотека глубокого обучения? Вы серьезно?! После того
как все пути сошлись в одну точку, которой стал всеобщий любимец
PyTorch, и сформировалась устойчивая экосистема вокруг Tensor
Flow, зачем мне беспокоиться о JAX? А если бы мне потребовалось
низкоуровневое средство разработки нейронной сети, то есть ста
рый добрый NumPy или его альтернативы с поддержкой GPU. По
чему я вообще должен обращать внимание на JAX?
История развития фреймворков глубокого обучения показывает,
что ни один фреймворк не существует вечно. Например, где теперь
Theano, который оказал значительное влияние на развитие этой
отрасли? Для меня это была первая библиотека глубокого обуче
ния в современном смысле после давно забытых PyBrain2 и про
чих. А где сейчас Caffe? Он был весьма популярен много лет назад,
особенно для развертывания промышленных решений. Много лет
назад я разработал инструмент помощника водителя для распозна
26
Глава 1
Когда и зачем используется JAX
вания дорожных знаков в реальном времени, работающий на ста
рых смартфонах с ОС Android, обладающих гораздо меньшей мощ
ностью, чем современные телефоны, и Caffe был лучшим вариантом
выбора для выполнения этой работы в то время.
А что можно сказать о старом добром Torch7, который многие
модели переноса стилей изображений использовали приблизи
тельно в 2015 году? Я тогда участвовал в одном таком проекте, и мы
использовали многие из этих моделей во внутренних (серверных)
компонентах. Кстати, а где сейчас Chainer, TensorFlow 1, CNTK, Caffe2
и многие другие?
Это были замечательные фреймворки; ими перестали пользо
ваться вовсе не потому, что они были плохими. Просто время летит
и требования меняются. Но эти фреймворки (или библиотеки; гра
ница иногда размыта) продолжают жить в своих потомках. Благода
ря Torch7 и Caffe2 у нас теперь есть PyTorch. Благодаря TensorFlow 1
мы получили Keras, который сейчас является API высокого уровня
по умолчанию для TensorFlow 2. Спасибо Chainer, DyNet и прочим за
то, что в нашем распоряжении теперь имеются динамические вы
числительные графы и немедленное выполнение.
Если вы предполагаете, что PyTorch или TensorFlow 2 – это «венец
эволюции», то подумайте еще раз. В этой области уже есть новый
персонаж: JAX и его богатая экосистема.
JAX – это библиотека, написанная на языке Python и разработан
ная компанией Google для высокопроизводительных вычислений
с использованием массивов. Некоторые называют JAX «NumPy на
стероидах», потому что эта библиотека используется для крупно
масштабных высокопроизводительных вычислений благодаря
своей способности компилировать и распараллеливать код не
зависимо от того, написан ли он для задач глубокого обучения
или, например, для имитации поведения океана (https://dionhae
fner.github.io/2021/12/supercharged-high-resolution-ocean-simu
lation-with-jax/), для прогнозирования погоды (https://arxiv.org/
abs/2212.12794) или для космологических исследований (https://
arxiv.org/abs/2302.05163; https://arxiv.org/abs/2305.06347). Многие
весьма интересные приложения и библиотеки на основе JAX созда
ны для физики, включая молекулярную динамику, гидрогазодина
мику, имитации абсолютно твердого тела, квантовые вычисления,
астрофизику, моделирование (поведения) океана и многое другое.
Существуют библиотеки для распределенного разложения (фак
торизации) матриц на множители, потоковой обработки данных,
сворачивания (укладки) белка (protein folding) и химического мо
делирования, а кроме того, постоянно появляются и другие новые
приложения.
Не вызывает удивления тот факт, что JAX широко применяется
в приложениях глубокого обучения. В действительности JAX часто
Когда и зачем используется JAX
27
определяется как фреймворк глубокого обучения, третий после Py
Torch и TensorFlow.
В общеизвестном отчете «State of AI 2021» (https://mng.bz/aEYx)
JAX назван новым сильным фреймворком-конкурентом, и действи
тельно – все больше и больше исследований в области глубокого
изучения управляются с помощью JAX. Самыми свежими приме
рами исследований являются весьма важные материалы компании
Google о применении трансформеров и современных версий много
слойных перцептронов (MLP – multilayer perceptron) для задач ма
шинного зрения (восприятия) (а именно Vision Transformer [ViT]
и MLP-Mixer; https://github.com/google-research/vision_transformer).
Компания DeepMind объявила об использовании JAX для ускорения
исследований, а кроме того, эта компания вносит существенный
вклад в развитие экосистемы JAX (https://deepmind.google/discover/
blog/using-jax-to-accelerate-our-research/). Рассмотрим еще несколь
ко областей, в которых JAX в настоящее время набирает обороты.
Поскольку JAX обеспечил высокую скорость проведения экспери
ментов с новыми алгоритмами и архитектурами, теперь он стано
вится основой многих недавних публикаций DeepMind. Среди них
я бы выделил следующие:
новая методика обучения с самоконтролем (self-supervised
learning – SSL) под названием BYOL (Bootstrap Your Own Latent;
https://arxiv.org/abs/2006.07733);
обобщенная архитектура на основе трансформеров для струк
турированных входных и выходных данных под названием
Perceiver IO (https://deepmind.google/discover/blog/building-ar
chitectures-that-can-handle-the-worlds-data/);
исследование больших языковых моделей (LLM – large language
model) на модели Gopher с 280 млрд параметров и на модели
Chinchilla с 70 млрд параметров (https://deepmind.google/dis
cover/blog/an-empirical-analysis-of-compute-optimal-large-lan
guage-model-training/).
Одна из первых больших GPT-подобных моделей (GPT – Genera
tive Pretrained Transformer – предобученный генеративный транс
формер) с открытым исходным кодом – трансформерная модель
(естественного) языка с 6 млрд параметров GPT-J-6B исследователь
ской группы EleutherAI тренировалась с использованием JAX в среде
Google Cloud. Авторы модели GPT-J утверждают, что это был наибо
лее подходящий комплект инструментальных средств для быстрой
разработки крупномасштабных моделей. Внутри компании Google
JAX используется для тренировок собственных больших языковых
моделей, таких как PaLM 2, Gemini и Gemma. В 2023 г. компания xAI
(https://x.ai/) начала тренировку своей модели Grok LLM (https://
github.com/xai-org/grok-1), используя специализированный стек
тренировки и инференса на основе Kubernetes, Rust и JAX.
28
Глава 1
Когда и зачем используется JAX
Кроме того, в 2021 г. компания Hugging Face сделала связку JAX
и Flax (высокоуровневую библиотеку поддержки нейронных сетей
на основе JAX) третьим официально поддерживаемым фреймвор
ком в своей широко известной библиотеке Transformers. Комплект
Hugging Face предварительно натренированных моделей по состоя
нию на июнь 2024 г. содержал количество моделей JAX, сравнимое
с количеством моделей TensorFlow. PyTorch пока опережает оба
фреймворка, но уже начался процесс переноса моделей из PyTorch
в JAX/Flax.
Возможно, сам по себе JAX не вполне подходит для развертыва
ния в производственной среде, поскольку в основном сконцентри
рован на решении научно-исследовательских проблем, но именно
так и происходило и с PyTorch. Вероятнее всего, разрыв между сфе
рой научных исследований и производством скоро исчезнет. Уже
существует пакет JAX2TF, обеспечивающий поддержку собственной
сериализации JAX и взаимодействие между JAX и TensorFlow. Мо
дель JAX можно преобразовать в Saved Model TensorFlow, чтобы по
явилась возможность работы с ней с помощью TensorFlow Serving,
TFLite или TensorFlow.js.
Учитывая авторитет Google и быстрое расширение сообщества,
я считаю, что JAX ожидает блестящее будущее. Этот фреймворк легко
внедрять, поскольку Python и NumPy широко используются и знако
мы большинству разработчиков. Его компонуемые преобразования
функций помогают поддерживать исследования в области машин
ного обучения. JAX имеет весьма жизнеспособную экосистему, вы
сокую степень выразительности и может обеспечить высокую про
изводительность и вычислительную эффективность. Он использует
общеизвестный API NumPy и продвигает методику функционально
го программирования. Код с применением JAX может эффективно
работать на нескольких внутренних аппаратных компонентах си
стемы, включая центральные, графические (GPU) и тензорные про
цессоры (TPU). А если вы разрабатываете LLM или обучаете другие
крупномасштабные нейронные сети, то JAX является правильным
вариантом выбора для этих целей. (Подробнее об LLM читайте в гла
ве 12, где я предоставляю вам список модулей JAX для высокопроиз
водительного обучения LLM.)
Почему выбрано название JAX
В документе, впервые представляющем JAX, «Compiling Machine Learning Programs Via High-Level Tracing» (https://mlsys.org/Conferences/
doc/2018/146.pdf), объясняется, что аббревиатура JAX означает «Just
After eXecution» (сразу после выполнения), поскольку для компиляции
функции мы сначала отслеживаем ее выполнение в среде Python.
Когда и зачем используется JAX
29
Прежде чем перейти к изучению конкретных способов использо
вания JAX, позвольте мне внести некоторое уточнение о содержании
этой книги. Хотя существует множество интереснейших областей,
в которых можно использовать JAX (включая космологию, укладку
(сворачивание) белков и моделирование (поведения) океана), эта
книга в основном посвящена глубокому обучению. Она ориентиро
вана на разработчиков, инженеров и исследователей, которые уже
понимают концепции глубокого обучения и имеют некоторый опыт
работы с одним из других фреймворков. Если вы новичок в этой об
ласти, то все равно можете использовать книгу для изучения JAX, но
вам понадобится отдельная книга, в которой описаны основы глу
бокого обучения.
Хотя многие примеры адаптированы для использования в про
цессе глубокого обучения, читатель, заинтересованный в приме
нении JAX для других целей, также может изучить этот фреймворк
в восьми главах, в которых подробно описано ядро JAX в части II.
Я добавил примечания, объясняющие концепции глубокого обуче
ния простым языком для таких читателей. Да простят меня экспер
ты по глубокому обучению, для которых эти дополнительные при
мечания, вероятно, покажутся слишком тривиальными, но я считаю,
что они могут оказаться чрезвычайно полезными для читателей, за
нятых в других сферах деятельности и желающих изучить JAX. Если
у вас есть опыт работы в области глубокого обучения, то просто про
пускайте все примечания в духе «Глубокое обучение 101».
Я преднамеренно не включил в книгу варианты, не относящиеся
к теме глубокого обучения, например астрономию. Хотя есть превос
ходные примеры использования JAX в других областях, я считаю, что
не стоит включать такие варианты в книгу. Они не принесут особой
пользы людям, занимающимся глубоким обучением, но и не ока
жут большой помощи людям, занимающимся физикой, поскольку
им, вероятнее всего, потребуется гораздо больше информации, ско
рее всего, отдельная книга или семинар по этой теме. Хотелось бы
полностью сосредоточиться на основах JAX, которые полезны везде,
и помочь людям, занимающимся глубоким обучением, попробовать
новый фреймворк и начать его применять на практике. Тем не ме
нее я многократно подчеркиваю применимость JAX в других сферах
деятельности и даю некоторые соответствующие ссылки (особенно
в последней главе об экосистеме).
Книга формирует для читателя базовое понимание JAX с самого
начала, глубоко погружаясь во все основные концепции, каждая из
которых объясняется с помощью многочисленных примеров и не
больших проектов. Иногда это может показаться довольно прими
тивным, но именно такой подход дает вам полное понимание того,
как все работает, и обеспечивает прочную основу для изучения бо
гатой экосистемы JAX. Когда вы начнете свое захватывающее путе
30
Глава 1
Когда и зачем используется JAX
шествие в мир JAX, пусть книга «Глубокое обучение с JAX» станет
вашим незаменимым проводником.
В этой главе я представлю JAX, его динамично развивающую
ся экосистему и объясню, что такое JAX и как он сравним с NumPy,
PyTorch и TensorFlow. Мы рассмотрим сильные стороны JAX, чтобы
понять, как их объединить, дабы получить мощный инструмент для
исследований в области глубокого обучения и высокопроизводи
тельных вычислений.
1.1
Почему нужно использовать JAX
Есть несколько причин, по которым вам может потребоваться ис
пользование JAX. Во-первых, это его выдающаяся производитель
ность и способность выполнять вычисления на нескольких аппа
ратных устройствах, включая CPU, GPU и TPU, применяя при этом
хорошо знакомый API в стиле NumPy для простоты освоения ис
следователями и инженерами. Во-вторых, это его функциональ
ная природа и набор компонуемых преобразований функций для
компиляции, векторизации, автоматического дифференцирования
и распараллеливания. В-третьих, это компонуемость его элементов
и богатая экосистема модулей и библиотек, созданных на основе JAX.
Прежде чем мы перейдем к более подробному объяснению вы
шеперечисленных трех причин использования JAX, важно отме
тить, что JAX не всегда является лучшим вариантом для решения
каждой задачи. Если вы работаете с весьма специфическими кон
фигурациями, такими как среда развертывания для мобильных или
встроенных устройств, то сохранение текущего фреймворка, такого
как TFLite, вероятнее всего, будет наилучшим вариантом (хотя при
необходимости можно воспользоваться конвертером JAX2TF). Если
уже существует обширная и стабильная кодовая база с использо
ванием текущего фреймворка, например PyTorch, то в переходе на
JAX, возможно, нет особого смысла (хотя если скорость очень важна,
такой переход все-таки может оказаться правильным вариантом).
Теперь обсудим преимущества JAX.
1.1.1
Производительность вычислений
В первую очередь следует отметить, что JAX обеспечивает превос
ходную производительность вычислений. Некоторые эталонные
тесты показывают, что JAX быстрее TensorFlow (https://dzone.com/
articles/accelerated-automatic-differentiation-with-jax-how). Другие
подтверждают, что «производительность JAX вполне конкуренто
способна на GPU и CPU. Он неизменно входит в число лучших реа
лизаций на обеих платформах» (https://github.com/dionhaefner/pyh
Почему нужно использовать JAX
31
pc-benchmarks). Даже самый быстрый в мире трансформер на 2020 г.
был создан с использованием JAX (https://cloud.google.com/blog/
products/ai-machine-learning/google-breaks-ai-performance-recordsin-mlperf-with-worlds-fastest-training-supercomputer).
В этом разделе можно привести множество фактов, включая воз
можность использования современного аппаратного оборудова
ния, такого как GPU или TPU (и JAX, вероятно, является наилучшим
выбором, если необходимо достичь максимальной производитель
ности на TPU), компиляция just-in-time (JIT) с помощью XLA (Ac
celerated Linear Algebra), автоматическая векторизация, простое
распараллеливание в кластере и возможность использования па
раллельного обучения данных и моделей при масштабировании
программы с процессора ноутбука до самого крупного тензорного
процессора TPU Pod в облаке. Мы обсудим все эти темы в различных
главах книги.
JAX можно использовать как ускоренный NumPy, если заменить
инструкцию import numpy as np на import jax.numpy as np в самом на
чале программы. В определенном смысле это аналог переключе
ния с NumPy на CuPy (для использования графических процессоров
с NVIDIA CUDA или платформ AMD ROCm), Numba (для обеспечения
поддержки JIT и GPU) или даже PyTorch, если требуется ускорение
аппаратных устройств для выполнения операций линейной алгеб
ры. Не все функции NumPy реализованы в JAX, поэтому иногда по
требуется нечто большее, чем простая замена инструкции импорта.
Эта тема рассматривается в главе 3.
Аппаратное ускорение с помощью GPU или TPU может повысить
скорость умножения матриц и других операций, которые могут по
лучить преимущество от работы на таком массово-параллельном
оборудовании. Чтобы начать использовать этот тип ускорения, не
обходимо только лишь выполнять вычисления с многомерными
массивами, размещенными в памяти ускорителя. В главе 3 показа
но, как управлять размещением данных.
Ускорение также может являться результатом JIT-компиляции
с помощью компилятора XLA, который оптимизирует граф вычис
лений и способен объединять последовательность операций в одну
эффективную вычислительную операцию или исключать некоторые
избыточные вычисления. Такой подход повышает производитель
ность даже на обычном процессоре без какого-либо другого аппарат
ного ускорения (хотя процессоры отличаются друг от друга и многие
современные ЦП предоставляют специальные инструкции, вполне
подходящие для приложений глубокого обучения).
На рис. 1.1 показан снимок экрана с кодом простой (и доволь
но-таки бесполезной) функции, содержащей некоторое количество
вычислительных операций, которые выполняются с применением
только NumPy, средствами JAX, использующими CPU и GPU (в моем
32
Глава 1
Когда и зачем используется JAX
варианте Tesla-P100), а также в JAX-версиях с компиляцией той же
функции для устройств CPU и GPU. Все подробности и особенности,
связанные с этой темой, будут описаны в главе 5. Соответствующий
код можно найти в репозитории этой книги: https://github.com/cheshr-cat/JAX-in-Action/blob/main/Chapter-1/JAX_in_Action_Chapter_1_
JAX_speedup.ipynb.
Рис. 1.1 Простая (и довольно-таки бесполезная) функция, содержащая
некоторое количество вычислительных операций, которые выполняются
с применением только NumPy, средствами JAX, использующими CPU и GPU,
а также в JAX-версиях с компиляцией той же функции для устройств CPU и GPU.
Мы определим время выполнения вычислений в функции f(x)
На рис. 1.2 сравниваются различные способы выполнения вычис
лений в функции f(x). Сейчас не обращайте внимания на некоторые
пока непонятные функции, такие как block_until_ready() или jax.
device_put(). Первая функция необходима, потому что JAX исполь
зует асинхронное диспетчирование и не обязан ждать завершения
вычислений. В этом случае измерение времени может оказаться
неправильным, так как учитывает только некоторые вычисления.
Вторая функция требуется для связывания массива с конкретным
устройством. Более подробно мы обсудим эти функции в главе 3.
В рассматриваемом здесь примере скомпилированная версия JAX
для CPU почти в пять раз быстрее исходного варианта с примене
нием только NumPy, хотя версия JAX для CPU без компиляции вы
полняется немного медленнее. Нет необходимости в компиляции
этой функции для использования GPU, так как по умолчанию все
Почему нужно использовать JAX
33
массивы создаются на первом устройстве GPU/TPU, если оно до
ступно. JAX-функция без компиляции, продолжающая использовать
GPU, в 5,6 раза быстрее, чем скомпилированная версия JAX для CPU.
А скомпилированная версия для GPU той же функции еще в 2,9 раза
быстрее версии без компиляции. Если сравнивать с исходной функ
цией, применяющей только NumPy, то получаем общее ускорение
почти в 77 раз – весьма значительное улучшение скорости без како
го-либо существенного изменения кода функции.
Рис. 1.2 Измерения времени для различных способов вычисления в функции
f(x). Строка 1 – реализация с применением только NumPy; строка 2 –
использование JAX на CPU; строка 3 – JIT-скомпилированная версия JAX на ЦП;
строка 4 – использование JAX на GPU; строка 5 – JIT-скомпилированная версия
JAX на GPU
Еще одно средство, способное повысить скорость выполнения, –
автоматическая векторизация, преобразующая функцию с одним
элементом в функцию, которая может обрабатывать пакет элемен
тов. Главная цель такого подхода – облегчение деятельности разра
ботчика и ускорение написания эффективного кода с применением
векторизации. К тому же появляется возможность предоставления
альтернативного способа ускорения вычислений, если имеющиеся
аппаратные ресурсы и логика программы позволяют одновременно
выполнять вычисления с несколькими (многими) элементами. Ав
томатическая векторизация рассматривается в главе 6.
Наконец, можно распараллелить код в кластере и выполнять
крупномасштабные вычисления в распределенном стиле, что не
возможно при использовании только NumPy, но может быть реали
зовано с помощью Dask, DistArray, Legate NumPy или других средств.
Тренировка крупномасштабных моделей (таких как GPT-подобных
34
Глава 1
Когда и зачем используется JAX
LLM или диффузионных моделей преобразования текста в изобра
жение) – это весьма актуальная тема. Возможность сделать это эф
фективно может кардинально изменить ситуацию, сэкономив часы
и дни вычислений, следовательно, и немалые денежные средства.
В главах 7 и 8 рассматриваются разнообразные аспекты распарал
леливания.
И конечно же, вы можете извлечь выгоду из всех вышеупомяну
тых вещей одновременно, что и происходит при обучении крупных
распределенных нейронных сетей, создании масштабных физиче
ских имитаций, выполнении распределенных эволюционных вы
числений и т. д.
1.1.2
Функциональный подход
В JAX все открыто и все в явном виде. Благодаря функционально
му подходу отсутствуют скрытые переменные и побочные эффек
ты (подробнее о побочных эффектах здесь: https://ericnormand.me/
podcast/what-are-side-effects). Код понятен; вы можете изменить
все, что захотите, а сделать что-то нестандартное гораздо проще.
Как я уже отмечал ранее, научным работникам очень нравится JAX,
и с ним проводится много новых исследований.
Такой подход требует от вас изменения некоторых привычек.
В PyTorch и TensorFlow код обычно организован в классы. Конкрет
ная нейронная сеть – это класс, все параметры которого являются ее
внутренним состоянием. Оптимизатор – это другой класс со своим
внутренним состоянием. Если вы работаете по направлению обуче
ния с подкреплением, то используемая среда – это обычно еще один
класс с собственным состоянием. Подобное приложение глубокого
обучения выглядит как объектно ориентированная программа с эк
земплярами классов и вызовами методов.
А вот в JAX код организован в форме функций. Тем не менее не
правильным является утверждение о том, что JAX не разрешает ис
пользование классов; все высокоуровневые библиотеки предостав
ляют пользователю абстракции классов по аналогии с обычными
фреймворками глубокого обучения. Различие заключается в том, что
классы и функции не используют какое-либо внутреннее состояние.
Требуется передача любого внутреннего состояния как параметра
функции, так что все параметры модели (набор весовых коэффи
циентов для нейронных сетей или даже начальные числа (seeds)
для генерации случайных чисел) передаются напрямую в функции.
В этом случае функции генерации случайных чисел требуют явно
го предоставления состояния генератора СЧ (генераторы СЧ в JAX
подробно рассматриваются в главе 9). Градиенты вычисляются явно
с помощью вызова специальной функции (полученной в результа
те преобразования grad() конкретной рассматриваемой функции –
Почему нужно использовать JAX
35
тема главы 4). Состояние оптимизатора и вычисленные градиенты
также являются параметрами функции оптимизатора и т. д. Нет ни
каких скрытых состояний, все можно увидеть. И побочные эффекты
недопустимы, потому что они не работают при JIT-компиляции.
JAX заставляет изменить привычное мышление в отношении опе
раторов if и циклов for из-за способа их компиляции. Эта тема об
суждается в главе 5.
Кроме того, функциональный подход создает богатые возмож
ности композиционности. JAX предоставляет мощные компонуе
мые трансформации функций (более подробно об этом в под
разделе 1.2.2), такие как автоматическое дифференцирование,
автоматическая векторизация, сквозная компиляция с помощью
XLA и распараллеливание, которые с легкостью объединяются. Бога
тые возможности композиционности и выразительность приводят
к формированию мощной экосистемы.
1.1.3
Экосистема JAX
JAX обеспечивает прочную основу для создания нейронных сетей, но
источником его истинной мощи является постоянно развивающая
ся экосистема. В совокупности со своей экосистемой JAX предостав
ляет весьма перспективную альтернативу двум самым высокотех
нологичным в настоящее время фреймворкам глубокого обучения:
PyTorch и TensorFlow.
Используя JAX, можно с легкостью объединять решения, получен
ные с помощью различных модулей. Нет необходимости в работе
с функционально монолитным фреймворком, таким как TensorFlow
или PyTorch, в который «все включено». Это мощные фреймворки,
но замена одного плотно интегрированного компонента на другой
иногда оказывается трудным делом. В JAX вы можете создавать спе
циализированные решения, комбинируя необходимые блоки в раз
ных сочетаниях.
Задолго до появления JAX среда глубокого обучения напоминала
конструктор Лего – пользователь получал в свое распоряжение на
бор блоков разнообразных форм и цветов для создания собственных
решений. Глубокое обучение – это действительно среда в стиле Лего
с разными уровнями: на нижних располагаются активационные
функции и типы слоев (нейросети), на более высоких – конструк
тивные блоки архитектурных примитивов (таких как уровень само
внимания), оптимизаторов, токенизаторов и т. п.
Применяя JAX, вы получаете еще бóльшую свободу для объедине
ния разнообразных блоков, наилучшим образом соответствующих
вашим потребностям. Например, можно взять нужную высокоуров
невую библиотеку поддержки нейронных сетей и отдельный модуль,
реализующий оптимизаторы, которые вы намерены использовать.
Глава 1
36
Когда и зачем используется JAX
Далее можно выполнить собственную настройку оптимизатора для
получения специализированных норм (интенсивности) обучения
для различных уровней (слоев), воспользоваться предпочитаемыми
загрузчиками данных PyTorch, применить отдельную библиотеку
поддержки обучения с подкреплением, добавить дерево поиска Mon
te-Carlo Tree Search, а также использовать некоторые другие компо
ненты для оптимизаторов метаобучения из другой библиотеки.
Комбинирование решений с использованием JAX и его экосисте
мы напоминает структуру Unix-подобных систем, основанную на
простом модульном проектном решении. Эрик Рэймонд (Eric Ray
mond) сказал: «Пишите программы, которые выполняют только одну
задачу и делают это хорошо. Пишите программы, которые работают
совместно. Пишите программы, которые обрабатывают текстовые
потоки, потому что это универсальный интерфейс» («Basics of the
Unix Philosophy», Addison-Wesley Professional, 2003 г.). Все сказанное
выше, кроме текстовых потоков, относится и к философии JAX. Но
и обмен текстовыми данными когда-нибудь также может стать уни
версальным интерфейсом между различными (крупными) нейрон
ными сетями. Кто знает? И как раз здесь возникает экосистема.
Экосистема JAX уже имеет огромный размер, и это только начало.
Она содержит превосходные модули, перечисленные ниже:
для высокоуровневого программирования нейронных сетей, из
множества которых можно особо выделить следующие:
– Flax компании Google (https://github.com/google/fax);
– Equinox с простым в использовании синтаксисом в стиле Py
Torch (https://github.com/patrick-kidger/equinox);
– Keras 3, новейший модуль с поддержкой многих внутрен
них компонентов/устройств (https://github.com/keras-team/
keras);
модуль с самыми современными оптимизаторами Optax
(https://github.com/deepmind/optax);
библиотеки поддержки обучения с подкреплением:
– RLax компании DeepMind (https://github.com/deepmind/rlax);
– Coax группы Microsoft Research (https://github.com/coax-dev/
coax);
библиотека поддержки графов нейронных сетей Jraph (https://
github.com/deepmind/jraph);
библиотека поддержки молекулярной динамики JAX, M.D.
(https://github.com/google/jax-md).
Следует учесть, что это далеко не полный список всех модулей, он
приведен, чтобы дать читателям общее представление об экосисте
ме JAX.
Экосистема JAX уже содержит сотни модулей и быстро развивает
ся. Постоянно появляются новые библиотеки. Совсем недавно я об
ратил внимание на некоторые из них:
Чем JAX отличается от NumPy
37
EvoJAX (https://github.com/google/evojax) и evosax (https://git
hub.com/RobertTLange/evosax) для эволюционных вычислений;
FedJAX для федеративного обучения (https://github.com/google/
fedjax);
Paxml (https://github.com/google/paxml) и MaxText (https://git
hub.com/google/maxtext) для тренировки больших языковых
моделей (LLM);
библиотека Scenic для исследований в области компьютерного
зрения (https://github.com/google-research/scenic).
Компания DeepMind уже разработала комплект библиотек на
основе JAX (некоторые из них были упомянуты выше). Новые биб
лиотеки доступны для Monte Carlo Tree Search, для верификации
нейронных сетей, для обработки изображений и для многих других
задач. Но это далеко не полная картина. Я не собираюсь перечислять
все модули и библиотеки в этой главе, потому что это невозможно
и не является целью. Первая глава не предназначалась для публи
кации абсолютно полного списка всего, что имеется в экосистеме.
Существует множество широко известных библиотек, кратко опи
санных в главе 12. Но все же существует немалая вероятность того,
что после выхода книги из печати появятся новые полезные библио
теки, которые будут использоваться повсеместно. Если вы понимае
те сущность JAX, то с легкостью начнете использовать эти модули.
Количество моделей JAX в репозитории Hugging Face постоянно
увеличивается, так как JAX входит в тройку фреймворков, поддер
живаемых этой компанией. Изучение данной темы на практике яв
ляется содержимым главы 11.
Итак, JAX набирает обороты, и его экосистема постоянно расши
ряется. Это весьма подходящее время, чтобы присоединиться.
1.2
Чем JAX отличается от NumPy
Строго говоря, JAX – это математическая библиотека на языке Py
thon с интерфейсом NumPy, разработанная компанией Google (точ
нее, группой Google Brain). Библиотека интенсивно используется
для исследований в сфере машинного обучения, но не ограничена
этой областью, поэтому с помощью JAX можно решать многие дру
гие задачи.
Создатели JAX описывают этот фреймворк как объединение Au
tograd и XLA для высокопроизводительных вычислений. Не стоит
беспокоиться, если вам незнакомы эти названия, вы могли и не слы
шать о них, особенно если вы новичок в этой области.
Autograd (https://github.com/hips/autograd) – это библиотека, по
зволяющая эффективно вычислять производные в коде NumPy
и фактически являющаяся предшественником JAX. Кстати, главные
38
Глава 1
Когда и зачем используется JAX
разработчики библиотеки Autograd теперь работают в группе JAX.
Короче говоря, само название Autograd означает, что вы можете ав
томатически вычислять градиенты в своих расчетах, что является
главным смыслом, квинтэссенцией глубокого обучения и многих
других областей науки, включая численную оптимизацию, физиче
ские имитации, а также, в более общем плане, дифференцируемое
программирование.
XLA (Accelerated Linear Algebra) – это специализированный для
конкретной предметной области компилятор для линейной алгеб
ры, созданный компанией Google. Он компилирует функции Python,
содержащие операции линейной алгебры, в высокопроизводитель
ный код для выполнения на графических (GPU) или тензорных (TPU)
процессорах. Начнем с NumPy.
1.2.1
JAX как NumPy
NumPy – это «рабочая лошадка» для выполнения числовых расчетов
на Python. Этот фреймворк настолько широко используется в про
изводстве и науке, что NumPy API стал стандартом де-факто для ра
боты с многомерными массивами в Python. JAX предоставляет со
вместимый с NumPy API, но при этом предлагает множество новых
функциональных возможностей, отсутствующих в NumPy. Именно
поэтому некоторые люди называют JAX «NumPy на стероидах».
JAX предоставляет структуру данных многомерного массива
с именем Array, которая реализует многие типовые свойства и ме
тоды массива numpy.ndarray. Также существует пакет jax.numpy, реа
лизующий NumPy API с множеством хорошо известных функций,
таких как abs(), conv(), exp() и т. п.
JAX пытается соответствовать NumPy API в как можно большей
степени, и во многих случаях можно переключаться с numpy на jax.
numpy без изменений в программе. Но при этом остаются действую
щими некоторые ограничения: не весь код NumPy можно исполь
зовать в JAX. JAX активно продвигает парадигму функционального
программирования и требует использования чистых функций без
побочных эффектов. Поэтому массивы JAX являются неизменяе
мыми, хотя программы с применением NumPy часто используют
прямое обновление «на месте», например arr[i] += 10. JAX предо
ставляет чисто функциональный альтернативный API, заменяющий
прямые обновления «на месте» на функцию индексированного об
новления. Для приведенного выше примера замена будет выглядеть
так: arr = arr.at[i].add(10). Существует еще несколько различий,
которые мы рассмотрим в главе 3.
Таким образом, вы можете использовать практически всю мощь
библиотеки NumPy и продолжать писать программы так, как при
выкли. Но при этом у вас появляются новые возможности.
Чем JAX отличается от NumPy
1.2.2
39
Компонуемые трансформации
JAX намного больше, чем просто NumPy. Он предоставляет набор
компонуемых трансформаций функций (composable function transformations) для кода Python + NumPy. По своей сущности JAX являет
ся расширяемой системой для трансформаций числовых функций
с четырьмя основными операциями трансформации (но это не ис
ключает возможности появления новых операций трансформации):
вычисление градиента в коде или дифференцирование. Это
основной смысл глубокого обучения. JAX использует методику
под названием автоматическое дифференцирование (automatic
differentiation, сокращенно: autodiff). Автоматическое диффе
ренцирование позволяет сосредоточиться на написании кода,
а не тратить время на вычисление производных вручную, об
этом позаботится фреймворк. Обычно эта операция выполня
ется с помощью функции grad(), но существуют и другие, более
продвинутые варианты. Контекст автоматического дифферен
цирования и прочие подробности этой операции рассматрива
ются в главе 4;
компиляция кода с помощью функции jit(), или JIT-компиля
ция. Используется XLA компании Google для компиляции и ге
нерации эффективного кода для GPU (обычно для графических
процессоров NVIDIA через CUDA, но также разрабатывается
поддержка платформы AMD ROCm) и TPU (Tensor Processing
Units компании Google). XLA представляет собой внутренний
компонент, увеличивающий мощь фреймворков машинного
обучения, изначально TensorFlow, при использовании на раз
нообразных устройствах, включая CPU, GPU и TPU. Этой теме
посвящена вся глава 5;
автоматическая векторизация кода с помощью vmap(), т. е.
функции векторизирующего отображения. Вероятно, вам из
вестно, что такое отображение (map), если вы знакомы с функ
циональным программированием. Но даже если незнакомы,
не беспокойтесь, немного позже я объясню, что это означает.
vmap() принимает на себя обязанности по управлению набо
рами измерений используемых массивов и может с легкостью
преобразовать код обработки одного элемента данных в код
обработки многочисленных элементов (называемых пакетом
(batch)) за один прием. Такой подход также можно назвать авто
матическим пакетированием. Выполняя эту операцию, вы век
торизуете вычисления, что, как правило, позволяет добиться
существенного ускорения на современном аппаратном обору
довании, способном эффективно распараллеливать матричные
вычисления. Автоматическая векторизация рассматривается
в главе 6;
Глава 1
40
Когда и зачем используется JAX
распараллеливание кода для выполнения на нескольких уско
ряющих устройствах, например на GPU или TPU. Это делается
с помощью pmap(), функции, помогающей создавать SIMD-прог
раммы (SIMD – single program, multiple data – одна программа
(или инструкция), множественный поток данных). pmap() ком
пилирует функцию с использованием XLA, затем реплициру
ет ее (размножает ее копии) и выполняет каждую реплику на
собственном XLA-устройстве в параллельном режиме. Эта тема
обсуждается в главах 7 и 8.
Каждая операция трансформации принимает некоторую функ
цию и возвращает другую функцию, так что вы можете взять лю
бую функцию и автоматически получить ее скомпилированную или
векторизованную версию, при этом не требуется писать много кода.
Все ваше внимание сосредоточено только на том, что действитель
но важно. При необходимости можно комбинировать различные
трансформации, если используются чистые функции с точки зрения
функционального программирования. Позже мы более подробно
обсудим эту тему, а сейчас ограничимся кратким определением:
функционально чистая функция – это функция, поведение кото
рой определяется только ее входными данными. Чистая функция
не имеет внутреннего состояния и не должна создавать какие-либо
побочные эффекты. Для читателей, основной сферой деятельности
которых является функциональное программирование, такое пове
дение воспринимается вполне естественно. Для других не должно
возникнуть трудностей при переходе на подобный способ написа
ния программ, а я помогу сделать такой переход.
Если вы соблюдаете эти ограничения, то можете объединять, вы
страивать в цепочку и выполнять вложения трансформаций и при
необходимости создавать сложные конвейеры из них. JAX обеспе
чивает возможность произвольной компоновки всех этих операций
трансформации.
Например, можно подготовить функцию для обработки изобра
жения с помощью нейронной сети, а затем автоматически сгене
рировать высокооптимизированную, векторизованную и распреде
ленную версию этой функции, которая может быть автоматически
дифференцируемой. Таким образом, вы получаете в свое распо
ряжение все необходимое для тренировки нейронной сети, напи
сав при этом всего лишь основную часть вычислений. С техниче
ской точки зрения это означает, что, используя исходную функцию
и vmap(), вы получаете другую функцию, обрабатывающую пакет
изображений. Затем с помощью jit() вы компилируете получен
ную функцию в эффективный код для выполнения на GPU или TPU
(или на нескольких таких процессорах в параллельном режиме с по
мощью pmap()). Наконец, вы генерируете функцию для вычисления
градиентов, применив grad(), для тренировки полученной функции
обработки изображений методом градиентного спуска. В NumPy вам
41
Чем JAX отличается от NumPy
пришлось бы написать огромный объем кода для выполнения этой
задачи. В дальнейшем мы увидим несколько весьма интересных
примеров применения такого подхода.
Нет необходимости в реализации всех описанных выше транс
формаций на чистом коде библиотеки NumPy. Вам не придется
вычислять производные вручную, потому что мощный фреймворк
позаботится об этом независимо от того, насколько сложны ваши
функции, и то же самое относится к автоматической векторизации
и распараллеливанию.
На рис. 1.3 (в левой части) показан NumPy как механизм для ра
боты с многомерными массивами с использованием множества по
лезных математических функций. JAX обеспечивает привычный
API, совместимый с NumPy, его многомерными массивами и много
численными функциями, поэтому для исследователей и инженеров
переход на JAX происходит легко. Кроме того, NumPy-подобный API
JAX предоставляет набор мощных трансформаций функций, позво
ляющих сэкономить немало времени и усилий по сравнению с реа
лизацией тех же трансформаций на чистом коде библиотеки NumPy.
Многомерные
массивы
Математические
функции
Автоматическое
дифференцирование
Распараллеливание
NumPy-подобный API
Автоматическое
пакетирование
Рис. 1.3 JAX (справа) – это намного больше, чем просто NumPy (слева). JAX
предоставляет набор компонуемых трансформаций функций с кодом Python +
NumPy для компиляции, векторизации, распараллеливания и автоматического
дифференцирования
В некотором смысле JAX похож на Julia. В Julia также предостав
лена возможность JIT-компиляции, неплохие возможности автома
тического дифференцирования, имеются библиотеки машинного
обучения и полноценная поддержка аппаратного ускорения в режи
42
Глава 1
Когда и зачем используется JAX
ме параллельных вычислений. Но, используя JAX, вы остаетесь в хо
рошо знакомом мире Python. Иногда это имеет большое значение.
1.3
Чем JAX отличается от TensorFlow
и PyTorch
В предыдущем подразделе мы узнали, чем JAX отличается от NumPy.
А теперь сравним JAX с двумя современными фреймворками глубо
кого обучения: PyTorch и TensorFlow.
Ранее было отмечено, что JAX продвигает функциональный под
ход, в отличие от объектно ориентированной методики, принятой
в PyTorch и TensorFlow. Это первый весьма ощутимый факт, с ко
торым вы сталкиваетесь, начиная программировать на JAX. Такой
подход изменяет структуру кода и требует несколько изменить при
вычное мышление. В то же время методика функционального про
граммирования предоставляет мощные трансформации функций,
заставляя вас писать понятный код, и обеспечивает богатую возмож
ностями композиционность. Если вам нравится методика функцио
нального программирования или необходимы предоставляемые ав
томатические трансформации функций, то, возможно, JAX является
правильным вариантом выбора. Библиотеки высокого уровня, пред
назначенные для поддержки нейронных сетей, на основе JAX, такие
как Flax, предоставляют набор функционально чистых классов, ко
торые выглядят хорошо знакомыми всем пользователям, имеющим
опыт работы с Keras или PyTorch. Вы сами убедитесь в этом в главе 11.
JAX-подобные компонуемые трансформации функций
для PyTorch
Компонуемые трансформации функций JAX в немалой степени повлияли на PyTorch. В марте 2022 г. в релизе PyTorch 1.11 разработчики
анонсировали бета-версию библиотеки functorch (https://github.com/
pytorch/functorch), предоставляющую JAX-подобные компонуемые
трансформации функций для PyTorch. Причина заключалась в том, что
в то время многие варианты использования были слишком сложны для
реализации в PyTorch, например вычисление градиентов по каждой выборке данных, запуск ансамбля моделей на одном компьютере, эффективное выполнение пакетных задач во внутренних циклах метаобучения
и эффективное вычисление матриц Якоби и Гессе (якобианов и гессианов), а также их пакетированных версий. Сейчас библиотека functorch
интегрирована в PyTorch как комплект API torch.func (https://pytorch.
org/docs/master/func.html), а оригинальные API functorch являются
устаревшими и не рекомендуемыми к применению, начиная с версии
PyTorch 2.0. Библиотека torch.func пока остается в стадии бета-версии.
Чем JAX отличается от TensorFlow и PyTorch
43
Еще один факт, обращающий на себя внимание практически
сразу, – JAX представляет собой довольно-таки минималистичный
фреймворк. Он не реализует все подряд. TensorFlow и PyTorch – два
широко распространенных и тщательно проработанных фреймвор
ка, в которые включены почти все возможные компоненты (нефор
мально называемые «батарейками»). В отличие от TensorFlow и Py
Torch, JAX чрезвычайно минималистичен, и даже трудно назвать его
фреймворком. Скорее, это библиотека.
Например, JAX не предоставляет какие-либо загрузчики данных
просто потому, что другие библиотеки (например, PyTorch или Ten
sorFlow) делают это лучше. Авторы JAX не ставили перед собой цель
заново реализовать все функции, а сосредоточились исключительно
на ядре. И это в точности тот случай, когда вы можете и даже долж
ны объединить JAX и другие фреймворки глубокого обучения. Впол
не нормально взять функции загрузки данных, скажем, из PyTorch
и воспользоваться ими. PyTorch предлагает превосходные загрузчи
ки данных, так почему бы не взять их на вооружение?
Возможно, в JAX вы почувствуете отсутствие некоторых функ
циональных средств, которые были полезны и предоставлялись ва
шим текущим фреймворком, и потребуется некоторое время, что
бы найти аналогичные средства в экосистеме JAX (или обеспечить
применение средств из текущего фреймворка, например загрузчи
ков данных). Но есть и хорошие новости: если в вашем исходном
фреймворке чего-то не хватало или вас не устраивала существую
щая функциональность, но с этим невозможно было ничего сделать,
то вам предоставляется возможность свободно объединять все, что
пожелаете, в JAX.
Еще один заслуживающий внимания факт: примитивы (элемен
тарные компоненты) JAX являются довольно-таки низкоуровне
выми сущностями, поэтому создание крупных нейронных сетей на
основе операций умножения матриц может потребовать слишком
больших затрат времени. Следовательно, необходим язык более вы
сокого уровня для определения таких моделей. Это похоже на тре
бование чего-то подобного модулю torch.nn вместо torch. JAX не
предоставляет такие API высокого уровня, выходящие за пределы
его функциональных возможностей (так же, как и TensorFlow вер
сии 1 до того, как в версию 2 был добавлен высокоуровневый Keras
API). В JAX не включены никакие дополнительные средства (те са
мые «батарейки»), но это не становится проблемой, поскольку для
экосистемы JAX существуют библиотеки высокого уровня. Flax, Equi
nox и Keras предоставляют все необходимые абстракции высокого
уровня, которые могут потребоваться, так что полученные в итоге
определения моделей, вероятнее всего, будут очень похожи на ана
логичные определения моделей в PyTorch или TensorFlow. Нет ника
кой необходимости писать код нейронных сетей с использованием
Глава 1
44
Когда и зачем используется JAX
NumPy-подобных примитивов. Описания библиотек высокого уров
ня приведены в главе 11, после того как мы рассмотрим во всех под
робностях ядро JAX.
На рис. 1.4 наглядно показаны различия между PyTorch/Tensor
Flow и JAX.
TF/PyTorch
Математи
ческое ядро
Внешние
библиотеки
Метаобучение
Библиотеки
высокого
уровня
поддержки
нейросетей
Обучение
с подкреплением
Оптимизаторы
Трансформеры
от Hugging Face
Библиотеки
обучения
с подкреплением
Библиотеки
высокого уровня
поддержки
нейросетей
Оптимизаторы
Инструменты
отладки
Компиляторы
Загрузчики
данных
XLA
Ядро JAX
• jax.numpy API
• jax.lax
• Трансформации
Инструменты
отладки
Средства
развертывания
Библиотеки
метаобучения
Загрузчики данных
TensorFlow/PyTorch
Инструменты
развертывания
Рис. 1.4 В отличие от PyTorch/TensorFlow, JAX представляет собой минималистичный
фреймворк. Тем не менее многие компоненты, включенные в TensorFlow/PyTorch, доступны
как отдельные модули из экосистемы JAX
Поскольку JAX является расширяемой системой для компонуемых
трансформаций функций, можно с легкостью создавать отдельные
модули для любых целей и объединять их в разнообразных требуе
мых сочетаниях.
Завершая это введение в JAX, следует отметить: если производи
тельность вычислений чрезвычайно важна, если вы высоко цените
функциональное программирование и понятный код, если вы участ
вуете в исследованиях в области глубокого обучения или заинтере
сованы в управлении каждым аспектом своего кода, то JAX – это аб
солютно правильный выбор. В то же время JAX не всегда становится
наилучшим вариантом решения каждой задачи, и в некоторых слу
чаях, таких как развертывания во встроенных системах или работа
с огромными старыми кодовыми базами, использующими другие
фреймворки, PyTorch или TensorFlow может оказаться более пра
вильным вариантом выбора.
Резюме
45
В следующей главе представлено более глубокое введение в JAX
с техническими подробностями для тех, кто предпочитает разби
раться в исходном коде. Мы рассмотрим пример проекта глубокого
обучения с применением JAX для классификации изображений.
Резюме
JAX – это библиотека низкого уровня на языке Python, созданная
компанией Google и весьма интенсивно используемая для иссле
дований в области машинного обучения (но также предоставля
ющая возможность применения в других областях, например для
физических имитаций и числовой оптимизации).
JAX предоставляет API, совместимый с NumPy, для работы с мно
гомерными массивами и математическими функциями NumPy.
JAX содержит полноценный комплект трансформаций функций,
включая автоматическое дифференцирование (autodiff), JIT-ком
пиляцию, автоматическую векторизацию и распараллеливание,
которые можно компоновать произвольно.
JAX обеспечивает высокую производительность вычислений бла
годаря своей способности использовать современное аппаратное
оборудование, такое как TPU и GPU, JIT-компиляцию с примене
нием XLA, автоматическую векторизацию и распараллеливание
в кластере без каких-либо затруднений.
JAX использует парадигму функционального программирования
и требует применения чистых функций без побочных эффектов.
Для JAX существует постоянно расширяющаяся экосистема мо
дулей, которые вы можете свободно комбинировать в различные
структурные блоки, соответствующие вашим потребностям наи
лучшим образом.
В отличие от TensorFlow и PyTorch, JAX представляет собой ми
нималистичный фреймворк, но благодаря непрерывно развива
ющейся экосистеме существует множество качественных библио
тек для удовлетворения конкретных потребностей пользователей.
2
Первая программа в JAX
Темы главы:
структура высокого уровня проекта глубокого обучения
с использованием JAX;
загрузка набора данных;
создание простой нейронной сети в JAX;
использование трансформаций JAX для автоматической
векторизации, вычисления градиентов и JIT-компиляции;
сохранение и загрузка модели;
чистые функции и функции с побочными эффектами.
JAX – это библиотека компонуемых трансформаций для программ
на языке Python с применением NumPy. Хотя сфера применения
этой библиотеки не ограничена исследованиями в области глубо
кого обучения, ее часто считают фреймворком глубокого обучения,
иногда называя третьим после PyTorch и TensorFlow. Поэтому мно
гие начинают именно с изучения JAX для создания приложений глу
бокого обучения.
В этой главе мы будем выполнять учебное упражнение Hello World
для глубокого обучения. Создадим простое приложение нейронной
сети, демонстрирующее подход JAX для формирования модели глу
Учебная задача машинного обучения: классификация рукописных цифр
47
бокого обучения. Это будет модель классификации изображений,
работающая с набором данных MNIST, содержащим рукописные
изображения цифр, – задача, которую вы, вероятно, видели или
даже сами решали с помощью PyTorch или TensorFlow. Проект пред
ставит вам дерево основных трансформаций JAX: grad() для работы
с градиентами, jit() для компиляции и vmap() для автоматической
векторизации. Даже используя только эти три трансформации, вы
можете создавать специализированные решения нейронных сетей,
не требующие распределения в кластере (для распределенных вы
числений существует отдельная трансформация pmap()).
Эта глава предлагает обобщенный комплексный взгляд на весь
фреймворк в целом, особо выделяя функциональные возможности
и основополагающие концепции JAX. В последующих главах эти
концепции будут описаны более подробно.
Код примеров этой главы можно найти в репозитории книги. Я ис
пользую блокнот Google Colab с механизмом времени выполнения
GPU. Это самый простой способ начать работу с JAX. В блокноте Co
lab JAX установлен изначально и использует аппаратное ускорение.
В дальнейшем мы также будем пользоваться Google TPU для работы
с JAX. Поэтому я рекомендую выполнять код примеров в блокноте
Colab. В приложении B вы найдете инструкции, демонстрирующие,
как начать работу в Colab.
Вы можете выполнять код примеров в любой другой системе, но
в этом случае придется установить JAX вручную. Описание установ
ки приведено в приложении A, где содержатся подробные пошаго
вые инструкции.
Начнем с определения учебной задачи, которую необходимо ре
шить.
2.1
Учебная задача машинного обучения:
классификация рукописных цифр
Задача классификации изображений – абсолютный лидер по частоте
упоминания среди задач компьютерного зрения: можно классифи
цировать продукты питания по фотографии прилавков магазинов,
определять типы галактик в астрономии и виды животных в биоло
гии или распознавать цифры в ZIP-коде.
Предположим, что у вас есть набор изображений, каждому из ко
торых присвоена метка, скажем «кот», «собака» или «человек», или
набор изображений с числами от 0 до 9 и с соответствующими мет
ками. Обычно это называют тренировочным набором (training set).
Затем вам нужна программа, принимающая новое изображение без
метки, и необходимо присвоить одну из предварительно опреде
48
Глава 2
Первая программа в JAX
ленных смысловых меток рассматриваемому изображению (фото
графии). Полученная в итоге модель обычно проверяется на тесто
вом наборе (test set), который не был предоставлен модели во время
тренировки. Иногда существует еще и контрольный валидационный
набор (validation set) для настройки гиперпараметров во время тре
нировки и выбора наилучших вариантов.
В новейшую эпоху глубокого обучения обычно вы тренируете
нейронную сеть для выполнения этой задачи. Затем интегрируе
те натренированную нейронную сеть в программу для насыщения
сети данными и интерпретируете полученные итоговые резуль
таты. Например, это может быть программа для классификации
бабочек. Каждая рассматриваемая фотография бабочки обрабаты
вается нейросетью, производящей некоторые выходные данные
(с технической точки зрения активизация соответствующего ней
рона выходного слоя предусмотрена для каждого класса бабочек,
известного этой конкретной нейросети). Затем программа интер
претирует полученные данные, например выбирая класс с макси
мальным уровнем активизации («возбуждения», если воспользо
ваться термином из области исследований человеческого мозга)
и выводя метку с наименованием бабочки в пользовательском
интерфейсе.
Задачи классификации и регрессии
Классификация является одной из стандартных задач машинного
обучения, которая в сочетании с регрессией относится к области обуче
ния с учителем (управляемого, или контролируемого, обучения). В обеих задачах имеется набор данных с примерами (тренировочный набор), предоставляющий контрольный сигнал (supervision signal) (отсюда
и происходит термин «контролируемое обучение» – supervised learning), определяющий, что является корректным для каждого примера.
При классификации контрольный сигнал – метка класса, поэтому обязательно нужно различать определенное фиксированное количество
классов. Это может быть классификация пород собак по фотографии,
или эмоциональная тональность (высказывания) по соответствующему
тексту, или пометка конкретной транзакции по банковской карте как
мошеннической по известным характеристикам и предыдущей истории.
Особый вариант множественной классификации (по множеству классов) называется двоичной классификацией (binary classification) и применяется, когда надо различать только два класса некоторого объекта.
Классы могут быть взаимоисключающими (например, виды животных)
или совмещаемыми (например, присваивание фотографиям предварительно определенных тегов). Первый вариант называется многоклассовой классификацией (multiclass classification), второй – многозначной
классификацией (multilabel classification).
Общий обзор проекта глубокого обучения с использованием JAX
49
При регрессии контрольный сигнал обычно имеет вид непрерывной последовательности чисел, и необходимо предсказать правильное число
для новых вариантов. Это может быть прогноз комнатной температуры
в некоторый момент времени по другим измерениям и факторам, цена
дома, определяемая по его характеристикам и месту расположения, или
размер порции еды на вашей тарелке на основе соответствующей фотографии.
Мы будем использовать широко известный набор данных MNIST
(https://www.tensorfow.org/datasets/catalog/mnist), состоящий из ру
кописных цифр. На рис. 2.2 показаны примеры изображений из это
го набора данных.
Теперь представим мысленную модель такого процесса, затем
рассмотрим все его шаги более подробно. В последующих главах эти
шаги описаны еще детальнее, и я постоянно буду обращать внима
ние читателей на их особенности в каждой конкретной главе.
2.2
Общий обзор проекта глубокого обучения
с использованием JAX
Обычный проект глубокого обучения с использованием JAX включа
ет следующие шаги:
1 выбор
набора данных для конкретной задачи. В нашем приме
ре используется набор данных MNIST;
2 создание загрузчика данных для считывания выбранного набо
ра данных и преобразования его в последовательность пакетов.
Мы будем использовать загрузчик данных из TensorFlow Data
sets;
3 определение модели для работы с отдельной точкой данных.
JAX требует использования чистых функций (более подробно
о них – немного позже в этой главе) без сохранения состоя
ния и побочных эффектов, поэтому необходимо отделить па
раметры модели от функции, применяющей нейронную сеть
к данным. Модель нейронной сети определяется как (a) набор
параметров модели и (b) функция, выполняющая вычисления
с параметрами и входными данными. Также можно воспользо
ваться библиотеками более высокого уровня, поддерживающи
ми нейронные сети, поверх JAX для определения конкретных
моделей;
4 определение модели для работы с пакетом данных. В JAX это
обычно выполняется с помощью автоматической векториза
ции функции модели из шага 3;
50
Глава 2
Первая программа в JAX
5 определение функции потерь (loss function), принимающей па
раметры модели и пакет данных. Функция вычисляет значение
уровня потерь (loss value), обычно представляющее собой не
которую ошибку, которую необходимо минимизировать. Мы
будем использовать категориальную функцию потерь пере
крестной энтропии (categorical cross-entropy loss), часто при
меняемую для многоклассовой классификации;
6 получение градиентов функции потерь с учетом параметров
модели. Функция градиента вычисляется по параметрам моде
ли и входным данным и вычисляет градиенты для каждого па
раметра модели. Для этого мы воспользуемся трансформацией
grad() библиотеки JAX;
7 реализация шага уточнения градиента. Градиенты используют
ся для обновления параметров модели с помощью некоторой
процедуры градиентного спуска (gradient descent). Параметры
модели можно обновлять напрямую (именно это мы сделаем
сейчас) или применить специальный оптимизатор из отдельной
библиотеки (например, Optax; этим мы займемся в главе 11);
8 реализация полного цикла тренировки;
9 компиляция модели для целевой аппаратной платформы с ис
пользованием JIT-компиляции. Этот шаг может существенно
ускорить вычисления;
10 также можно распределить процесс тренировки модели по кла
стеру компьютеров;
11 после выполнения цикла тренировки для нескольких эпох (т. е.
всего обучающего множества) получаем натренированную мо
дель (с обновленным набором параметров), которую можно ис
пользовать для прогнозов или для выполнения любой другой
задачи;
12 сохранение натренированной модели;
13 использование модели. В зависимости от конкретного вариан
та можно развернуть модель в некоторой производственной
инфраструктуре или просто загрузить ее и выполнять вычисле
ния без специализированной инфраструктуры:
a для сохранения и восстановления весов (весовых коэффици
ентов) модели можно воспользоваться стандартными сред
ствами языка Python, например pickle, или более безопас
ными решениями, скажем safetensors. Библиотеки более
высокого уровня на основе JAX также могут предоставить
собственные инструментальные средства для загрузки/сохра
нения моделей;
b для развертывания существует несколько доступных вари
антов. Например, можно преобразовать модель в TensorFlow
или TFLite и использовать их тщательно проработанную эко
систему для развертывания модели.
51
Загрузка и подготовка набора данных
Шаг 1 нашего процесса уже завершен. Продолжим его пошаговое
выполнение, начиная с загрузки набора данных и его подготовки
для решения, разрабатываемого в текущей главе. На рис. 2.1 изобра
жены все необходимые шаги.
Распределенная
тренировка
pmap()
Цикл тренировки
Функция
потерь
grad()
Загрузчик
данных
Набор
данных
Функция
градиента
Пакеты
данных
Оптими
затор
Скомпилированная
и автовекторизованная
модель
jit(vmap(model))
Параметры
модели
Функция
модели
Градиенты
Натренированная
модель
Параметры
модели
Развер
тывание
Функция модели
AWS
SageMaker
TFLite
JAX2TF
Рис. 2.1 Структура высокого уровня для проекта JAX с процедурами загрузки данных,
тренировки и развертывания для работы в производственной среде
2.3
Загрузка и подготовка набора данных
В главе 1 было отмечено, что в JAX не включены какие-либо загрузчи
ки данных, так как этот фреймворк сосредоточен главным образом
на преимуществах функциональности своего ядра. Поэтому можно
без каких-либо затруднений воспользоваться одним из загрузчиков
данных TensorFlow или PyTorch, с которым предпочитаете иметь
дело и хорошо знаете его возможности. Официальная документа
ция JAX содержит примеры для обоих вариантов. Для рассматривае
мого здесь примера воспользуемся API загрузки данных TensorFlow
Datasets. В документации JAX можно найти пример практического
применения загрузчиков данных PyTorch.
TensorFlow Datasets содержит версию набора данных MNIST
с вполне предсказуемым именем mnist. Всего в наборе содержит
ся 70 000 изображений. Этот набор разделен на отдельную группу
Глава 2
52
Первая программа в JAX
тренировки с 60 000 изображений и тестовую часть с 10 000 изобра
жений. Изображения выполнены в оттенках серого цвета и имеют
размер 28×28 пикселов.
Листинг 2.1
Загрузка набора данных
import tensorflow as tf
import tensorflow_datasets as tfds
data_dir = '/tmp/tfds'
data, info = tfds.load(name="mnist",
data_dir=data_dir,
as_supervised=True,
with_info=True)
data_train = data['train']
data_test = data['test']
❶
❶
❷
❸
❸
❸
❸
❹
❹
❶ Импорт необходимых модулей из TensorFlow.
❷ Временный каталог для загрузки данных.
❸ Загрузка набора данных MNIST с использованием функции из TensorFlow Data-
sets. Параметр as_supervised=True определяет возврат данных в виде кортежа
(image, label) вместо словаря.
❹ Извлечение отдельных групп для тренировки и тестирования из загруженного
набора данных.
После загрузки данных можно проверить выборки из набора
MNIST с помощью кода, показанного в листинге 2.2.
Листинг 2.2 Визуальная проверка выборок из загруженного набора данных
import numpy as np
import matplotlib.pyplot as plt
plt.rcParams['figure.figsize'] = [10, 5]
ROWS = 3
COLS = 10
i = 0
fig, ax = plt.subplots(ROWS, COLS)
for image, label in data_train.take(ROWS*COLS):
ax[int(i/COLS), i%COLS].axis('off')
ax[int(i/COLS), i%COLS].set_title(str(label.numpy()))
ax[int(i/COLS), i%COLS].imshow(np.reshape(image, (28,28)), cmap='gray')
i += 1
❶
❶
❷
❷
❸
❸
❸
plt.show()
❶ Импорт модуля рисования из библиотеки Matplotlib и установка размера рабочей области.
❷ Параметры для размещения изображений: необходимо представить изображения в сетке из
3 строк и 10 столбцов.
53
Загрузка и подготовка набора данных
❸ Вывод каждого примера в соответствующей позиции в сетке с отключением осей координат
и предоставлением метки класса как названия изображения. Каждое изображение выводится
в палитре оттенков серого цвета.
Приведенный в листинге 2.2 код генерирует общее изображение,
показанное на рис. 2.2.
Рис. 2.2 Примеры из набора данных MNIST. Каждое рукописное изображение
цифры имеет метку класса, показанную над ним
Поскольку все изображения имеют одинаковый размер, с ними
можно работать одинаковым способом и объединять несколько изо
бражений в пакет (batch). Единственная операция предварительной
обработки, которая может потребоваться, – нормализация, т. е. пре
образование значений байта пиксела (uint8) из целых чисел в диапа
зоне [0, 255] в тип с плавающей точкой (float32) в диапазоне [0, 1].
Листинг 2.3 Предварительная обработка набора данных
и разделение его на пакеты
HEIGHT = 28
WIDTH = 28
CHANNELS = 1
NUM_PIXELS = HEIGHT * WIDTH * CHANNELS
NUM_LABELS = info.features['label'].num_classes
def preprocess(img, label):
"""Resize and preprocess images."""
return (tf.cast(img, tf.float32)/255.0), label
train_data = tfds.as_numpy(
data_train.map(preprocess).batch(32).prefetch(1))
test_data = tfds.as_numpy(
data_test.map(preprocess).batch(32).prefetch(1))
❶ Параметры изображений и набора данных.
❶
❶
❶
❶
❶
❷
❸
❸
54
Глава 2
Первая программа в JAX
❷ Эта функция выполняет преобразование целочисленного значения в число с пла-
вающей точкой float32 и деление его на 255, максимальное целочисленное значение в наборе данных, чтобы получать значения в диапазоне [0, 1].
❸ Применение функции предварительной обработки к тренировочной и тестовой
группам набора данных, генерация потока пакетов с 32 изображениями в каждом и предварительное (упреждающее) извлечение одного пакета.
Мы сообщаем загрузчику данных о необходимости применения
функции preprocess к каждому примеру, распределяем все изобра
жения по пакетам из 32 элементов и с упреждением извлекаем но
вый пакет, не ожидая завершения обработки предыдущего пакета
на GPU.
Шаг 2 завершен, и на текущий момент этого достаточно. Теперь
можно переключиться на разработку нашей первой нейронной сети
с использованием JAX. На протяжении всего процесса разработки
будут постоянно подчеркиваться различия между JAX и более при
вычными фреймворками, такими как NumPy, TensorFlow и PyTorch.
2.4
Простая нейронная сеть
с использованием JAX
Переходим к шагу 3. Здесь мы используем нейронную сеть с пря
мой связью (с прямым распространением сигнала), известную как
многослойный перцептрон (multilayer perceptron – MLP). Это весьма
простая (и далеко не самая лучшая) сеть, выбранная для демонстра
ции важных концепций без излишнего усложнения. Более продви
нутое решение будет использоваться в главе 11.
Нашим решением будет двухслойный перцептрон – это обычный
учебный пример для нейронных сетей. Схема разрабатываемой
нейронной сети показана на рис. 2.3.
Изображение «выпрямляется» в одномерный массив из 784 зна
чений (так как изображение размером 28×28 содержит 784 пиксе
ла), и этот массив становится входным для нейронной сети. Входной
слой (input layer) отображает каждый из 784 пикселов изображения
в отдельную единицу входных данных. Далее располагается полно
связный (или плотный) слой с 512 нейронами. Его называют скры
тым слоем (hidden layer), потому что он размещен между входным
и выходным слоями. Каждый из 512 нейронов «рассматривает» все
входные элементы одновременно. За скрытым слоем следует еще
один полносвязный слой, который называют выходным слоем (out
put layer). Он содержит 10 нейронов, т. е. столько же, сколько клас
сов в наборе данных. Каждый выходной нейрон отвечает за соот
ветствующий ему класс. Например, нейрон #0 выдает вероятность
принадлежности классу 0, нейрон #1 – аналогичную вероятность для
Простая нейронная сеть с использованием JAX
55
класса 1 и т. д. Таким образом, выходной слой формирует распреде
ление вероятностей со значениями вероятностей, предсказанными
нейронами в выходном слое, в сумме равными 1.
Изображение
Вероятности
соответствия
классу
«Выпрям
ление»
10 нейронов
512 нейронов
784 пиксела
Рис. 2.3 Структура нейронной сети. Изображение размером 28×28 пикселов
«выпрямляется» в последовательность из 784 пикселов (не имеющую
двумерной структуры) и передается во входной слой сети. Далее располагается
полносвязный скрытый слой с 512 нейронами, за ним следует еще один
полносвязный выходной слой с 10 нейронами, создающий активации целевого
класса
В каждом слое с прямой связью реализуется простая функция
y = f(x × w + b), состоящая из весов w, умножаемых на входящие дан
ные x, и систематической ошибки (необъективности) b, прибавляе
мой к произведению. Функция активации f() – нелинейная функция,
применяемая к результату умножения и сложения.
Глава 2
56
Первая программа в JAX
Процесс создания нейронной сети в JAX отличается от аналогич
ного процесса в PyTorch/TensorFlow в нескольких аспектах, а именно
использованием генераторов случайных чисел для инициализации
параметров модели и структурой кода модели и ее параметров.
Первое отличие: для сохранения функциональной чистоты гене
раторы случайных чисел в JAX требуют предоставления своего со
стояния извне (на рис. 2.4 показано, что эту роль играет PRNGKey).
Второе отличие: функция прямого прохода также не должна иметь
собственного сохраненного состояния и обязана быть функциональ
но чистой, поэтому параметры модели передаются в нее как неко
торые входные данные. В этом и состоит отличие от PyTorch и Ten
sorFlow, где параметры модели хранятся внутри некоторых объектов
вместе с кодом. В других аспектах функции прямого прохода прак
тически одинаковы во всех трех фреймворках.
Наглядная схема процесса создания модели показана на рис. 2.4.
Вывод
результата
Подготовка функции
прямого прохода
def predict(params, image)
Начальные
параметры
модели
Процедура
тренировки
Инициализация
параметров нейросети
params = …
Данные
для тренировки
Применение нейросети
predict(params, image)
Тренированные
параметры
модели
Входные
данные
PRNGKey
Рис. 2.4
в JAX
Процесс инициализации и практического применения нейронной сети
В первую очередь необходимо инициализировать параметры слоя.
2.4.1
Инициализация нейронной сети
Прежде чем начать тренировку нейронной сети, необходимо ини
циализировать все параметры b и w случайными числами.
Простая нейронная сеть с использованием JAX
Листинг 2.4
57
Инициализация нейронной сети
from jax import random
LAYER_SIZES = [28*28, 512, 10]
PARAM_SCALE = 0.01
def init_network_params(sizes, key=random.PRNGKey(0), scale=1e-2):
"""Initialize all layers for a fully-connected
neural network with given sizes"""
# Инициализация всех слоев для полносвязной нейросети
# с заданными размерами.
def random_layer_params(m, n, key, scale=1e-2):
"""A helper function to randomly initialize
weights and biases of a dense layer"""
# Вспомогательная функция для случайной инициализации
# весов и необъективностей плотного (полносвязного) слоя.
w_key, b_key = random.split(key)
return scale * random.normal(w_key, (n, m)),
➥scale * random.normal(b_key, (n,))
keys = random.split(key, len(sizes))
return [random_layer_params(m, n, k, scale)
➥for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
❶
❷
❸
❹
❺
❶
❷
❸
❹
❺
Список размеров слоев.
Параметр для масштабирования случайных значений.
Генерация случайных ключей (более подробно об этом в главе 7).
Генерация случайных значений для параметров слоя w и b.
Запуск процедуры генерации для всех слоев.
Работа со случайными числами в JAX отличается от NumPy, так
как JAX требует применения чистых функций, а генераторы случай
ных чисел (ГСЧ/RNG – random number generator) NumPy не являют
ся чистыми, поскольку используют скрытое внутреннее состояние.
ГСЧ TensorFlow и PyTorch обычно также используют внутреннее
состояние. В JAX реализованы чисто функциональные генераторы
случайных чисел, они подробно рассматриваются в главе 9. Сейчас
достаточно понимать, что необходимо предоставить для каждого
вызова функции рандомизации состояние ГСЧ (RNG), которое в на
шем случае называется ключом (key), и вы должны использовать
каждый ключ только один раз, поэтому каждый раз, когда нужен но
вый ключ, старый разделяется на требуемое количество новых клю
чей. Шаг 3a завершен.
58
2.4.2
Глава 2
Первая программа в JAX
Нейронная сеть с прямой связью
Далее необходима функция, выполняющая все вычисления ней
ронной сети с прямой связью (с прямым распространением сигна
ла). Это почти та же функция, что и forward(self, x) в PyTorch или
call(self, x) в TensorFlow/Keras, за исключением способа передачи
параметров модели. В JAX такая функция выглядит приблизительно
так: predict(params, x). Имя функции не имеет значения, ее мож
но назвать forward(), call() или как-то еще. Здесь самое главное
свойство заключается в том, что функция не является членом како
го-либо класса, а принимает параметры модели как свои параметры
(параметры функции) (прошу извинить за тавтологию). Именно по
этому вместо указателя на объект self используется имя params.
Для обеспечения прямой связи у нас уже имеются начальные
значения параметров b и w. Единственная отсутствующая деталь –
функция активации. Воспользуемся широко известной функцией
активации Swish из библиотеки jax.nn.
Функции активации
Функции активации являются чрезвычайно важными компонентами
в мире глубокого обучения. Они обеспечивают нелинейность вычислений нейронной сети. При отсутствии нелинейности многослойная нейросеть с прямым распространением сигнала становится равнозначной
одному слою. По правилам простой математики линейная комбинация
линейных комбинаций входных данных остается линейной комбинацией входных данных; именно это и делает единственный нейрон. Известно, что возможности одного нейрона для решения сложных задач
классификации ограничены линейно разделимыми (под)задачами (вероятно, вы слышали о хорошо известной проблеме XOR, которую невозможно разрешить с помощью линейного классификатора). Поэтому
функции активации обеспечивают экспрессивность нейронной сети
и предотвращают ее свертывание в более простую модель.
В настоящее время предлагается множество найденных ранее разнообразных функций активации. В этой области исследования начались
с простых и понятных функций, таких как сигмоида (S-образная кривая)
и гиперболический тангенс. Это гладкие функции, обладающие свойствами, которые очень нравятся математикам, например дифференцируемость в каждой точке.
Затем появилась новая разновидность функции, ReLU (rectified linear
unit – блок линейной ректификации). Функция ReLU не являлась глад-
Простая нейронная сеть с использованием JAX
59
кой, так как в точке x = 0 ее производная не существует. Тем не менее специалисты-практики обнаружили, что нейронные сети быстрее
обучаются при использовании функции ReLU.
ReLU(x) = max(0, x)
После этого было обнаружено множество других функций активации –
некоторые экспериментальным путем, другие в результате конструктивных разработок. Наиболее широко известными разработанными функциями являются линейные блоки кривой ошибок Гаусса (Gaussian error
linear units – GELU, https://arxiv.org/abs/1606.08415) и линейные блоки
масштабируемой экспоненциальной кривой (scaled exponential linear
units – SELU, https://arxiv.org/abs/1706.02515).
Среди самых последних тенденций в сфере глубокого обучения
можно выделить автоматическое обнаружение (automatic discovery),
которое обычно называют NAS (Neural Architecture Search – поиск
нейронной архитектуры). Основная идея этой методики заключается
в проектировании обширного, но управляемого пространства поиска,
описывающего компоненты, интересующие исследователя. Компонентами могут быть функции активации, типы слоев, формулы обновления оптимизаторов и т. п. Затем запускается автоматическая процедура интеллектуального поиска в этом пространстве. Другие методики
также могут использовать обучение с подкреплением, эволюционные
вычисления или даже метод градиентного спуска. Именно таким способом была обнаружена функция активации Swish (https://arxiv.org/
abs/1710.05941).
Swish(x) = x · sigmoid(βx)
Методика NAS имеет захватывающую историю, и я верю в то, что богатая возможностями экспрессивность JAX может внести существенный
вклад в развитие этой области исследований. Возможно, кое-кто из читателей этой книги станет автором весьма впечатляющего достижения
в сфере глубокого обучения.
Здесь мы разрабатываем функцию прямого прохода, часто назы
ваемую функцией прогноза, или предсказания (predict function). Она
принимает изображение для классификации и выполняет все вы
числения прямого прохода для создания активизаций в выходном
слое нейронов. Нейрон с наивысшим уровнем активизации опре
деляет класс входного изображения (т. е. если наивысший уровень
активизации находится в нейроне 5, то в соответствии с предельно
понятной методикой нейронная сеть определила, что входное изо
бражение содержит рукописную цифру 5).
Глава 2
60
Листинг 2.5
Первая программа в JAX
Прямой проход нейронной сети
import jax.numpy as jnp
from jax.nn import swish
def predict(params, image):
"""Function for per-example predictions."""
# Функция прогнозирования для каждого образца.
activations = image
for w, b in params[:-1]:
outputs = jnp.dot(w, activations) + b
activations = swish(outputs)
final_w, final_b = params[-1]
logits = jnp.dot(final_w, activations) + final_b
return logits
❶
❷
❸
❹
❺
❺
❻
❻
Импорт функции активации Swish.
Обратите внимание: в функцию передается изображение и параметры нейросети.
Инициализация активизаций пикселами входного изображения.
Циклы от первого до предпоследнего слоя.
Успешное обновление активизаций с использованием выходных данных каждого слоя.
❻ Для последнего слоя функция активации не применяется.
❶
❷
❸
❹
❺
Обратите внимание: здесь мы передаем список параметров. Этот
подход отличается от обычной программы на PyTorch или Tensor
Flow, где те же параметры обычно скрыты внутри класса, а функция
использует переменные (члены) класса для доступа к ним.
Рассмотрим повнимательнее, как структурированы вычисления
нейронной сети. В JAX обычно имеются две функции для нейронных
сетей: одна для инициализации параметров, другая для применения
конкретной нейронной сети к некоторым входным данным. Первая
функция возвращает параметры как некую структуру данных (в на
шем примере: список массивов; в дальнейшем это будет специаль
ная структура данных под названием pytree). Вторая функция при
нимает параметры и данные, а возвращает результат применения
нейронной сети к полученным данным. В будущем этот паттерн
будет появляться многократно, даже во фреймворках нейронных
сетей высокого уровня.
Вот и все. Можно использовать эту новую функцию для прогно
зирования по каждому представленному образцу. Далее в листинге
2.6 мы генерируем изображение того же размера, что и в использу
емом наборе данных со случайными значениями пикселов. Затем
передаем их в функцию predict(). Также можно воспользоваться
реально существующим изображением из того же набора данных.
Не следует ожидать хороших результатов, поскольку нейронная сеть
пока еще не тренирована. Сейчас нас интересует только выводимая
vmap: автоматически векторизованные вычисления для обработки пакетов
61
форма, и мы видим, что функция прогнозирования выводит кортеж
из 10 активизаций для 10 классов.
Листинг 2.6
Формирование прогноза
random_flattened_image = \
random.normal(random.PRNGKey(1), (28*28*1,))
preds = predict(params, random_flattened_image)
print(preds.shape)
>>> (10,)
❶
❷
❸
❶ Генерация изображения размером 28×28 пикселов с одним цветовым каналом
и случайными значениями пикселов.
❷ Передача изображения в функцию прогнозирования.
❸ Прогноз содержит кортеж активизаций для 10 классов.
Шаг 3b также завершен.
Итак, все выглядит неплохо, но требуется обработка пакетов изо
бражений, а написанная функция предназначена для работы только
с одним изображением. Здесь нам поможет автоматическое пакети
рование.
2.5
vmap: автоматически векторизованные
вычисления для обработки пакетов
Нельзя обойти вниманием тот факт, что написанная в предыдущем
разделе функция predict() была предназначена для обработки од
ного элемента, поэтому не будет работать, если передать ей пакет
изображений. Для проверки того, что происходит при передаче па
кетов, можно сгенерировать случайный пакет из 32 изображений
с размером каждого 28×28 пикселов и передать его в функцию predict(). Измененный код показан в листинге 2.7. Мы также добавили
обработку исключений, но только для того, чтобы сократить размер
сообщения об ошибке и выделить самую важную часть.
Листинг 2.7
Выполнение прогнозирования для пакета
random_flattened_images = \
random.normal(random.PRNGKey(1), (32, 28*28*1))
try:
preds = predict(params, random_flattened_images)
except TypeError as e:
print(e)
>>> dot_general requires contracting dimensions
❶
❷
Глава 2
62
Первая программа в JAX
to have the same shape, got (784,) and (32,).
# dot_general требует соответствующие контракту измерения
# для обеспечения одинаковой формы, получено (784,) и (32,).
❸
❶ Генерация пакета из 32 случайных изображений размером 28×28 пикселов и од-
ним цветовым каналом для каждого изображения.
❷ Передача пакета изображений в функцию прогнозирования, обрабатывающую
один элемент.
❸ Сообщение об ошибке показывает, что функция, обрабатывающая один элемент,
не может обработать пакет изображений.
Получение этого сообщения об ошибке не вызывает удивления,
так как функция predict() представляет собой упрощенную реали
зацию матричных вычислений, предполагающих конкретные фор
мы массива. В сообщении об ошибке указано, что измерение, по
которому вычисляется скалярное произведение, должно иметь оди
наковую форму. Ожидается 784 числа для вычисления скалярного
произведения по массиву весов, а новое измерение пакета (в данном
примере его размер равен 32) приводит к ошибке. Необходимо чтото изменить, чтобы устранить возникшую проблему и адаптировать
программу для работы с таким дополнительным измерением.
Тензоры, матрицы, векторы и скаляры
В сфере глубокого обучения многомерные массивы являются основными структурами данных, используемыми для обмена информацией
между нейронными сетями и их слоями. Такие структуры также называют тензорами (tensors). В математике и физике понятие тензора имеет
более строгий и сложный смысл, поэтому не стоит беспокоиться, если
что-то вдруг покажется вам слишком сложным в строгих определениях
тензоров. А здесь, в глубоком обучении, это просто синонимы многомерных массивов. И если вы работали с NumPy, то вам уже известно почти
все, что необходимо знать о тензорах.
Существуют конкретные формы тензоров, или многомерных массивов.
Матрица – это тензор с двумя измерениями (или тензор ранга-2), вектор – тензор с одним измерением (тензор ранга-1), а скаляр (или просто
число) представляет собой тензор с нулем измерений (тензор ранга-0).
Таким образом, тензоры являются обобщением для скаляров, векторов
и матриц и далее до произвольного количества измерений (ранга).
Например, значением выбранной конкретной функции потерь является
скаляр (только одно число). Массив вероятностей определения класса
на выходе классификационной нейронной сети для одного элемента
входных данных – вектор размером k (количество классов) с одним измерением (не следует путать размер и ранг). Массив таких прогнозов
для пакета данных (при одновременном вводе нескольких элементов
vmap: автоматически векторизованные вычисления для обработки пакетов
63
данных) представляет собой матрицу размером k×m (где k – количество
классов, а m – размер пакета). RGB-изображение – это тензор ранга-3
с тремя измерениями (ширина, высота и цветовые каналы). Пакет RGBизображений становится тензором ранга-4 (добавляется измерение пакета). Поток кадров видео также можно считать тензором ранга-4 (здесь
добавляется новое измерение времени). Пакет видео – это уже тензор
ранга-5 и т. д. При глубоком обучении вы обычно работаете с тензорами,
количество измерений которых не превышает четырех-пяти.
Какими возможностями мы располагаем, чтобы устранить воз
никшую проблему?
Во-первых, существует простейшее решение. Можно написать
цикл, выполняющий разделение пакета на отдельные изображения
и обеспечивающий их последовательную обработку. Это будет ра
ботать, но такой подход был бы неэффективным, поскольку боль
шинство современных аппаратных устройств способны выполнять
гораздо больше вычислений в единицу времени. В этом случае ап
паратура становится существенно недозагруженной. Если вы уже
работали с MATLAB, NumPy или аналогичными фреймворками, то
оценили в полной мере преимущества векторизации. Такой подход
стал бы эффективным решением проблемы.
Поэтому существует второй вариант: переписать и вручную век
торизовать функцию predict() так, чтобы обеспечить прием па
кетов данных в качестве ввода. Обычно это означает, что входные
тензоры дополняются измерением пакета, следовательно, необхо
димо переписать код вычислений с учетом нового измерения. Это
легко сделать для простых вычислений, но все становится гораздо
сложнее для функций с хитроумным содержимым. Такой способ
обычно применяется при написании нейронных сетей на основе
NumPy или при использовании примитивов низкого уровня Ten
sorFlow и PyTorch. Библиотеки более высокого уровня в экосистеме
TensorFlow/PyTorch могут предоставить интерфейс, скрывающий
подобные сложности.
Переходим к третьему варианту – автоматической векторизации.
JAX предоставляет трансформацию vmap(), преобразующую функ
цию, работающую с одним элементом, в функцию, способную об
рабатывать пакеты данных. Именно такую возможность вы будете,
вероятнее всего, использовать бóльшую часть времени в JAX, так как
это наиболее удобный способ, к тому же обеспечивающий превос
ходную производительность. Думаю, он вам очень понравится. Тем
не менее ничто и никто не запрещает вам пользоваться и другими
вариантами.
Глава 2
64
Листинг 2.8
Первая программа в JAX
Автоматическая векторизация функции
from jax import vmap
batched_predict = vmap(predict, in_axes=(None, 0))
❶
❷
❶ Импорт трансформации vmap().
❷ Создание функции, обрабатывающей пакеты данных, на основе функции, работа-
ющей с одним элементом, с помощью трансформации vmap().
Операция выполняется всего лишь в одной строке. Вы можете
пропустить следующий абзац, потому что подробное описание будет
приведено в главе 6. Для заинтересованных читателей ниже кратко
описан смысл кода в листинге 2.8.
Параметр in_axes определяет, по каким осям входного массива
производится отображение (т. е. векторизация). Их длина обязатель
но должна быть равна количеству позиционных аргументов функ
ции. Значение None указывает, что не требуется отображение какихлибо осей, и в нашем примере это соответствует первому параметру
функции predict(), т. е. params. Этот параметр остается одинаковым
для любого прямого прохода (прямого распространения сигнала),
поэтому нет необходимости в его пакетировании (хотя если бы мы
использовали отдельные веса нейросети для каждого вызова, то вос
пользовались бы этой опцией). Второй элемент в кортеже in_axes со
ответствует второму параметру функции predict(), т. е. image. Нуле
вое значение указывает, что необходимо пакетирование по первому
(нулевому) измерению, содержащему различные изображения. Если
предположить вариант, в котором пакетируемое измерение будет
располагаться в другой позиции в тензоре, то нам бы пришлось за
менить нулевое значение на соответствующий индекс.
Теперь можно применять векторизованную функцию к пакетам
данных и получать корректные выходные результаты.
Листинг 2.9
Использование функции после трансформации vmap()
batched_preds = batched_predict(params, random_flattened_images)
print(batched_preds.shape)
❶
>>> (32, 10)
❷
❶ Передача пакета изображений в функцию batched_predict(), полученную с по
мощью трансформации vmap() из исходной функции predict().
❷ Теперь вывод корректен и содержит 10 активизаций классов для каждого из
32 элементов принятого пакета.
Обратите внимание на один весьма важный факт. Мы не измени
ли исходную функцию. Мы создали новую функцию.
vmap() – чрезвычайно полезная трансформация, поскольку она
полностью освобождает пользователя от выполнения векторизации
Autodiff: как вычислять градиенты, ничего не зная о производных
65
вручную. Возможно, векторизация является не самой интуитивно
понятной процедурой, так как приходится воспринимать ее в кон
тексте матриц или тензоров и их измерений. Это не так-то просто
для каждого обычного человека, и процесс векторизации может
стать потенциальным источником ошибок, поэтому наличие ав
томатической векторизации в JAX – это истинное чудо. Вы пише
те функцию для обработки одного экземпляра данных, затем с по
мощью vmap() превращаете ее в функцию, работающую с пакетами
данных. Осталась нерассмотренной только одна часть общего про
цесса – тренировка, но это еще одна весьма интересная часть. Нам
нужна тренировка.
2.6
Autodiff: как вычислять градиенты, ничего
не зная о производных
Для тренировки нейронной сети обычно используется процедура
градиентного спуска (gradient descent). Хотя общий принцип оста
ется тем же, что и в PyTorch/TensorFlow, сам процесс существенно
отличается от процедуры в этих фреймворках. Полное описание
и сравнение с PyTorch/TensorFlow представлено в главе 4.
Мы будем использовать один из самых простых мини-батч1 ме
тодов градиентного спуска с экспоненциально убывающей нормой
(скоростью) обучения без инерции, похожий на простой базовый
стохастический метод градиентного спуска (stochastic gradient de
scent – SGD) оптимизатора в любом фреймворке глубокого обучения.
Процедура градиентного спуска
Метод градиентного спуска (gradient descent) – это простая итеративная
процедура поиска локальных минимумов дифференцируемой функции.
Для дифференцируемой функции можно найти градиент, т. е. направление наиболее существенного изменения функции. Если двигаться
противоположно градиенту, то это будет направление наискорейшего
(наикратчайшего) спуска. Мы достигаем локального минимума функции,
выполняя повторяющиеся шаги.
Нейронные сети – это дифференцируемые функции, определяемые
своими параметрами (весами, или весовыми коэффициентами). Необходимо найти такое сочетание весов, которое минимизирует некоторую
функцию потерь (loss function), которая вычисляет величину несовпадения (расхождения) между прогнозом модели и эталонными данными.
1
Мини-батч (mini-batch) – небольшое подмножество тренировочного на
бора. – Прим. перев.
Глава 2
66
Первая программа в JAX
Чем меньше несовпадение, тем точнее прогнозы. Для решения нашей
задачи можно применить метод градиентного спуска.
Потери
Мы начинаем работу с некоторыми случайно взятыми весами, выбранной конкретной функцией потерь и тренировочным набором данных.
Затем многократно вычисляем градиент функции потерь с учетом текущих значений весов для тренировочного набора данных (или некоторого пакета (подмножества) из этого набора). После вычисления градиента для каждого веса в нейронной сети каждый вес может быть обновлен
в противоположном направлении посредством вычитания некоторой
части градиента из значения веса. Процедура останавливается после
завершения предварительно определенного количества итераций, при
прекращении увеличения потерь или по какому-либо другому условию.
Наглядная схема этого процесса показана на рис. 2.5.
Градиент
Кривая
потерь
(ландшафт
отбора)
Начальная
точка
Шаг
Локальный
минимум
Локальные
минимумы
Глобальный
минимум
W0 W1
Wk
Вес, w
Рис. 2.5 Графическая схема выполнения шагов градиентного спуска
в сочетании с кривой потерь (ландшафтом отбора)
Функция потерь – это кривая, на которой существует значение потери,
соответствующее каждому значению веса. Такую кривую также называют кривой потерь (loss curve), или ландшафтом отбора (fitness landscape) (последний вариант имеет смысл использовать в более сложных
случаях с количеством измерений более одного).
Сейчас мы начали с некоторого начального случайно выбранного веса
(W0), а после выполнения последовательности шагов пришли к глобальному минимуму, соответствующему конкретному значению веса (Wk).
Специальный параметр, называемый скоростью (нормой) обучения
(learning rate), определяет, насколько большое или малое значение градиента мы получаем.
67
Autodiff: как вычислять градиенты, ничего не зная о производных
Повторяя эти шаги, мы проходим по траектории, ведущей к локальному
минимуму функции потерь. При этом мы надеемся, что этот локальный
минимум совпадает с глобальным или по крайней мере незначительно
отличается от него. Достаточно странно, но такая методика работает для
нейронных сетей. Причина успешной работы этой методики – тема, интересная сама по себе.
На рис. 2.5 для наглядности выбрана «хорошая» начальная точка, из
которой легко достигается глобальный минимум. Начало процедуры из
других точек ландшафта отбора может приводить к локальным минимумам (несколько локальных минимумов показаны на рис. 2.5).
Для такой базовой процедуры градиентного спуска существует множество усовершенствований, в том числе методы градиентного спуска
с инерцией и адаптивного градиентного спуска, такие как Adam, Adadelta, RMSProp, LAMB и т. д. Многие из них также помогают исключить из
рассмотрения некоторые локальные минимумы.
Для реализации выбранной процедуры градиентного спуска не
обходимо начать с некоторой произвольно выбранной точки в про
странстве параметров (мы уже сделали это, когда инициализирова
ли параметры нейронной сети в предыдущем разделе).
2.6.1
Функция потерь
Теперь нам нужна функция потерь для вычисления текущих значе
ний набора параметров на тренировочном наборе данных. Функция
потерь вычисляет несоответствие (величину расхождения) между
прогнозом модели и эталонными значениями из меток тренировоч
ного набора. Существует множество разнообразных функций потерь
для конкретных задач машинного обучения, а мы воспользуемся
простой функцией потерь, подходящей для многоклассовой класси
фикации, – категорийной функцией потерь перекрестной энтропии
(categorical cross-entropy function).
Это почти такая же функция потерь, что и в других фреймворках,
но с одним отличием: подобно функции прямого прохода, она не
пременно должна быть функционально чистой, поэтому для нее по
требуется предоставление параметров модели.
Листинг 2.10
Реализация функции потерь
from jax.nn import logsumexp
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits = batched_predict(params, images)
❶
❷
68
Глава 2
Первая программа в JAX
log_preds = logits - logsumexp(logits)
return -jnp.mean(targets*log_preds)
❸
❹
❶ Импорт функции logsumexp.
❷ Генерация активизаций нейросети (часто называемых логитами (logits)) для
входного пакета изображений.
❸ Вычисление логарифмов вероятностей с помощью функции logsumexp().
❹ Вычисление значения категорийной потери перекрестной энтропии.
Здесь мы используем функцию logsumexp() – это общепринятый
прием в машинном обучении для нормализации вектора логариф
мов вероятностей с целью исключения проблем числового перепол
нения или потери значимости. Если вы хотите получить более под
робную информацию по этой теме, то см.: https://mng.bz/BdJq.
Для функции loss() требуются эталонные значения, или цели
(targets), чтобы вычислять несоответствия между прогнозом моде
ли и контрольными эталонными данными. Прогнозы модели уже
представлены в форме активизаций классов, где каждый выходной
нейрон выдает некоторую (числовую) оценку для соответствующего
класса. Изначально целями являются просто номера классов, т. е. для
класса 0 это число 0, для класса 1 – число 1 и т. д. Необходимо преоб
разовать эти числа в активизации, и для таких случаев использует
ся специальная схема прямого кодирования с одним активным со
стоянием (one-hot encoding). Класс «0» создает массив активизаций
с числом 1 в позиции 0 и с нулями во всех прочих позициях. Класс
«1» размещает 1 в позиции 1 и т. д. Такое преобразование выполня
ется за пределами функции loss().
После определения функции потерь завершается шаг 5. Мы гото
вы к реализации обновления процедуры градиентного спуска.
2.6.2
Определение градиентов
Логика проста. Требуется вычислять градиенты функции потерь
с учетом параметров модели на основе текущего пакета данных.
Здесь нам поможет трансформация grad(), которая принимает не
которую функцию (в нашем случае функцию потерь) и создает функ
цию, вычисляющую градиент функции потерь с учетом конкретно
заданного параметра; по умолчанию это первый параметр исходной
функции (в данном случае params). Таким образом, шаг 6 завершен.
2.6.3
Шаг обновления градиента
Обновление градиента в JAX отличается от аналогичной процеду
ры в других фреймворках, таких как TensorFlow и PyTorch. В этих
фреймворках вы обычно получаете градиенты после выполнения
прямого прохода, и фреймворк отслеживает все операции, выпол
Autodiff: как вычислять градиенты, ничего не зная о производных
69
ненные с интересующими пользователя тензорами. JAX применяет
другой подход. Он выполняет преобразование исходной функции
и генерирует другую функцию, вычисляющую градиенты. Затем вы
вычисляете градиенты, предоставляя все требуемые параметры,
веса нейросети и данные в эту специализированную функцию.
В рассматриваемом здесь примере мы вычисляем градиенты, за
тем обновляем все параметры в направлении, противоположном
вычисленному градиенту (следовательно, необходим знак минус
в формулах обновления весов). Все градиенты масштабируются с ис
пользованием параметра скорости обучения, который зависит от ко
личества эпох (одна эпоха – это полный проход по тренировочному
набору данных). Мы сформировали экспоненциально убывающую
скорость обучения, поэтому для более поздних эпох скорость обуче
ния будет ниже, чем для ранних.
Листинг 2.11
Реализация шага обновления градиента
from jax import grad
INIT_LR = 1.0
DECAY_RATE = 0.95
DECAY_STEPS = 5
def update(params, x, y, epoch_number):
grads = grad(loss)(params, x, y)
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)]
❶
❷
❸
❹
❺
❻
❼
❼
Импорт трансформации grad().
Начальная скорость обучения.
Параметр снижения скорости обучения.
Этот параметр определяет, сколько эпох должно пройти, прежде чем скорость
обучения снизится в очередной раз.
❺ Мы генерируем функцию вычисления градиентов с помощью трансформации
grad() и немедленно применяем ее к текущим параметрам и данным (params, x,
y) для получения значений градиентов.
❻ Вычисление скорости обучения для текущего шага.
❼ Возврат обновленных параметров с принятием небольшого шага (определяемого
параметром скорости обучения lr) в направлении, противоположном градиенту.
❶
❷
❸
❹
В рассматриваемом здесь примере вы не вычисляете напрямую
выбранную функцию потерь. Вместо этого вычисляются только гра
диенты. Во многих случаях также требуется отслеживание значений
потерь, и JAX предоставляет еще одну функцию value_and_grad(),
которая вычисляет значение и градиент функции. Поэтому можно
изменить функцию update() соответствующим образом, как пока
зано в листинге 2.12.
Глава 2
70
Первая программа в JAX
Листинг 2.12 Реализация шага обновления градиента с вычислением
значений потерь и градиента
from jax import value_and_grad
def update(params, x, y, epoch_number):
loss_value, grads = value_and_grad(loss)(params, x, y)
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in
zip(params, grads)], loss_value
❶ Импорт функции value_and_grad.
❷ Вычисление значений потерь и градиента.
❸ Возврат обновленных параметров и значения потерь.
❶
❷
❸
Шаг 7 завершен.
2.6.4
Цикл тренировки
Теперь необходимо выполнить цикл тренировки с заданным коли
чеством эпох. Для этого потребуется еще несколько вспомогатель
ных функций, вычисляющих точность (прогнозов) и обеспечива
ющих некоторое журналирование для отслеживания всей важной
информации, получаемой во время тренировки.
Листинг 2.13
Реализация процедуры градиентного спуска
from jax.nn import one_hot
num_epochs = 25
❶
❷
def batch_accuracy(params, images, targets):
❸
images = jnp.reshape(images, (len(images), NUM_PIXELS))
predicted_class = jnp.argmax(batched_predict(params, images), axis=1)
return jnp.mean(predicted_class == targets)
def accuracy(params, data):
accs = []
for images, targets in data:
accs.append(batch_accuracy(params, images, targets))
return jnp.mean(jnp.array(accs))
❹
import time
for epoch in range(num_epochs):
start_time = time.time()
losses = []
for x, y in train_data:
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = one_hot(y, NUM_LABELS)
❺
❻
Autodiff: как вычислять градиенты, ничего не зная о производных
params, loss_value = update(params, x, y, epoch)
losses.append(loss_value)
epoch_time = time.time() - start_time
start_time = time.time()
train_acc = accuracy(params, train_data)
test_acc = accuracy(params, test_data)
eval_time = time.time() - start_time
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
print("Eval in {:0.2f} sec".format(eval_time))
print("Training set loss {}".format(jnp.mean(jnp.array(losses))))
print("Training set accuracy {}".format(train_acc))
print("Test set accuracy {}".format(test_acc))
71
❼
❽
❾
❿
❿
❶ Еще одна вспомогательная функция для генерации прямого кодирования с од❷
❸
❹
❺
❻
❼
❽
❾
❿
ним активным состоянием для метки класса.
Определение количества эпох в тренировке.
Вычисление точности (процента правильных ответов) для пакета данных.
Вычисление точности для всего набора данных с множеством пакетов.
«Выпрямление» входного изображения в массив идентификаторов (ID).
Преобразование метки класса в объект прямого кодирования с одним активным
состоянием.
Обновление параметров по одному шагу процедуры градиентного спуска.
Сохранение значения потерь для накопления статистических данных о трени
ровке.
Фиксация в журнале времени (эпохи) для накопления статистических данных
о тренировке.
Вычисление для каждой эпохи точности по тренировочным и тестовым данным.
Мы готовы к выполнению первого тренировочного цикла:
Epoch 0 in 36.39 sec
Eval in 8.06 sec
Training set loss 0.41040700674057007
Training set accuracy 0.9299499988555908
Test set accuracy 0.931010365486145
Epoch 1 in 32.82 sec
Eval in 6.47 sec
Training set loss 0.37730318307876587
Training set accuracy 0.9500166773796082
Test set accuracy 0.9497803449630737
Epoch 2 in 32.91 sec
Eval in 6.35 sec
Training set loss 0.3708733022212982
Training set accuracy 0.9603500366210938
Test set accuracy 0.9593650102615356
Epoch 3 in 32.88 sec
...
Epoch 23 in 32.63 sec
Eval in 6.32 sec
Training set loss 0.35422590374946594
Глава 2
72
Первая программа в JAX
Training set accuracy 0.9921666979789734
Test set accuracy 0.9811301827430725
Epoch 24 in 32.60 sec
Eval in 6.37 sec
Training set loss 0.354021817445755
Training set accuracy 0.9924833178520203
Test set accuracy 0.9812300205230713
Шаг 8 завершен, и мы натренировали свою первую нейронную
сеть с помощью JAX. Похоже, что все работает, и решена задача клас
сификации рукописных цифр из набора MNIST с точностью 98,12 %.
Выглядит неплохо.
Созданное решение требует более 30 с на каждую эпоху и допол
нительные 6 с на оценочный прогон каждой эпохи. Это быстро или
медленно? Проверим, можно ли улучшить результат с помощью JITкомпиляции, т. е. выполним шаг 9 нашей общей процедуры.
2.7
JIT: компиляция кода для ускорения
его выполнения
Мы только что реализовали полноценную нейронную сеть для клас
сификации рукописных цифр. Возможно, она даже воспользуется
графическим процессором (GPU), если видеокарта с таким процес
сором установлена в вашем компьютере, поскольку по умолчанию
все тензоры размещаются в GPU. И все же мы можем сделать полу
ченное решение еще более быстрым. В предыдущей версии не ис
пользовалась JIT-компиляция и ускорение, предоставляемое XLA.
Так давайте сделаем это.
Обеспечить компиляцию существующих функций просто. Можно
воспользоваться трансформацией функции jit() или аннотацией @
jit. Сейчас мы применим второй способ.
Компилируем две функции с наибольшим потреблением ресур
сов – update() и batch_accuracy(). Нужно всего лишь добавить анно
тацию @jit перед определениями этих функций.
Листинг 2.14
Добавление JIT-компиляции в код
from jax import jit
❶
@jit
❷
def update(params, x, y, epoch_number):
loss_value, grads = value_and_grad(loss)(params, x, y)
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)], loss_value
JIT: компиляция кода для ускорения его выполнения
73
@jit
def batch_accuracy(params, images, targets):
images = jnp.reshape(images, (len(images), NUM_PIXELS))
predicted_class = jnp.argmax(batched_predict(params, images), axis=1)
return jnp.mean(predicted_class == targets)
❶ Импорт трансформации jit().
❷ Использование трансформации jit() как аннотации функции.
Теперь и шаг 9 завершен.
Здесь мы пропускаем шаг 10, потому что для такой простой зада
чи распределенная тренировка не требуется. Существует множество
способов распределения вычислений в JAX, и это тема глав 7 и 8.
После повторной инициализации параметров нейронной сети
и перезапуска тренировочного цикла получаем следующие резуль
таты:
Epoch 0 in 2.15 sec
Eval in 2.52 sec
Training set loss 0.41040700674057007
Training set accuracy 0.9299499988555908
Test set accuracy 0.931010365486145
Epoch 1 in 1.68 sec
Eval in 2.06 sec
Training set loss 0.37730318307876587
Training set accuracy 0.9500166773796082
Test set accuracy 0.9497803449630737
Epoch 2 in 1.69 sec
Eval in 2.01 sec
Training set loss 0.3708733022212982
Training set accuracy 0.9603500366210938
Test set accuracy 0.9593650102615356
Epoch 3 in 1.67 sec
...
Epoch 23 in 1.69 sec
Eval in 2.07 sec
Training set loss 0.35422590374946594
Training set accuracy 0.9921666979789734
Test set accuracy 0.9811301827430725
Epoch 24 in 1.69 sec
Eval in 2.06 sec
Training set loss 0.3540217876434326
Training set accuracy 0.9924833178520203
Test set accuracy 0.9812300205230713
Качество то же самое, но скорость заметно увеличилась. Для эпо
хи потребовалось около 1,7 с вместо 32,6 с в предыдущей версии,
а оценочный прогон занял почти 2 с вместо 6,3 с. Это весьма суще
ственное улучшение.
Глава 2
74
Первая программа в JAX
Вероятно, вы также заметили, что первые итерации выполнялись
дольше, чем последующие. Поскольку компиляция происходит во
время первого прогона функции, этот этап выполняется медленнее.
Последующие прогоны используют скомпилированную функцию,
поэтому они быстрее. Более подробно JIT-компиляция рассматри
вается в главе 5.
Шаг 11 завершен, и мы почти решили поставленную задачу ма
шинного обучения.
2.8
Сохранение и развертывание модели
После создания модели необходимо сразу же сохранить ее, а затем
в зависимости от конкретного варианта использования можно либо
развернуть модель в какой-либо рабочей среде, либо просто загру
жать ее по мере необходимости и использовать для текущих вычисле
ний, когда не требуется наличие производственной инфраструктуры.
JAX не предоставляет каких-либо специальных инструменталь
ных средств для сохранения модели, потому что с технической точки
зрения JAX ничего не знает о ней. Это просто фреймворк для вы
числения тензоров с высокой производительностью. Самый простой
способ сохранения модели – применение стандартных средств язы
ка Python, например pickle. Но использование pickle не является
правильным решением с точки зрения обеспечения безопасности,
поэтому лучше поискать более надежные варианты. Пакет safetensors (https://github.com/huggingface/safetensors) предлагает более
надежный способ сохранения.
В листинге 2.15 демонстрируется сохранение и восстановление
параметров модели с использованием встроенной в Python струк
туры данных под названием pytree, которой полностью посвящена
глава 10.
Листинг 2.15
Сохранение и загрузка параметров модели
import pickle
model_weights_file = 'mlp_weights.pickle'
❶
❷
with open(model_weights_file, 'wb') as file:
pickle.dump(params, file)
❸
with open(model_weights_file, 'rb') as file:
restored_params = pickle.load(file)
❹
❶
❷
❸
❹
Импорт модуля pickle.
Имя файла для сохраняемых параметров.
Сохранение параметров модели.
Восстановление параметров модели.
Сохранение и развертывание модели
75
Чтобы воспользоваться восстановленной моделью для получения
результатов в отдельной программе, необходимо создать еще и ко
пии кода функции predict() или batched_predict(), так как файл со
хранения содержит только структуру с весами модели.
Библиотеки более высокого уровня поддержки нейронных сетей
на основе JAX, такие как Flax или Equinox, содержат собственные
концепции модели и предоставляют встроенные инструментальные
средства для сериализации/десериализации модели. Эта тема об
суждается в главе 11, где используется библиотека Flax.
Итак, шаг 12 завершен.
Для выполнения шага 13 рассмотрим только загрузку модели из
сохраненного состояния в некоторой контрольной точке. Если тре
буется производственная инфраструктура, а это отдельная тема, то
рекомендуется начать с изучения инструмента JAX2TF (https://www.
tensorfow.org/guide/jax2tf), предоставляющего простой способ пре
образования модели JAX в TensorFlow SavedModel. Далее вы можете
выполнять множество разнообразных операций, например:
выполнение логического вывода (inference) на сервере, исполь
зуя TensorFlow (TF) Serving, на устройстве, применяя TFLite, или
в веб-среде с помощью TensorFlow.js;
более точную настройку (fine-tuning), продолжая тренировку
модели, ранее тренированной в JAX, в TensorFlow с собствен
ными имеющимися тренировочными данными и параметрами
настройки;
слияние (fusion), объединяя части модели, которые ранее были
тренированы с использованием JAX, с компонентами, трениро
ванными с применением TensorFlow.
Еще один проект компании Google под названием Saxml, или Sax
(https://github.com/google/saxml), представляет собой эксперимен
тальную систему, обслуживающую JAX, Paxml (фреймворк машин
ного обучения на основе JAX для тренировки крупномасштабных
моделей, таких как LLM; см. главу 11) и моделей PyTorch для логиче
ских выводов. Ячейка Sax (также известная как Sax-кластер) состоит
из сервера администрирования и exchange model группы серверов
моделей. Сервер администрирования следит за серверами моделей,
распределяет публикуемые модели по серверам моделей для обслу
живания и помогает клиентам размещать серверы моделей, обслу
живающие конкретные публикуемые модели. Недавно был выпущен
релиз 1.0.0 этого проекта.
Некоторые инструментальные средства работают в противопо
ложном направлении, позволяя использовать JAX для моделей, соз
данных в других фреймворках.
В первую очередь следует отметить экспериментальную библио
теку TF2JAX компании DeepMind (https://github.com/deepmind/tf
2jax) для преобразования функций/графов TensorFlow в функции
76
Глава 2
Первая программа в JAX
JAX, что позволяет повторно использовать существующие модели
TensorFlow и точно настраивать их в кодовых базах JAX.
Существует инструментальный комплект JAX ONNX Runtime
(https://github.com/google/jaxonnxruntime), поддерживающий бес
проблемное выполнение открытых обменных моделей нейросетей
с использованием JAX в качестве внутреннего компонента.
Также существуют фреймворки, унифицирующие различные
внутренние компоненты, обычно TensorFlow, JAX, PyTorch и NumPy.
Среди них можно особо выделить Keras 3 (https://keras.io/) и Ivy
(https://github.com/ivy-llc/ivy). Например, Ivy способен перекомпи
лировать код модели из PyTorch в JAX (https://mng.bz/d6Az).
Рассмотрим несколько подробнее общие отличия JAX от осталь
ных фреймворков, а также попробуем понять, почему чистый функ
циональный подход JAX в совокупности с компонуемыми трансфор
мациями функций настолько важен.
2.9
Чистые функции и компонуемые
трансформации: почему они так важны
Мы создали и натренировали свою первую нейронную сеть с исполь
зованием JAX. В ходе этого процесса особо отмечались некоторые
существенные различия между JAX и уже ставшими более широко
распространенными фреймворками, такими как PyTorch и Tensor
Flow. Основу этих различий составляет функциональный подход,
принятый в JAX.
Я уже неоднократно отмечал, что функции JAX обязательно долж
ны быть чистыми. Это означает, что их поведение определяется
только их входными данными, т. е. одни и те же входные данные
всегда должны производить неизменно одинаковый результат на
выходе. Не допускается наличие какого-либо внутреннего состоя
ния, воздействующего на вычисления. Также запрещены какие бы
то ни было побочные эффекты.
Существует множество причин, по которым чистые функции наи
более предпочтительны. Среди прочих можно отметить простое рас
параллеливание, кеширование и возможность создания функцио
нальных композиций, например jit(vmap(grad(some_function))).
К тому же отлаживать чистые функции проще.
Одно отмеченное ранее весьма важное различие касается случай
ных чисел. Генераторы случайных чисел (ГСЧ) NumPy не являются
чистыми, поскольку имеют внутреннее состояние. Собственные ГСЧ
JAX – явно чистые функции. Состояние передается в функцию, кото
рая требует случайности. Поэтому, принимая некоторое состояние,
вы всегда получаете в результате одни и те же «случайные» числа.
Упражнение 2.1
77
Поэтому будьте осторожны и внимательны. Более подробно ГСЧ об
суждаются в главе 9.
Еще одно чрезвычайно важное различие заключается в том, что
параметры нейронной сети не скрыты внутри некоторого объекта,
а всегда передаются в явной форме. Многие вычисления для ней
ронных сетей структурированы следующим образом: сначала вы
генерируете или инициализируете требуемые параметры, затем
передаете их в функцию, использующую конкретные значения па
раметров для вычислений. Этот паттерн мы увидим и в высокоуров
невых библиотеках поддержки нейросетей на основе JAX, таких как
Flax и Equinox. Градиенты также вычисляются и применяются явно,
без каких-либо скрытых «магических» действий.
Параметры нейронной сети становятся сами по себе отдельной
сущностью. Такая структура предоставляет вам гораздо больше сво
боды во всем, что вы делаете. Можно реализовать специализирован
ные обновления, с легкостью сохранять и восстанавливать их, а так
же создавать разнообразные функциональные композиции.
Отсутствие побочных эффектов особенно важно при компиля
ции функций с применением JIT. Если игнорировать чистоту при
компиляции jit() и кешировании функции, то можно получить
совершенно неожиданные результаты. Если на поведение функ
ции воздействует некоторое состояние или функция создает ка
кие-либо побочные эффекты, то скомпилированная версия может
сохранить вычисления, выполненные во время первого прогона,
и воспроизводить их при последующих вызовах, а это, вероятнее
всего, совсем не то, что вам нужно. Об этом мы подробно погово
рим в главе 5.
Поздравляю! Мы завершили работу. Отдельные части проекта
можно заменить на модули из экосистемы даже в таком простейшем
проекте. Если вы работаете над более продвинутыми темами в об
ласти машинного обучения, такими как обучение с подкреплени
ем, графы нейросетей, метаобучение или эволюционные вычисле
ния, то вы наверняка добавите более специализированные модули
из экосистемы JAX и/или замените некоторые части общей схемы,
предложенной в этой главе.
На этом и закончим. Теперь вы готовы к более глубокому погру
жению в ядро JAX.
Упражнение 2.1
1 Измените
функцию predict() так, чтобы она принимала список
функций активации для скрытого слоя.
2 Измените архитектуру нейронной сети и/или процесса трениров
ки для улучшения качества классификации. Подсказка: изменяй
Глава 2
78
Первая программа в JAX
те архитектуру, функции активации, скорость обучения, оптими
затор, что-то еще.
3 Реализуйте другой конвейер машинного обучения, например ра
ботающий с различными типами данных (скажем, классификатор
эмоционального тона сообщений в соцсети на основе предложен
ного набора данных: https://mng.bz/rV7E).
Резюме
В JAX нет собственного загрузчика данных; можно воспользовать
ся внешним загрузчиком из PyTorch или TensorFlow.
В JAX параметры нейронной сети обычно передаются как внеш
ние параметры в функцию, выполняющую все вычисления, а не
хранятся внутри какого-либо объекта, как это принято в Tensor
Flow/ PyTorch.
Параметры модели хранятся во встроенной в язык Python струк
туре данных под названием pytree.
Генераторы случайных чисел в JAX не имеют сохраняемого состоя
ния, поэтому необходимо предоставлять им внешнее состояние
(PRNGKey).
Трансформация vmap() выполняет преобразование функции с оди
ночным элементом входных данных в функцию, работающую с па
кетом данных.
Градиент функции можно вычислять с помощью функции grad().
Если требуется значение и градиент функции, то можно восполь
зоваться функцией value_and_grad().
Трансформация jit() компилирует функцию с применением ком
пилятора линейной алгебры XLA и создает оптимизированный
код, способный выполняться на CPU, GPU или TPU.
Можно с легкостью сохранять и загружать веса (весовые коэффи
циенты) модели, используя стандартные библиотеки Python, на
пример pickle, или более безопасные модули, такие как safetensors компании Hugging Face.
Можно использовать экосистему TensorFlow для развертывания
моделей JAX с помощью пакета JAX2TF.
Необходимо использовать чистые функции без внутреннего со
стояния и побочных эффектов, чтобы все трансформации работа
ли корректно.
Часть II
Ядро JAX
В
части II мы глубже погружаемся во внутренние механизмы JAX
и изучаем функциональные средства ядра, которые делают JAX не
вероятно мощным инструментом для глубокого обучения и науч
ных вычислений. Эта часть включает восемь глав, в которых пред
ставлено подробнейшее описание возможностей JAX – от работы
с массивами и вычисления градиентов до компиляции, векториза
ции и распараллеливания кода. Каждая глава посвящена одному из
основополагающих аспектов JAX и содержит практические примеры
и обстоятельные всесторонние обсуждения, помогающие укрепить
полученные знания и практические навыки.
Главы с 3 по 10 предназначены для скрупулезного изучения всех
тонкостей и хитроумных особенностей JAX, обеспечивая овладение
в полной мере их самыми мощными функциональными возможно
стями. Мы начнем с освоения работы с массивами (глава 3), чтобы
понять все нюансы, отличающие JAX от NumPy, и узнать, как исполь
зовать эти различия с пользой для себя. Потом вы перейдете к под
робному изучению вычисления градиентов (глава 4) с применением
средств автоматического дифференцирования JAX для упрощения
и ускорения тренировки нейронных сетей. Главы 5 и 6 познакомят
вас с just-in-time-компиляцией и автоматической векторизацией
в JAX с описанием стратегий, позволяющих существенно улучшить
производительность.
Но на этом путь через ядро JAX не заканчивается. Вы узнаете о том,
как распараллеливать свои вычисления и использовать сегментиро
вание (sharding) тензоров (главы 7 и 8) для масштабирования при
ложений с распределением их по нескольким устройствам – это
чрезвычайно важный навык для управления крупномасштабными
80
Ядро JAX
моделями и наборами данных. Изучение случайных чисел в JAX
(глава 9) позволит понять функциональный подход к случайности
и обеспечит воспроизводимость и эффективность стохастических
(вероятностных) операций. Глава 10, посвященная работе с pytree,
научит вас аккуратно управлять сложными структурами данных,
расширяя возможности создания более сложных моделей и алгорит
мов с помощью JAX.
После завершения изучения части II вы будете в полной мере
понимать функциональность ядра JAX и приобретете знания для
работы с продвинутыми проектами глубокого обучения. Эта часть
чрезвычайно важна для всех желающих освоить весь потенциал JAX
и применять его в исследованиях или в промышленных приложе
ниях.
3
Работа с массивами
Темы главы:
работа с массивами NumPy;
работа с массивами JAX на CPU/GPU/TPU;
адаптация кода к различиям между массивами NumPy
и JAX;
использование интерфейсов высокого и низкого уровня:
jax.numpy и jax.lax.
В предыдущей главе мы разработали простую нейронную сеть на
основе JAX. В этой главе мы начинаем более глубоко изучать ядро
JAX и начнем с массивов (или тензоров – мы будем использовать эти
термины как взаимозаменяемые).
Тензор, или многомерный массив, – это основная структура дан
ных во фреймворках глубокого обучения и научных вычислений.
В каждой программе используется некоторая форма тензора – одно
мерный вектор, двумерная матрица или массив с бóльшим числом
измерений. Изображения рукописных цифр из предыдущей главы,
промежуточные активизации и итоговые прогнозы нейросети – все
это тензоры. NumPy предоставляет тип numpy.ndarray; в JAX имеется
тип Array (ранее известный как DeviceArray).
82
Глава 3
Работа с массивами
Массивы NumPy (тип numpy.ndarray) и их API стали промышленным
стандартом де-факто, принятым во многих других фреймворках. JAX
предоставляет NumPy-совместимый (в основном) API, поэтому пере
ход от NumPy к JAX не должен создавать каких-либо затруднений, и во
многих случаях вам даже не придется что-то менять в своем коде, за
исключением инструкции импортирования. Тем не менее существу
ют некоторые различия, и мы уделим им особое внимание.
В этой главе рассматриваются массивы и соответствующие опе
рации с ними в NumPy и JAX. Мы займемся реализацией примера
обработки изображений с использованием матричных фильтров.
Сначала мы будем работать с чистым NumPy, используя его абстрак
цию для многомерных массивов numpy.ndarray. Затем перейдем
к JAX, введем в употребление структуру данных Array и объясним
все нюансы, специфические для JAX, особенно операции, связанные
с устройствами. В конце главы рассматриваются различия между API
NumPy и JAX.
В главе много исходного кода, но в мои планы не входило наме
рение перегружать вас его объемом, поэтому код представлен не
большими фрагментами и подробно аннотирован. Изобилие ис
ходного кода вызвано тем, что я хотел бы продемонстрировать все
самые важные концепции. При изучении исходного кода вы полу
чаете гораздо больше полезной информации по сравнению с чте
нием только обычного текста. Я уверен, что программисты намного
быстрее схватывают основные идеи и принципы, когда их описание
сопровождается исходным кодом. Кроме того, хотелось бы, чтобы вы
могли следить за идеями, излагаемыми в тексте, даже без доступа
к компьютеру (хотя гораздо лучше, если у вас есть возможность по
экспериментировать с кодом во время чтения главы).
3.1
Обработка изображений с использованием
массивов NumPy
Начнем с реальной задачи обработки изображений с использовани
ем чистого NumPy. Изображения представляют собой превосходные
наглядные примеры тензоров, или многомерных массивов, поэто
му работа с изображениями обеспечит полное понимание тензоров
и некоторых важных операций с ними более простым способом.
Предположим, что имеется набор фотографий, требующих об
работки. На некоторых фото слишком много «лишнего» простран
ства, которое нужно обрезать, другие содержат искажения и помехи
(«шум»), которые необходимо удалить. Многие изображения имеют
хорошее качество, но вам хотелось бы применить к ним некоторые
художественные эффекты. Для упрощения общей задачи сосредото
Обработка изображений с использованием массивов NumPy
83
чимся только на устранении помех и дефектов («шума») в изображе
ниях, как показано на рис. 3.1.
Изображение с помехами
Восстановленное изображение
Рис. 3.1 Пример обработки изображения, которое необходимо реализовать
Чтобы реализовать такой способ обработки, необходимо загру
зить изображения в некоторую структуру данных и написать функ
ции их обработки. Предположим, что имеется фотография (я выбрал
одну из статуй котов, созданных Фернандо Ботеро (Fernando Botero)
в архитектурно-монументальном комплексе «Каскад» в Ереване
(Армения)), на которой во время съемки появились посторонние
«шумовые» помехи. Необходимо удалить эти помехи и, возможно,
применить к изображению некоторые художественные фильтры.
Разумеется, существуют многочисленные инструменты обработки
изображений, но мы предпочитаем реализовать собственный кон
вейер в демонстрационных (и учебных) целях. Вероятно, вы также
захотите в будущем добавить в этот конвейер еще один этап обра
ботки нейросетью, например реализацию сверхвысокого разреше
ния или дополнительные художественные фильтры. Применяя JAX,
это сделать достаточно просто.
Пример обработки изображений будет состоять из нескольких
шагов:
1 загрузка
фотографии из файла в тензор, размещенный в памя
ти (здесь: массив NumPy);
2 предварительная обработка изображения, если это необходи
мо, и преобразование значений пикселов в числа с плавающей
точкой;
3 генерация зашумленной версии фотографии для имитации по
мех, созданных камерой при съемке;
4 фильтрация изображения для снижения уровня помех («шума»)
и увеличения резкости;
5 сохранение обработанного изображения в файле.
Глава 3
84
Работа с массивами
Каждый шаг включает выполнение различных операций с тензо
ром и помогает вам овладеть основными навыками работы с тензо
рами. Эти навыки окажутся полезными в дальнейшей деятельности
в области глубокого обучения, поскольку они необходимы всегда
и везде: при создании конвейера обработки и пополнении данных,
при инициализации весов модели, при применении нейросетей
к данным, при сохранении весов или полученных результатов моде
ли – тензоры используются везде.
Сначала мы реализуем пример с применением чистого NumPy.
Затем переключимся на JAX (спойлер: изменив всего лишь пару ин
струкций импортирования). Начнем с загрузки изображения.
3.1.1
Загрузка изображения в массив NumPy
В первую очередь необходимо загрузить изображения. Изобра
жения представляют собой превосходные примеры многомерных
объектов. Они имеют два пространственных измерения (ширина
и высота) и обычно еще одно измерение с цветовыми каналами
(как правило, red (красный), green (зеленый), blue (синий) и иногда
альфа-канал). Таким образом, изображения естественным образом
представляются в форме многомерных массивов, и массив NumPy
является вполне подходящей структурой для хранения изображе
ний в памяти компьютера. Общая схема такого массива показана на
рис. 3.2.
Высота
Цвет
Ширина
Рис. 3.2 Тензор изображения имеет три измерения:
ширина, высота и цветовые каналы
Обработка изображений с использованием массивов NumPy
85
Соответствующий тензор изображения также имеет три измере
ния: ширину, высоту и цветовые каналы. В черно-белых изображе
ниях или в изображениях с оттенками серого цвета измерение цве
товых каналов может отсутствовать. В системе индексации NumPy
первое измерение соответствует строкам, второе – столбцам. Цве
товое измерение может быть размещено до или после простран
ственных измерений. С чисто технической точки зрения цветовое
измерение можно также разместить между измерениями высоты
и ширины, но это не имеет никакого смысла. В scikit-image каналы
располагаются в порядке RGB (reg (красный) первый, потом green
(зеленый), за ним blue (синий)). В других библиотеках, например
OpenCV, может использоваться другая схема, скажем BGR, с обрат
ным порядком цветовых каналов.
Загрузим изображение и рассмотрим повнимательнее его свой
ства. Я поместил файл изображения с именем The_Cat.jpg в текущий
каталог. Вы можете разместить в текущем каталоге файл с любым
другим изображением, но не забудьте изменить имя файла в исход
ном коде.
Мы воспользуемся широко распространенной библиотекой sci
kit-image для загрузки и сохранения изображений. В среде Google
Colab эта библиотека уже установлена предварительно, но при ее от
сутствии в вашей системе установка выполняется по инструкциям
https://scikit-image.org/docs/stable/user_guide/install.html. Исходный
код примеров и файлы изображений можно найти в репозитории
GitHub этой книги: https://github.com/che-shr-cat/JAX-in-Action/blob/
main/Chapter-3/JAX_in_Action_Chapter_3_Image_Processing.ipynb. Код
в листинге 3.1 загружает изображение в память.
Листинг 3.1
Загрузка изображения в массив NumPy
import numpy as np
from scipy.signal import convolve2d
from matplotlib import pyplot as plt
from skimage.io import imread, imsave
from skimage.util import (img_as_float32,
img_as_ubyte, random_noise)
img = imread('The_Cat.jpg')
from matplotlib import pyplot as plt
plt.figure(figsize = (6, 10))
plt.imshow(img)
❶
❶
❶
❶
❶
❷
❸
❸
❸
❶ Импорт всех компонентов, необходимых здесь и в дальнейшем.
❷ Загрузка изображения.
❸ Вывод изображения.
Глава 3
86
Работа с массивами
Выведенное изображение должно выглядеть приблизительно
так, как показано на рис. 3.3. В зависимости от разрешения вашего
экрана, возможно, потребуется изменение размеров изображения.
Это можно с легкостью сделать, используя параметр figsize, пред
ставляющий собой кортеж (width, height), где значения заданы
в дюймах.
Рис. 3.3 Цветное изображение (в печатной версии показано в оттенках
серого) представлено как трехмерный массив с измерениями: высота,
ширина, цвет
Шаг 1 завершен, и мы можем проверить тип тензора изображения:
type(img)
>>> numpy.ndarray
❶
❶ Это тип массива ndarray NumPy.
Тензоры обычно описываются по их форме кортежем с количест
вом элементов, равным рангу тензора (т. е. числу измерений). Каж
дый элемент кортежа представляет количество позиций индекса по
Обработка изображений с использованием массивов NumPy
87
соответствующему измерению. На этот кортеж ссылается свойство
shape. Число измерений можно узнать из свойства ndim.
Например, цветное изображение размером 1024×768 можно пред
ставить такой формой: (768, 1024, 3) или (3, 768, 1024).
При работе с пакетом изображений добавляется новое измерение
пакета, обычно первым с индексом 0.
Пакет из 32 цветных изображений размером 1024×768 представ
ляется в следующей форме: (32, 768, 1024, 3) или (32, 3, 768, 1024):
img.ndim
>>> 3
img.shape
>>> (667, 500, 3)
❶
❷
❶ Тензор изображения содержит три измерения.
❷ Размеры этих измерений: 667 (высота), 500 (ширина) и 3 (цветовые каналы).
Здесь можно видеть, что изображение из рассматриваемого при
мера представлено трехмерным массивом, в котором первое изме
рение – высота (667 пикселов), второе – ширина (500 пикселов), тре
тье – цветовые каналы (3 канала).
Схема размещения тензора в памяти: сравнение NCHW
и NHWC
Тензоры изображений обычно представлены в памяти в двух обобщенных форматах: NCHW и NHWC. Эти буквы верхнего регистра (строчные)
обозначают смысл осей тензора, где N соответствует размерности пакета, C – размерности канала (Channel), H – высоте (Height), W – ширине
(Width). Описанный таким способом тензор содержит пакет, состоящий
из N изображений с C цветовыми каналами, при этом каждое изображение имеет высоту H и ширину W. При наличии в пакете изображений
разного размера возникает проблема, но специальные структуры данных, такие как тензор с элементами разной длины, или невыровненный тензор (ragged tensor), способны хранить подобные изображения.
Другой вариант: выравнивание объектов различной длины до одного
размера с помощью некоторого элемента-заполнителя.
Фреймворки и библиотеки предпочитают использовать разные форматы. JAX (https://docs.jax.dev/en/latest/_autosummary/jax.lax.conv_
general_dilated.html) и PyTorch (https://discuss.pytorch.org/t/whydoes-pytorch-prefer-using-nchw/83637/4) используют NCHW по умолчанию. В TensorFlow (https://www.tensorflow.org/api_docs/python/tf/
nn/conv2d), Flax (https://flax-linen.readthedocs.io/en/latest/api_refer
ence/flax.linen/layers.html#flax.linen.Conv) и Haiku (https://dm-hai
ku.readthedocs.io/en/latest/api.html#haiku.Conv2D) принят формат
NHWC, и почти во всех библиотеках имеется функция для преобразова-
Глава 3
88
Работа с массивами
ния между этими форматами или параметр, определяющий, какой тип
тензора изображения должен передаваться в функцию.
С точки зрения математики описанные выше представления равнозначны, но на практике они могут быть различными. Например, свертки (convolutions), реализованные в NVIDIA Tensor Cores, требуют схемы размещения NHWC и работают быстрее, если во входных тензорах
применяется формат NHWC (https://docs.nvidia.com/deeplearning/
performance/dl-performance-convolutional/index.html#tensorlayout).
Grapper, принятая по умолчанию система оптимизации графов в механизме времени выполнения TensorFlow (https://www.tensorflow.org/
guide/graph_optimization?hl=ru), может автоматически преобразовывать NHWC в NCHW во время процедуры оптимизации схемы размещения (https://research.google/pubs/pub48051/).
Кроме того, существуют два полезных свойства, связанных с раз
мером. Свойство size возвращает количество элементов в тензоре.
Его значение равно произведению измерений массива. Свойство
nbytes возвращает общее количество байтов, потребляемых эле
ментами массива. В это количество не включается объем памяти,
потребляемой атрибутами объекта массива, не являющимися соб
ственно элементами. Для тензора, включающего значения типа
uint8, эти свойства возвращают одинаковое число, так как для каж
дого элемента требуется только один байт.
img.size
>>> 1000500
img.nbytes
>>> 1000500
❶
❷
❶ В тензоре изображения содержится 1 000 500 элементов.
❷ Тензор занимает 1 000 500 байт (каждый элемент является однобайтовым значе-
нием).
Для значений типа float32 число занятых байтов будет в четы
ре раза больше, чем количество элементов. Итак, шаг 1 завершен,
изображение загружено. Теперь необходимо выполнить некоторую
предварительную обработку.
3.1.2
Выполнение простых операций предварительной
обработки с изображением
Мы находимся на шаге 2, где выполняется предварительная обра
ботка изображения. А зачем вообще нужен этот шаг?
Обработка изображений с использованием массивов NumPy
89
Во-первых, может потребоваться обрезка изображения для уда
ления некоторых несущественных подробностей у краев. В данном
случае в этом нет необходимости (изображение вполне подходя
щее), но я просто хочу рассказать вам подробнее о мощной функции
под названием slicing (вырезка), которую можно применить непо
средственно для обрезки изображения (часто называемую кадриро
ванием). С помощью вырезки можно выбрать конкретные элементы
тензора по каждой оси. Например, вы выбираете определенную пря
моугольную подобласть изображения или берете только выбранные
цветовые каналы.
Листинг 3.2
Вырезка из массивов NumPy
cat_face = img[80:220, 190:330, 1]
cat_face.shape
>>> (140, 140)
plt.figure(figsize = (3,4))
plt.imshow(cat_face, cmap='gray')
❶
❷
❶ Вырезаются пикселы со строки 80 по строку 220 (не включая саму эту строку)
по высоте, столбцы с 190 по 330 и второй цветовой канал (с индексом 1, так как
нумерация индексов начинается с 0).
❷ В результате получен тензор высотой в 140 пикселов и шириной в 140 пикселов.
Код в листинге 3.2 выбирает пикселы, относящиеся к голове кота,
только из канала зеленого цвета изображения (это средний канал
с индексом 1). Далее мы выводим полученное изображение в оттен
ках серого цвета, поскольку оно содержит информацию только об
одном цветовом канале. При желании можно выбрать другую цве
товую палитру. Полученное изображение показано на рис. 3.4. Также
можно с легкостью реализовать любое другое простое преобразова
ние, например зеркальное отображение, с помощью операции вы
резки из массива.
Рис. 3.4 Обрезанное изображение
с одним цветовым каналом
Глава 3
90
Работа с массивами
Код img = img[:,::-1,:] изменяет порядок пикселов на противопо
ложный по горизонтальному измерению, сохраняя при этом верти
кальную ось и ось каналов. Также можно воспользоваться функцией
flip() из библиотеки NumPy. Вращение выполняется с помощью
функции rot90() с заданным числом поворотов (параметр k=2), как
показано в коде img = np.rot90(img, k=2, axes=(0,1)).
Еще одно важное средство предварительной обработки относит
ся к типам данных. Каждый тензор имеет тип данных, связанный
со всеми его элементами и обозначенный как dtype. Изображения
обычно представлены либо значениями с плавающей точкой в диа
пазоне [0, 1], либо беззнаковыми целыми числами в диапазоне [0,
255]. Значениями в нашем тензоре являются беззнаковые 8-битовые
целые числа (uint8).
В данном случае необходимо преобразовать элементы изображе
ния из типа данных uint8 в тип float32. Для нас проще работать со
значениями с плавающей точкой в диапазоне [0.0, 1.0]. Воспользу
емся функцией img_as_float() из пакета scikit-image.
Листинг 3.3 Преобразование значений пикселов изображения
в числа с плавающей точкой
img.dtype
>>> dtype('uint8')
img = img_as_float32(img)
img.dtype
>>> dtype('float32')
❶
❷
❸
❶ Изображение закодировано с помощью беззнаковых 8-битовых целых чисел.
❷ Преобразование изображения из байтов (беззнаковых 8-битовых целых чисел)
в 32-битовые значения с плавающей точкой.
❸ Тип данных тензора изменен на 'float32’.
Шаг 2 завершен.
3.1.3
Добавление шума в изображение
Шаг 3 необходим только в демонстрационных целях. Чтобы ими
тировать изображение с помехами («зашумленное» изображение),
воспользуемся гауссовым шумом (Gaussian noise), часто появляю
щимся в цифровых камерах при условиях слабого освещения и вы
сокой светочувствительности по стандарту ISO. Применим функцию
random_noise из того же пакета scikit-image.
Обработка изображений с использованием массивов NumPy
Листинг 3.4
91
Генерация зашумленной версии изображения
img_noised = random_noise(img, mode='gaussian')
plt.figure(figsize = (6, 10))
plt.imshow(img_noised)
❶
❶ Мы используем функцию из пакета scikit-image для добавления в изображение
случайных помех (шума) заданного типа.
Код в листинге 3.4 генерирует зашумленную версию исходного
изображения, показанную на рис. 3.5. Вы можете поэксперимен
тировать с другими интересными типами шума, такими как «соль
и перец» (черно-белый; salt-and-pepper) шум или импульсные поме
хи (impulse noise), возникающими как разреженные минимальные
и максимальные значения пикселов.
Рис. 3.5 Зашумленная версия исходного изображения
Шаг 3 завершен. Теперь мы готовы к реализации некоторой более
продвинутой обработки изображения.
92
3.1.4
Глава 3
Работа с массивами
Реализация фильтрации изображения
Шаг 4 является самым главным в рассматриваемом примере обра
ботки изображения. Он состоит из двух подшагов:
1 создание ядра фильтра (более подробно о ядрах см. ниже в при
мечании «КИХ-фильтры и свертка»);
2 применение ядра фильтра к изображению (или собственно вы
полнение шага фильтрации).
Создание ядра фильтра
Фильтр размытия по Гауссу (Gaussian blur filter) удаляет имеющий
ся в рассматриваемом здесь изображении тип шума. Фильтры раз
мытия по Гауссу принадлежат к большому семейству матричных
фильтров, также называемых фильтрами с конечной импульсной
характеристикой (КИХ-фильтров; finite impulse response – FIR) из об
ласти обработки цифровых сигналов (digital signal processing – DSP).
Работу матричных фильтров еще можно наблюдать в приложениях
обработки изображений, например в Photoshop или GIMP.
КИХ-фильтры и свертка
КИХ-фильтр (FIR-фильтр) описывается собственной матрицей, также
называемой ядром (kernel). Матрица содержит веса (весовые коэффициенты), с которыми берутся пикселы изображения, когда фильтр
перемещается («скользит») по изображению. Во время каждого шага
все пикселы окна, через которое ядро «смотрит» на часть изображения,
умножаются на соответствующие веса ядра. Затем полученные произведения суммируются, чтобы получить единственное число – итоговая
интенсивность для любого пиксела в выводе (отфильтрованного) изображения. Такая операция, принимающая некоторый сигнал и ядро
(другой сигнал) и производящая отфильтрованный сигнал, называется
сверткой (convolution). Процесс свертки наглядно показан на рис. 3.6.
Свертка может выполняться по любому числу измерений. Для одномерных входных данных это одномерная свертка, для двумерных (таких
как изображения) – двумерная свертка и т. д. Эта операция интенсивно
используется в сверточных (конволюционных) нейронных сетях (convolutional neural networks – CNN). Мы будем отдельно применять двумерную свертку к каждому каналу изображения. В CNN все каналы изображения обычно обрабатываются одновременно с применением ядра
того же измерения канала, что и для изображения.
Если вы хотите узнать больше о цифровых фильтрах и операциях свертки, рекомендуется прочитать превосходную книгу о методах обработки
цифровых сигналов: https://dspguide.com/.
Обработка изображений с использованием массивов NumPy
93
Рецептивное
поле
Фильтр
свертки (3×3)
Целевой
пиксел
Изображение
после свертки
Рис. 3.6 КИХ-фильтр реализуется через операцию свертки (изображение
взято из книги Мохамеда Эльгенди (Mohamed Elgendy) «Deep Learning for
Vision Systems» (изд. Manning), 2020 г.)
Начнем с простого фильтра размытия (не гауссова) для демонстра
ции процесса создания ядра. Затем создадим ядро для фильтра гауссо
ва размытия, который применим для удаления помех из изображения.
Фильтр размытия содержит матрицу равных значений, означаю
щую, что каждый пиксел, располагающийся по соседству с целевым
пикселом, берется с одинаковым весом. Это равнозначно усреднению
всех значений внутри рецептивного поля ядра. Возможно, вы слыша
ли, что такой механизм называют простым фильтром скользящего
среднего (moving average filter). Код в листинге 3.5 демонстрирует
процедуру создания ядра фильтра размытия размером 5×5 пикселов.
Листинг 3.5
Матрица (или ядро) для простого фильтра размытия
kernel_blur = np.ones((5,5))
kernel_blur /= np.sum(kernel_blur)
kernel_blur
❶
❷
>>> array([[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04]])
❸
❸
❸
❸
❸
❶ Генерация матрицы 5×5 из единиц.
❷ Деление каждого элемента матрицы на сумму элементов.
❸ Итоговый массив содержит одинаковые числа, которые в сумме дают единицу.
Глава 3
94
Работа с массивами
Фильтр гауссова размытия представляет собой более сложную
версию простого фильтра размытия; его матрица содержит раз
личные значения, при этом более высокие значения располагаются
ближе к центру. Широко известная функция Гаусса генерирует ядро
гауссова размытия.
Листинг 3.6
Ядро гауссова размытия
def gaussian_kernel(kernel_size, sigma=1.0, mu=0.0):
""" A function to generate Gaussian 2D kernel """
# Функция генерации двумерного ядра гауссова размытия.
❶
center = kernel_size // 2
x, y = np.mgrid[
❷
-center : kernel_size - center,
❷
-center : kernel_size - center]
❷
d = np.sqrt(np.square(x) + np.square(y))
koeff = 1 / (2 * np.pi * np.square(sigma))
kernel = koeff * np.exp(-np.square(d-mu) /
(2 * np.square(sigma)))
❸
return kernel
kernel_gauss = gaussian_kernel(5)
kernel_gauss
>>> array([[0.00291502, 0.01306423,
0.01306423, 0.00291502],
>>>
[0.01306423, 0.05854983,
0.05854983, 0.01306423],
>>>
[0.02153928, 0.09653235,
0.09653235, 0.02153928],
>>>
[0.01306423, 0.05854983,
0.05854983, 0.01306423],
>>>
[0.00291502, 0.01306423,
0.01306423, 0.00291502]])
❶
❷
❸
❹
0.02153928,
0.09653235,
0.15915494,
0.09653235,
0.02153928,
Поиск центральной позиции ядра.
Генерация значений сетки X и Y.
Генерация коэффициентов ядра по заданной формуле.
Итоговое ядро.
❹
❹
❹
❹
❹
Теперь необходимо реализовать функцию для применения соз
данного фильтра к изображению.
Применение ядра фильтра к изображению
Это самая сложная часть главы. На рис. 3.7 изображена схема про
цесса применения фильтра к изображению.
Фильтр применяется отдельно к каждому цветовому каналу. Функ
ция должна применить двумерную свертку с ядром фильтра внутри
каждого цветового канала. Кроме того, усекаются полученные в ре
95
Обработка изображений с использованием массивов NumPy
зультате значения, чтобы ограничить диапазон до [0.0, 1.0]. Затем
обработанные цветовые каналы объединяются для формирования
обработанного изображения. Функция предполагает, что измерение
каналов является последним в тензоре изображения. Преобразова
ние схемы на рис. 3.7 в код теперь становится простым и понятным.
Входное
изображение
Разделение
по каналам
Канал
красного
цвета
Применение
фильтра
Усечение
значений
Канал
зеленого
цвета
Применение
фильтра
Усечение
значений
Канал
синего
цвета
Применение
фильтра
Усечение
значений
Выходное
изображение
Объединение
каналов
Ядро
фильтра
Рис. 3.7 Схема процесса применения фильтра к изображению
Листинг 3.7
Функция для применения фильтра к изображению
def color_convolution(image, kernel):
""" A function to apply a filter to an image"""
# Функция для применения фильтра к изображению.
channels = []
for i in range(3):
color_channel = image[:,:,i]
filtered_channel = convolve2d(
color_channel, kernel, mode="same")
filtered_channel = np.clip(
filtered_channel, 0.0, 1.0)
channels.append(filtered_channel)
final_image = np.stack(channels, axis=2)
return final_image
❶
❷
❸
❹
❶
❷
❸
❹
Извлечение канала с помощью операции вырезки.
Применение фильтра к извлеченному каналу.
Усечение значений до диапазона [0.0, 1.0].
Генерация итогового изображения посредством объединения отфильтрованных
каналов.
Мы готовы применить созданный фильтр к зашумленному изо
бражению.
Листинг 3.8
Фильтрация зашумленного изображения
img_blur = color_convolution(img_noised, kernel_gauss)
plt.figure(figsize = (12,10))
plt.imshow(np.hstack((img_blur, img_noised)))
❶
❶ Применение фильтра гауссова размытия к зашумленному изображению.
96
Глава 3
Работа с массивами
Полученное в результате освобожденное от шума изображение
показано на рис. 3.8. Лучше сравнивать изображения на компью
тере, но можно видеть, что уровень помех существенно снизился,
хотя изображение стало более размытым (нерезким). Этого можно
было ожидать (в конце концов, мы же применили фильтр размы
тия), но хотелось бы сделать изображение более резким, если это
возможно.
Рис. 3.8 Отфильтрованная (слева) и зашумленная (справа) версии изображения
Вполне ожидаемо, что существует матричный фильтр для увели
чения резкости изображений. Его ядро содержит большое положи
тельное число в центре и отрицательные числа по соседству. Цель
этого фильтра – усиление контраста между центральной точкой и ее
соседями.
Нормализация значений ядра обеспечивается делением каждо
го элемента на сумму всех значений. Такая операция ограничива
ет значения допустимым диапазоном после применения фильтра.
Внутри функции применения фильтра имеется функция clip(), но
это исключительная мера, позволяющая усекать каждое значение,
выходящее за пределы допустимого диапазона, до его границы. Бо
лее точная методика нормализации должна сохранять больше ин
формации в обрабатываемом сигнале.
Обработка изображений с использованием массивов NumPy
97
Листинг 3.9 Ядро фильтра увеличения резкости
kernel_sharpen = np.array(
[[-1, -1, -1, -1, -1],
[-1, -1, -1, -1, -1],
[-1, -1, 50, -1, -1],
[-1, -1, -1, -1, -1],
[-1, -1, -1, -1, -1]], dtype=np.float32
)
kernel_sharpen /= np.sum(kernel_sharpen)
kernel_sharpen
❶
❶
❶
❶
❶
❶
❶
❷
>>> array([[-0.03846154, -0.03846154, -0.03846154, -0.03846154,
>>>
[-0.03846154, -0.03846154, -0.03846154, -0.03846154,
>>>
[-0.03846154, -0.03846154, 1.9230769 , -0.03846154,
>>>
[-0.03846154, -0.03846154, -0.03846154, -0.03846154,
>>>
[-0.03846154, -0.03846154, -0.03846154, -0.03846154,
>>>
dtype=float32)
-0.03846154],
-0.03846154],
-0.03846154],
-0.03846154],
-0.03846154]],
❶ Создание ядра с большим положительным значением в центре и небольшими отрицательны-
ми значениями, окружающими его.
❷ Нормализация ядра.
Применим созданный фильтр увеличения резкости к размытому
изображению.
Листинг 3.10
Увеличение резкости изображения
img_restored = color_convolution(
img_blur, kernel_sharpen)
plt.figure(figsize = (12,20))
plt.imshow(np.vstack(
(np.hstack((img, img_noised)),
np.hstack((img_restored, img_blur)))
))
❶
❷
❷
❷
❷
❶ Применение фильтра повышения резкости к размытому изображению.
❷ Вывод всех четырех версий изображения (по часовой стрелке): исходная фото-
графия, зашумленная версия, размытое изображение, резкое или восстановленное изображение.
На рис. 3.9 показаны четыре различных изображения (разме
щенные по часовой стрелке): исходное изображение сверху слева,
зашумленное изображение сверху справа, очищенная от помех раз
мытая версия внизу справа и изображение с увеличением резкости
внизу слева. Мы в определенной степени восстановили резкость изо
бражения и в то же время успешно удалили некоторые имеющиеся
помехи (шум).
98
Глава 3
Работа с массивами
Рис. 3.9 Исходное изображение (сверху слева), зашумленная (сверху справа),
очищенная, но размытая (внизу справа) версии и изображение с увеличенной
резкостью (внизу слева)
Шаг 4 завершен, и мы почти закончили работу.
3.1.5
Сохранение тензора как файла изображения
Заключительный шаг под номером 5 – сохранение итогового изо
бражения. Перед сохранением в файле необходимо «отменить» не
которые операции предварительной обработки, сделанные на шаге 2
99
Массивы в JAX
и относящиеся к преобразованию типов данных. Тогда мы преобра
зовали тензор байтов в тензор чисел с плавающей точкой, упростив
работу с ним. Теперь для сохранения изображения требуется обрат
ное преобразование тензора в байты. Именно эта операция выпол
няется в листинге 3.11.
Листинг 3.11
Сохранение изображения из массива NumPy
image_modified = img_as_ubyte(img_restored)
imsave('The_Cat_modified.jpg', arr=image_modified)
❶
❷
❶ Преобразование значений типа float32 обратно в 8-битовые целые числа.
❷ Сохранение массива как изображения в формате JPEG.
Мы завершили выполнение примера обработки изображения. Вы
можете попробовать применить многие другие интересные фильт
ры, например рельефное (выпуклое) изображение (emboss), обнару
жение и выделение контуров элементов изображения (edge detec
tion) или собственный специализированный фильтр. Я разместил
некоторые фильтры в соответствующем блокноте Colab. Также мож
но объединять несколько ядер фильтров в одно ядро. Но сейчас мы
остановимся на достигнутом и проанализируем, что было сделано.
Мы начали с загрузки изображения и узнали, как реализовать ос
новные операции обработки изображения, такие как обрезка и зер
кальное отображение с помощью операции вырезки. Затем мы
создали зашумленную версию изображения и познакомились с мат
ричными фильтрами, потом выполнили некоторые операции фильт
рации помех и увеличения резкости изображения с применением
матричных фильтров.
Все это было реализовано с использованием библиотеки NumPy,
а теперь наступило время для того, чтобы узнать, что изменилось
с появлением фреймворка JAX.
3.2
Массивы в JAX
Перепишем нашу программу обработки изображения для работы на
основе JAX вместо NumPy. Этот раздел представляет собой общую
схему перевода NumPy-программы в среду JAX.
У нас уже есть работающий код решения, загружающий изобра
жение, добавляющий случайные помехи, создающий цифровой
фильтр, применяющий этот фильтр к зашумленному изображению,
а затем сохраняющий результат обработки. Части загрузки и сохра
нения остаются неизменными, а изменения в основном будут про
исходить в главной части, содержащей цифровые фильтры. Но их
будет не так уж много.
Глава 3
100
3.2.1
Работа с массивами
Переход на NumPy-подобный API JAX
Отличная новость: вы можете заменить всего лишь пару инструкций
импорта, и весь остальной код будет работать с использованием JAX.
Попробуйте сами.
Листинг 3.12 Замена инструкций импорта модулей NumPy
на модули JAX
## NumPy
#import numpy as np
#from scipy.signal import convolve2d
## JAX
import jax.numpy as np
from jax.scipy.signal import convolve2d
❶ Инструкции импорта NumPy и SciPy.
❷ Инструкции импорта JAX.
❶
❶
❷
❷
В JAX имеется NumPy-подобный API, импортируемый из моду
ля jax.numpy. Также существуют некоторые функции более высоко
го уровня из SciPy, заново реализованные в JAX. Модуль jax.scipy
предоставляет не такую богатую функциональность, как вся биб
лиотека SciPy, но используемая нами функция convolve2d() в нем
имеется.
Иногда в JAX обнаруживается отсутствие нужной соответству
ющей функции. Например, мы можем воспользоваться функцией
gaussian_filter() из scipy.ndimage для фильтрации по Гауссу. Такой
функции нет в модуле jax.scipy.ndimage.
В подобных случаях можно продолжать пользоваться функцией
NumPy в среде JAX и включать две инструкции импорта: одну из
NumPy, другую для интерфейса NumPy из JAX. Обычно это делается
так, как показано в листинге 3.13.
Листинг 3.13
Совместное использование NumPy и JAX
## NumPy
import numpy as np
## JAX
import jax.numpy as jnp
❶
❶
❶ Для библиотек используются различные имена, чтобы различать их и выбирать
нужную для использования.
Мы используем функцию из NumPy с префиксом np, а функцию из
JAX с префиксом jnp. Это позволяет предотвратить использование
101
Массивы в JAX
некоторых функциональных средств JAX с функциями NumPy либо
потому, что они реализованы на C++ (Python предоставляет только
привязки), либо потому, что они не являются функционально чис
тыми. Другой вариант: вы можете реализовать собственные новые
функции, как это было сделано для фильтрации по Гауссу.
Если вы выполните наш пример фильтрации изображения, изме
нив инструкции импорта для модулей JAX, то увидите, что код нор
мально работает благодаря наличию NumPy-совместимого API в JAX.
Единственное, на что, возможно, следует обратить внимание, – в тех
местах, где создаются массивы, тип numpy.ndarray будет заменен на
тип Array из JAX (более конкретно: на тип jaxlib.xla_extension.ArrayImpl), например при создании ядра фильтра.
Листинг 3.14 Матрица (или ядро) для простого фильтра размытия
при использовании JAX
kernel_blur = np.ones((5,5))
kernel_blur /= np.sum(kernel_blur)
kernel_blur
>>> Array([[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04],
>>>
[0.04, 0.04, 0.04, 0.04, 0.04]], dtype=float32)
type(kernel_blur)
>>> jaxlib.xla_extension.ArrayImpl
❶ Массив NumPy заменяется на массив JAX.
❷ Теперь для данных явно объявляется тип float32.
❸ Полное имя типа массива JAX.
❶
❶
❶
❷
❸
Также можно видеть, что теперь данные имеют тип float32.
В NumPy в большинстве случаев это был бы тип float64. Все осталь
ное работает как обычно. Мы уделим больше внимания типам дан
ных с плавающей точкой немного позже, в подразделе 3.3.2.
Здесь мы немного отвлечемся от примера обработки изображе
ния и сосредоточимся на том, что представляет собой тип Array. Тип
jax.Array (вместе со своим псевдонимом jax.numpy.ndarray) являет
ся основным типом хранения тензоров, или многомерных массивов
в JAX. Вы будете постоянно работать с массивами типа Array в JAX
так же, как и в NumPy. Понимание свойств этого типа и его отличий
от массивов NumPy заслуживает того, чтобы потратить некоторое
время на его изучение.
102
3.2.2
Глава 3
Работа с массивами
Что такое Array?
Array – это тип, принятый по умолчанию для представления масси
вов в JAX. Он может использовать различные бэкенды – CPU, GPU
и TPU. Array равнозначен типу numpy.ndarray на основе буфера па
мяти на одном устройстве, а также на нескольких устройствах (см.
главу 8). В общем смысле устройство (device) – это некая сущность,
используемая JAX для выполнения вычислений.
DeviceArray и Array
В JAX до версии 0.4.1 реализацией массива по умолчанию был тип DeviceArray. Начиная с версии 0.4.1 JAX перешел на внутреннюю реализацию массива нового типа jax.Array.
В будущем jax.Array останется единственным типом массива в JAX.
Это универсальный тип массива, охватывающий типы DeviceArray,
ShardedDeviceArray и GlobalDeviceArray в JAX. Новый тип Array помогает организовать режим распараллеливания ядра JAX, упрощает
и унифицирует внутренние компоненты JAX, а также позволяет унифицировать JIT (это тема главы 5) и pjit() (см. приложение D). jax.Array
предоставляет распределенные массивы сразу в готовом к применению
виде и позволяет с легкостью использовать автоматизированное распараллеливание.
Если в вашем коде используются более старые типы, то необходимо
перевести их в тип jax.Array в соответствии с инструкциями: https://
docs.jax.dev/en/latest/jax_array_migration.html.
Часто нет необходимости создавать экземпляры объектов Array
вручную (и мы этого не делали). Вы будете создавать через функции
jax.numpy, такие как array(), linspace() и т. п.
Заметное отличие от NumPy заключается в том, что NumPy обыч
но принимает списки или кортежи языка Python как входные дан
ные для своих функций API (не включая конструктор array()). JAX
преднамеренно отказывается от приема списков или кортежей в ка
честве входных данных своих функций, так как это может привести
к скрытому ухудшению производительности, которое трудно обна
ружить.
Если необходимо передать список Python в функцию JAX, то вы
должны явно преобразовать его в массив. Код в листинге 3.15 де
монстрирует работу со списками Python в функциях JAX. Исходный
код примеров в этом и следующих разделах размещен в репозито
рии книги: https://github.com/che-shr-cat/JAX-in-Action/tree/main/
Chapter-3.
103
Массивы в JAX
Листинг 3.15 Использование списков или кортежей Python
в функциях JAX
import numpy as np
import jax.numpy as jnp
np.array([1, 42, 31337])
>>> array([
1,
❶
42, 31337])
jnp.array([1, 42, 31337])
>>> Array([
1,
42, 31337], dtype=int32)
np.sum([1, 42, 31337])
❷
❸
>>> 31380
try:
jnp.sum([1, 42, 31337])
except TypeError as e:
print(e)
❹
>>> sum requires ndarray or scalar arguments, got <class 'list'> at
position 0.
jnp.sum(jnp.array([1, 42, 31337]))
❺
>>> Array(31380, dtype=int32)
❶
❷
❸
❹
❺
Массив NumPy nparray можно создать из списка Python.
Массив JAX Array также можно создать из списка Python.
Функция NumPy sum() может работать со списками Python.
Функция jax.numpy sum() не принимает списки Python.
Необходимо сначала создать Array, чтобы использовать функцию jax.numpy
sum().
Обратите внимание: при вызове jnp.sum() в возвращаемом типе
указано, что скаляры также упаковываются в тип Array.
Для JAX Array существует список свойств, аналогичный типу мас
сива в NumPy. В официальной документации содержится полный
перечень методов и свойств (https://docs.jax.dev/en/latest/_autosum
mary/jax.Array.html).
Листинг 3.16
Использование стандартных NumPy-подобных свойств
arr = jnp.array([1, 42, 31337])
arr.ndim
>>> 1
arr.shape
❶
❷
Глава 3
104
>>> (3,)
arr.dtype
>>> dtype('int32')
arr.size
>>> 3
arr.nbytes
>>> 12
❶
❷
❸
❹
❺
❻
Работа с массивами
❸
❹
❺
❻
Создание массива из трех целочисленных элементов.
Итоговый тензор содержит одно измерение (как обычные массивы).
Это форма тензора.
Элементами являются 32-битовые целые числа.
Тензор содержит три элемента.
Эти три элемента требуют 12 байт для хранения (так как каждый элемент имеет
длину 4 байта).
Объекты массивов предназначены для беспроблемной работы
с инструментальными средствами стандартной библиотеки Python,
где это уместно. Например, если copy.copy() или copy.deepcopy() из
встроенного модуля copy Python встречает тип Array, это равнознач
но вызову метода copy(), который создает копию буфера на том же
устройстве, где находится исходный массив.
Массивы также можно сериализовать или преобразовывать
в формат, сохраняемый в файле, с помощью встроенного модуля
pickle. По аналогии с объектами numpy.ndarray массив Array будет
сериализован через компактное представление битов. При выпол
нении обратной операции десериализации (unpickling) результатом
становится новый объект Array на устройстве, принятом по умолча
нию, потому что десериализация может выполняться в другой среде
с другими устройствами.
3.2.3
Операции, связанные с устройствами
Специализированные вычислительные устройства, такие как GPU
или TPU, помогают ускорить выполнение кода. Иногда скорость мо
жет увеличиваться на порядок. Это особенно важно для тренировки
крупных нейронных сетей или для реализации крупномасштабных
имитаций. Поэтому вы должны знать, как воспользоваться мощ
ными возможностями аппаратного ускорения. Для использования
ускоренных вычислений необходимо, чтобы все данные, участву
ющие в вычислениях (собственно тензоры), размещались в памяти
устройства-ускорителя (GPU или TPU). Следовательно, первым ша
гом в процедуре использования аппаратного ускорения является
изучение способов передачи данных между устройствами.
Массивы в JAX
105
Разумеется, существует ряд методов размещения тензоров на
устройствах. Доступными могут оказаться многочисленные устрой
ства. В рассматриваемых ниже примерах будет использоваться
блокнот Colab со средой времени выполнения GPU для демонстра
ции некоторых операций, связанных с аппаратными устройствами.
Типы аппаратных устройств: CPU, GPU, TPU
CPU – центральный процессор (ЦП), т. е. обычный процессор, производимый Intel, AMD или Apple (эта компания в настоящее время использует собственные процессоры ARM). Это универсальное вычислительное
устройство общего назначения, хотя во многих новых процессорах имеются специализированные инструкции для повышения производительности рабочих процессов машинного обучения.
GPU – графический процессор (graphics processing unit), специализированный процессор с высокой степенью распараллеливания, изначально
созданный для выполнения задач компьютерной графики. Современные
GPU содержат множество (до нескольких тысяч) простых процессоров
(ядер) и обеспечивают высокую степень распараллеливания, что делает
их весьма эффективными инструментами выполнения некоторых алгоритмов, в том числе операций умножения матриц – основы глубокого
обучения. Наиболее широко распространенными и поддерживаемыми
наилучшим образом являются GPU компании NVIDIA, хотя AMD и Intel
выпускают собственные GPU.
TPU – тензорный процессор (tensor processing unit) компании Google,
самый известный пример ASIC (application-specific integrated circuit –
интегральная микросхема специального назначения). Микросхема ASIC
специально спроектирована для конкретного варианта использования
вместо выполнения задач общего профиля, как CPU. Микросхемы ASIC
даже более специализированы, чем GPU, поскольку GPU, по существу,
остается полностью параллельным процессором с тысячами вычислительных элементов, способным выполнять множество разнообразных
алгоритмов, тогда как ASIC – это процессор, предназначенный для выполнения весьма небольшого набора вычислительных инструкций
(например, только инструкций умножения матриц). Но ASIC делает это
очень хорошо. Существует множество других микросхем ASIC для глубокого обучения и задач искусственного интеллекта, но в настоящее время их поддержка чрезвычайно ограничена во фреймворках глубокого
обучения.
Более подробная информация о микросхемах ASIC доступна здесь:
https://moocaholic.medium.com/hardware-for-deep-learning-part4-asic-96a542fe6a81.
Поскольку я использую систему с GPU, становится доступным
дополнительное аппаратное устройство. В рассматриваемых ниже
106
Глава 3
Работа с массивами
примерах будут использоваться как устройства CPU, так и GPU. Если
в вашей системе нет устройств, кроме CPU, попробуйте воспользо
ваться блокнотом Google Colab, который предоставляет облачные
GPU даже в бесплатном сегменте.
Локальные и глобальные устройства
Прежде всего существует хост (host). Хост – это CPU, управляющий
несколькими устройствами. Один хост может управлять немноги
ми устройствами (обычно до восьми), поэтому для использования
большего количества устройств требуется конфигурация со многи
ми хостами (также являющаяся мультипроцессной, как в данном
случае, где будет существовать множество JAX Python процессов,
выполняющихся независимо на каждом хосте).
JAX различает локальные и глобальные устройства. Локальное
для процесса устройство – то, к которому этот процесс может обра
титься напрямую и запустить на нем вычисления. Такое устройство
подключено непосредственно к хосту (или компьютеру), где выпол
няется программа JAX, например CPU, локальный GPU или восемь
ядер TPU, напрямую соединенных с хостом. Функция jax.local_devices() показывает локальные устройства процесса. Функция jax.
local_device_count() возвращает количество устройств, к которым
может обратиться напрямую текущий процесс. Обе функции при
нимают параметр для внутреннего компонента XLA, значением
которого может быть 'cpu', 'gpu' или 'tpu'. По умолчанию этот
параметр содержит значение None, соответствующее внутренне
му устройству, назначенному по умолчанию (GPU или TPU, если
устройство доступно).
Глобальное устройство доступно для всех процессов. Это важно
в многохостовых и многопроцессных средах. Когда каждый про
цесс запускает вычисление на своих локальных устройствах, вычис
лительная процедура может привлечь общедоступные устройства
и использовать коллективные операции (это тема глав 6 и 7) через
прямые коммуникационные связи между устройствами (обычно
высокоскоростные соединения между Cloud TPU или GPU). Функ
ция jax.devices() показывает все доступные глобальные устрой
ства, а функция jax.device_count() возвращает общее количество
устройств, известных всем процессам.
Мы подробно рассмотрим многохостовые и многопроцессные
среды в главе 7. А сейчас сосредоточимся только на средах с един
ственным хостом. В этом случае список глобальных устройств пол
ностью соответствует списку локальных устройств. Код в листин
ге 3.17 демонстрирует способы получения информации о локальных
и глобальных устройствах.
Массивы в JAX
Листинг 3.17
Получение информации об устройствах
import jax
jax.devices()
>>> [gpu(id=0)]
jax.local_devices()
>>> [gpu(id=0)]
jax.devices('cpu')
>>> [CpuDevice(id=0)]
jax.device_count('gpu')
>>> 1
❶
❷
❸
❹
❺
107
❶
❷
❸
❹
❺
Запрос внутреннего устройства по умолчанию. В моем случае это GPU.
Имеется одно устройство GPU.
Запрос только локальных устройств.
Прямой запрос о внутренних устройствах CPU.
Запрос о количестве устройств GPU.
Для некоторых устройств, не являющихся CPU, можно также уви
деть атрибут process_index. В текущий момент такой атрибут су
ществует для TPU (и мы увидим его немного позже в этой главе),
а в предыдущих версиях JAX такой же атрибут был предусмотрен
и для GPU. Каждый процесс JAX может получить собственный индекс
с помощью функции jax.process_index(). В большинстве случаев он
равен 0, но для многопроцессных конфигураций его значение будет
различным для каждого процесса.
Зафиксированные и незафиксированные данные
В JAX за вычислениями следует размещение данных. Существуют
два различных свойства размещения:
устройство, на котором размещаются данные;
зафиксированы данные на этом устройстве или нет. Если дан
ные зафиксированы, то иногда их называют приклеенными
к устройству (sticky to the device).
Место размещения данных можно узнать с помощью метода device().
По умолчанию объекты JAX Array размещаются как незафикси
рованные (uncommitted) на устройстве по умолчанию. Устройство
по умолчанию является первым элементом в списке, возвращаемом
при вызове функции jax.devices() (jax.devices()[0]). То есть это
первый графический процессор GPU или тензорный процессор TPU,
если он имеется, иначе – центральный процессор CPU:
Глава 3
108
Работа с массивами
arr = jnp.array([1, 42, 31337])
arr.device()
>>> gpu(id=0)
❶
❶ Исходный тензор находится на GPU, но он не зафиксирован.
Можно воспользоваться менеджером контекста jax.default_device() для временного замещения устройства по умолчанию для
операций JAX, если это необходимо. Также можно задействовать
переменную среды JAX_PLATFORMS или флаг командной строки --jax_
platforms. Кроме того, есть возможность установить порядок при
оритетов при представлении списка платформ в переменной среды
JAX_PLATFORMS.
Вычисления, включающие незафиксированные данные, выпол
няются на устройстве по умолчанию, и результаты также остаются
незафиксированными на устройстве по умолчанию. Допустим, что
требуется специализированный GPU для выполнения вычислений
со специфическим тензором. Можно явно разместить данные на
конкретном устройстве, используя для этого вызов функции jax.
device_put() с параметром device. В этом случае данные становят
ся зафиксированными (committed) на указанном устройстве. Если
в параметре device передается значение None, то операция поведет
себя как функция тождественности, если операнд уже находится на
любом устройстве. Иначе данные будут переданы на устройство по
умолчанию без фиксации (uncommitted data):
arr_cpu = jax.device_put(arr, jax.devices('cpu')[0])
arr_cpu.device()
❶
>>> CpuDevice(id=0)
❷
arr.device()
>>> gpu(id=0)
❶ Размещаем копию тензора на первом устройстве CPU.
❷ Проверяем, действительно ли новый тензор находится на CPU.
❸ Локацией исходного тензора остается GPU.
❸
Всегда следует помнить о функциональной сущности JAX. Функ
ция jax.device_put() создает копию исходных данных на заданном
устройстве и возвращает ее. Исходные данные остаются неизмен
ными.
Существует и обратная операция jax.device_get() для передачи
данных из устройства в процесс Python на хосте. Возвращаемые дан
ные представлены в форме массива ndarray библиотеки NumPy:
arr_host = jax.device_get(arr)
type(arr_host)
❶
109
Массивы в JAX
>>> numpy.ndarray
arr_host
>>> array([
1,
42, 31337], dtype=int32)
❶ Передача исходного тензора из GPU на хост в форме массива ndarray библиотеки
NumPy.
Вычисления с фиксированными данными выполняются на фикси
рованном устройстве, и результаты фиксируются на том же устрой
стве. При вызове операции с аргументами, зафиксированными на
различных устройствах, возникает ошибка (но если некоторые ар
гументы не являются зафиксированными, то ошибки не будет). Код
в листинге 3.18 демонстрирует размещение данных на устройстве.
Листинг 3.18
Размещение данных на устройстве
arr = jnp.array([1, 42, 31337])
arr.device()
>>> gpu(id=0)
arr_cpu = jax.device_put(arr, jax.devices('cpu')[0])
arr_cpu.device()
>>> CpuDevice(id=0)
2,
❷
❸
arr + arr_cpu
>>> Array([
❶
84, 62674], dtype=int32)
arr_gpu = jax.device_put(arr, jax.devices('gpu')[0])
try:
arr_gpu + arr_cpu
except ValueError as e:
print(e)
❹
❹
❺
❻
>>> Received incompatible devices for jitted computation.
Got argument x1 of jax.numpy.add with shape int32[3] and
device ids [0] on platform GPU and argument x2 of
jax.numpy.add with shape int32[3] and device ids [0]
on platform CPU
# Приняты несовместимые устройства для вычислений с jit-компиляцией.
# Получен аргумент x1 для jax.numpy.add с формой int32[3] и
# идентификаторы устройств [0] на платформе GPU и аргумент x2 для
# jax.numpy.add с формой int32[3] и идентификаторы устройств [0]
# на платформе CPU
❶ Исходный тензор находится на GPU, но он не зафиксирован.
❷ Мы помещаем копию исходного тензора на первое устройство CPU.
❸ Проверяем, действительно ли новый тензор находится на CPU.
Глава 3
110
Работа с массивами
❹ Вызов бинарной операции для незафиксированного и зафиксированного тензо-
ров, размещенных на различных устройствах (ОК).
❺ Фиксация нового тензора на другом устройстве.
❻ Еще один вызов бинарной операции для двух тензоров, зафиксированных на
различных устройствах (ошибка).
В более старых версиях JAX до реализации запроса на включение
изменений #6002 (https://github.com/google/jax/pull/6002) существо
вали некоторые ленивые действия при создании массива, сохраняв
шиеся во всех операциях создания константных массивов (zeros,
ones, eye и т. п.). То есть при вызове, например, jax.device_put(jnp.
ones(...), jax.devices()[1]) создавался массив с нулями1 не на
устройстве, соответствующем jax.devices()[1], а на устройстве по
умолчанию, и только потом копировался на jax.devices()[1]. В со
временных версиях JAX такая оптимизация исключена для упроще
ния реализации.
Pallas: язык ядра JAX
Для JAX существует расширение с именем Pallas, позволяющее писать
специализированные ядра для GPU и TPU с применением модели, похожей на Triton (https://triton-lang.org/main/index.html), компилятор
GPU, созданный и поддерживаемый компанией OpenAI.
О проектном решении Pallas можно узнать в соответствующем разделе
официальной документации (https://jax.readthedocs.io/en/latest/pal
las/design.html) и из примеров использования (https://docs.jax.dev/
en/latest/pallas/quickstart.html).
Среди множества примеров реального применения Pallas можно выделить код RecurrentGemma (https://github.com/google-deepmind/
recurrentgemma), включающий специализированное ядро Pallas
для выполнения линейной рекурсии на TPU (https://github.com/
google-deepmind/recurrentgemma/blob/main/recurrentgemma/jax/
pallas.py).
Вы получили общее представление о простом способе выполне
ния вычислений в стиле NumPy на GPU. Просто помните, что это
только лишь часть большой темы, касающейся увеличения произво
дительности. В главе 5 речь пойдет о JIT, методе динамической ком
пиляции, предоставляющем еще бóльшие возможности повышения
производительности. А в главе 8 рассматриваются распределенные
массивы и автоматическое распараллеливание.
1
Видимо, все-таки «массив с единицами», поскольку первым аргументом
является jnp.ones(...). – Прим. перев.
Массивы в JAX
3.2.4
111
Асинхронная диспетчеризация
Весьма важным аспектом внутренней работы JAX, о котором следует
знать, является применение асинхронной диспетчеризации (asyn
chronous dispatch). Это означает, что при выполнении какой-либо
операции JAX не ждет ее завершения и возвращает управление про
грамме на языке Python. JAX возвращает массив (Array), который
с технической точки зрения является «будущим массивом». В бли
жайшей перспективе значение не становится доступным немедлен
но, а, как следует из названия, будет сформировано через некоторое
время на устройстве-акселераторе. Тем не менее такой «будущий
массив» уже содержит форму и тип, и его можно также передавать
в последующие вычисления JAX.
Ранее мы не обращали внимания на этот факт, потому что при
просмотре выводимых результатов вычислений или преобразова
ний в массив NumPy JAX автоматически заставляет среду Python
ждать завершения вычислений. Если необходимо явное ожидание
результата, то можно воспользоваться методом block_until_ready()
объекта Array.
Асинхронная диспетчеризация чрезвычайно полезна и удобна,
поскольку позволяет среде Python продолжать работу без ожида
ния акселератора, помогая коду Python не попадать в критический
путь. Если код Python передает в очередь вычисления на устройстве
быстрее, чем их можно выполнить, и если при этом нет необходимо
сти в промежуточной проверке вычисляемых значений, то Pythonпрограмма может использовать акселератор более эффективно, не
заставляя его ждать.
Отсутствие знаний об асинхронной диспетчеризации может
привести к неверным умозаключениям во время выполнения эта
лонных тестов, и вы с большой вероятностью получите слишком
оптимистичные результаты. Именно поэтому в разделе 1.1.1 при
меняется метод block_until_ready() при эталонном тестировании
вычислительной функции на различных внутренних устройствах
с использованием JIT-компиляции и без нее.
В приведенном ниже примере (листинг 3.19) особое внимание
сосредоточено на различии по времени при измерении с блокиров
кой и без нее. При отсутствии блокировки измеряется только время
диспетчеризации работы без учета вычислений. Кроме того, мы из
меряем время вычисления на GPU и CPU. Это делается посредством
фиксации тензоров данных на соответствующих устройствах.
Листинг 3.19
Работа с применением асинхронной диспетчеризации
a = jnp.array(range(1000000)).reshape((1000,1000))
a.device()
>>> gpu(id=0)
Глава 3
112
Работа с массивами
%time x = jnp.dot(a,a)
>>> CPU times: user 757 µs, sys: 0 ns, total: 757 µs
>>> Wall time: 770 µs
%time x = jnp.dot(a,a).block_until_ready()
>>> CPU times: user 1.34 ms, sys: 65 µs, total: 1.41 ms
>>> Wall time: 4.33 ms
❶
❷
a_cpu = jax.device_put(a, jax.devices('cpu')[0])
a_cpu.device()
>>> CpuDevice(id=0)
%time x = jnp.dot(a_cpu,a_cpu).block_until_ready()
>>> CPU times: user 272 ms, sys: 0 ns, total: 272 ms
>>> Wall time: 150 ms
❸
❶ Измеряется только время диспетчеризации работы.
❷ Измеряется полное время вычислений на GPU.
❸ Измеряется полное время вычислений на CPU.
Здесь можно видеть, что вычисления на GPU выполняются в 30 раз
быстрее, чем на CPU (4,33 мс против 150 мс).
При чтении документации вы, возможно, обратили внимание на
то, что некоторые функции явно описаны как асинхронные. Напри
мер, для функции jax.device_put() отмечено «This function is always
asynchronous, i.e. returns immediately» (https://docs.jax.dev/en/latest/_
autosummary/jax.device_put.html#jax.device_put) («Эта функция всег
да асинхронна, т. е. возврат из нее происходит немедленно»). Теперь
вы понимаете, что это означает.
Последняя по порядку, но не по степени важности тема: я покажу
вам, как подготовить Cloud TPU и выполнить код JAX на этом аксе
лераторе. Мы уже изменили исходный код для работы с JAX, и ядро
работает на GPU. Мы не сделали ничего особенного для перевода
вычислений на GPU, но устройством по умолчанию стал процессор
GPU, когда код рассматриваемого здесь примера был запущен в си
стеме с GPU. Это настоящее чудо – все работает прямо «из коробки».
Теперь сделаем еще один трюк: выполним наш код на TPU. Зачем?
А просто потому, что есть такая возможность.
3.2.5
Выполнение вычислений на TPU
Выполнение кода на TPU не является необходимостью для рас
сматриваемого здесь примера обработки изображений, но такая
возможность может оказаться весьма полезной для тренировки
реальных больших нейронных сетей. Поэтому сейчас самое время
Массивы в JAX
113
продемонстрировать, как установить соединение с TPU и выполнить
код JAX на этом типе устройства. А в дальнейшем при работе с при
мерами в этой книге вы сможете с легкостью переключаться между
локальными устройствами CPU, GPU и TPU.
Такая конфигурация предполагает, что вы продолжаете запускать
код в Google Colab (или в локальном блокноте Jupyter), но не пользуе
тесь средами времени выполнения Colab TPU. Вместо этого вы раз
вертываете свой хост с Cloud TPU в Google Colab и устанавливаете со
единение Colab с тем, что называется локальной средой выполнения
(local runtime), т. е. со средой выполнения, которую предоставляете
сами себе.
За использование Cloud TPU придется заплатить некоторую де
нежную сумму (текущие расценки можно узнать здесь: https://
cloud.google.com/tpu/pricing), поэтому если такая перспектива вас не
устраивает по какой-либо причине, то вы можете просто пропустить
этот раздел.
Сравнение Colab TPU с Cloud TPU
Существуют два типа TPU-систем, с которыми вы можете встретиться.
Здесь речь идет не о версии TPU, а скорее о том, как организована такая
система.
Во-первых, имеется среда времени выполнения Colab TPU, доступная
в блокноте Colab, но уже не поддерживаемая фреймворком JAX начиная с версии 0.4. Эта архитектура устанавливает соединения Cloud TPU
удаленно через сеть с использованием gRPC с хост-компьютером, на
котором работает Colab. Пользователь не получает доступ к хосту TPU,
его компьютер подключается к TPU-акселераторам.
Для использования Colab TPU требуется шаг настройки, выполняемый
до любой операции JAX:
import jax.tools.colab_tpu
jax.tools.colab_tpu.setup_tpu()
Во-вторых, существует система Cloud TPU, доступная в Google Cloud. Это
новая архитектура с виртуальными машинами Cloud TPU Virtual Machines (VM), работающими на хост-компьютерах TPU, напрямую подключенных к TPU-акселераторам. Такая новая системная архитектура Cloud
TPU более простая и гибкая. В дополнение к основным преимущест
вам использования вы можете получить прирост производительности,
поскольку вашему коду больше не приходится совершать длительное
путешествие по сети центра обработки данных, чтобы добраться до TPU.
Более подробное описание новой архитектуры Cloud TPU можно найти
здесь: https://cloud.google.com/blog/products/compute/introducingcloud-tpu-vms.
114
Глава 3
Работа с массивами
В этой книге мы будем использовать Cloud TPU. Для этого потребуются отдельные работающие виртуальные машины (VM) и их соединения
с Google Colab, устанавливаемые вручную с использованием локальной
среды выполнения Colab.
Теперь мы должны узнать, как запускаются виртуальные машины
Cloud TPU.
Подготовка TPU к работе
Для начала работы требуется подготовка Cloud TPU. В приложении C
содержится описание всех подробностей этого процесса, а здесь вы
делены наиболее важные шаги.
Мы запускаем виртуальную машину с акселератором TPU v2-8 (вы
можете выбрать другой вариант) и устанавливаем ее соединение
с используемой локальной средой выполнения Colab:
$gcloud compute tpus tpu-vm create node-jax \
--zone us-central1-b --accelerator-type v2-8 \
--version tpu-vm-base
$gcloud compute tpus tpu-vm list --zone us-central1-b
$gcloud compute tpus tpu-vm ssh \
--zone us-central1-b node-jax -- \
-L 8888:localhost:8888
Открывается командная оболочка SSH с SSH-туннелем, перена
правляющим трафик из локального порта 8888 в порт с таким же но
мером на удаленном компьютере, где мы запустим Jupyter. На этом
компьютере можно выполнить все требуемые операции установки
(см. приложение C).
Здесь мы пропустим процедуру установки JAX, так как это можно
сделать из блокнота, и займемся установкой всего необходимого для
запуска сервера Jupyter:
$pip install -U jinja2
$pip install notebook
# Измените этот путь в соответствии со своим именем пользователя.
$export PATH=$PATH:/home/grigo/.local/bin
$pip install jupyter_http_over_ws
$jupyter notebook \
--NotebookApp.allow_origin='https://colab.research.google.com' \
--port=8888 \
--NotebookApp.port_retries=0
Сервер Jupyter начинает работу и предоставляет нам ссылку на
соединение с блокнотом (что-то вроде http://localhost:8888/?token
Массивы в JAX
115
=ac7ca95a1a2ebc0239a17f07dd4b8cb469d56ea1d1f51b36). Скопируй
те эту ссылку и используйте ее для установления соединения Colab
с локальной средой выполнения (более подробные инструкции со
скриншотами см. в приложении C).
Теперь наш блокнот Colab готов к работе с Cloud TPU. Вы можете
удалить виртуальную машину Cloud TPU, когда она станет ненужной
(не забывайте делать это, иначе будете оплачивать время соедине
ния, даже если не используете виртуальную машину!):
$gcloud compute tpus tpu-vm delete node-jax --zone us-central1-b
Выполнение вычислений на TPU
После подключения Cloud TPU к блокноту Colab вы получаете воз
можность использовать аппаратное ускорение, предоставляемое
тензорными процессорами TPU. В листинге 3.20 воспроизводится
код из листинга 3.19, но с учетом вычислений, выполняемых на TPU.
В репозитории этой книги также содержится копия блокнота с при
мерами jax.Array, но с использованием внутреннего устройства TPU
вместо GPU.
Листинг 3.20
Работа с Cloud TPU
!pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/
libtpu_releases.html
❶
from jax.lib import xla_bridge
print(xla_bridge.get_backend().platform)
>>> tpu
❷
❷
import jax
jax.local_devices()
>>> [TpuDevice(id=0, process_index=0,
coords=(0,0,0), core_on_chip=0),
>>> TpuDevice(id=1, process_index=0,
coords=(0,0,0), core_on_chip=1),
>>> TpuDevice(id=2, process_index=0,
coords=(1,0,0), core_on_chip=0),
>>> TpuDevice(id=3, process_index=0,
coords=(1,0,0), core_on_chip=1),
>>> TpuDevice(id=4, process_index=0,
coords=(0,1,0), core_on_chip=0),
>>> TpuDevice(id=5, process_index=0,
coords=(0,1,0), core_on_chip=1),
>>> TpuDevice(id=6, process_index=0,
coords=(1,1,0), core_on_chip=0),
>>> TpuDevice(id=7, process_index=0,
coords=(1,1,0), core_on_chip=1)]
❸
❸
❸
❸
❸
❸
❸
❸
Глава 3
116
Работа с массивами
import jax.numpy as jnp
a = jnp.array(range(1000000)).reshape((1000,1000))
a.device()
>>> TpuDevice(id=0, process_index=0,
coords=(0,0,0), core_on_chip=0)
%time x = jnp.dot(a,a)
>>> CPU times: user 1.07 ms, sys: 674 µs, total: 1.74 ms
>>> Wall time: 953 µs
%time x = jnp.dot(a,a).block_until_ready()
>>> CPU times: user 1.85 ms, sys: 1.17 ms, total: 3.02 ms
>>> Wall time: 2.07 ms
❶
❷
❸
❹
❺
❻
❹
❺
❻
Установка JAX с поддержкой TPU.
Проверка: какое внутреннее устройство использует JAX.
Мы видим, что доступны восемь устройств TPU.
Массив размещается в одном ядре TPU (из восьми доступных).
Измеряется только время диспетчеризации работы.
Измеряется время крупномасштабной операции на TPU.
Здесь можно видеть наличие восьми устройств TPU с идентифи
каторами (id) от 0 до 7, поскольку каждое устройство Cloud TPU пред
ставляет собой TPU-плату с четырьмя микросхемами TPU, каждая
из которых содержит два ядра. В кортеже coords указаны двоичные
координаты микросхемы TPU на плате, а атрибут core_on_chip обо
значает номер ядра внутри двухъядерной микросхемы TPU.
Также можно видеть, что устройством, на котором размещен
тензор, теперь становится TPU. Еще интереснее тот факт, что это
одно конкретное ядро TPU из восьми (доступных) ядер (первое
устройство, которое было устройством по умолчанию). Вычисле
ние скалярного произведения также происходит на этом конкрет
ном ядре.
Возможно, вы обратили внимание и на атрибут process_index.
Здесь все устройства TPU соединены с единственным процессом
JAX, имеющим индекс 0. Для мультипроцессных конфигураций вы,
вероятнее всего, увидите различные индексы. Каждый процесс JAX
может получить собственный индекс с помощью функции jax.process_index().
ПРЕДУПРЕЖ ДЕНИЕ Не забывайте останавливать и удалять
виртуальные машины Cloud GPU или TPU после завершения
их использования. Иначе это может привести к существен
ным расходам денежных средств.
Отличия от NumPy
117
Мы закончили работу с TPU. Как видите, пользоваться этими
устройствами не так уж сложно, обычно некоторого времени требует
только самый первый подготовительный этап.
Теперь мы знаем, как JAX работает с тензорами и как преобразо
вать NumPy-программу в JAX. Вы даже можете перенести нашу про
цедуру обработки изображения на TPU, хотя это не обеспечит какихлибо заметных преимуществ для такого простого примера.
Теперь необходимо уделить особое внимание различиям между
JAX и NumPy.
3.3
Отличия от NumPy
Возможно, вы все-таки будете использовать только чистый код
NumPy в тех случаях, когда не требуются никакие преимущества JAX,
особенно при выполнении небольших одноразовых вычислений. Но
в других ситуациях, когда преимущества, предоставляемые JAX, не
обходимы, вы можете перейти с NumPy на JAX. При этом потребует
ся внесение некоторых изменений в код.
Хотя NumPy-подобный API JAX пытается соответствовать насто
ящему API NumPy настолько точно, насколько это возможно, все же
существуют некоторые важные различия. Нам уже известно самое
очевидное различие в поддержке акселераторов. Тензоры могут раз
мещаться на различных внутренних устройствах (CPU, GPU, TPU),
и вам предоставляется возможность точного управления размеще
нием тензоров на устройствах. Асинхронная диспетчеризация также
относится к этой категории, поскольку изначально была предназна
чена для эффективного использования ускоренных вычислений.
Еще одно различие, которое упоминалось ранее, – поведение при
входных данных, не являющихся массивом, обсуждаемое в подраз
деле 3.2.2. Следует напомнить, что многие функции JAX не при
нимают в качестве входных данных списки или кортежи, чтобы
предотвратить снижение производительности. К прочим различи
ям относятся неизменяемость (данных) и особые темы, связанные
с поддержкой типов данных и повышением (расширением) типа.
Рассмотрим эти темы более подробно.
3.3.1
Неизменяемость
Массивы JAX являются неизменяемыми (immutable). Возможно,
раньше вы не обращали на это внимания, но попробуйте изменить
любой тензор – и увидите сообщение об ошибке. Почему возникает
ошибка? Давайте попытаемся изменить конкретный тензор и по
смотрим, что происходит.
Глава 3
118
Листинг 3.21
Работа с массивами
Обновление элемента тензора
a_jnp = jnp.array(range(10))
a_np = np.array(range(10))
a_np[5], a_jnp[5]
>>> (5, Array(5, dtype=int32))
a_np[5] = 100
a_np[5]
❶
>>> 100
try:
a_jnp[5] = 100
except TypeError as e:
print(e)
❷
>>> '<class 'jaxlib.xla_extension.ArrayImpl'>'
object does not support item assignment. JAX
arrays are immutable. Instead of ``x[idx] = y``,
use ``x = x.at[idx].set(y)`` or another .at[] method: ...
# Объект '<class 'jaxlib.xla_extension.ArrayImpl'>'
# не поддерживает присваивание элементу массива.
# Массивы JAX являются неизменяемыми.
# Вместо ``x[idx] = y`` используйте ``x = x.at[idx].set(y)``
# или другой метод .at[]: ...
❶ Операция присваивания с заменой для элемента массива NumPy (ОК).
❷ Операция присваивания с заменой для элемента массива JAX (не разрешена).
Напомню, JAX спроектирован на основе парадигмы функцио
нального программирования. Именно поэтому трансформации
JAX обладают такой мощью. Есть несколько превосходных книг по
функциональному программированию, например «Grokking Func
tional Programming» (https://www.manning.com/books/grokking-func
tional-programming), «Grokking Simplicity» (https://www.manning.
com/books/grokking-simplicity) и другие, поэтому я даже не пытаюсь
полностью охватить такую обширную тему в этой книге. Но основы
функциональной чистоты я все же напомню: код не должен созда
вать побочные эффекты. Код, изменяющий исходные аргументы, не
является функционально чистым. Единственный способ создания
измененного тензора – создание другого тензора на основе исход
ного. Возможно, вы наблюдали такое поведение в других системах
и языках, следующих функциональной парадигме, например в Spark.
Такой подход противоречит некоторым практическим методикам
программирования, принятым в NumPy. В NumPy обычной опера
цией является обновление индекса, происходящее при изменении
значения в тензоре по индексу, например изменение значения пя
119
Отличия от NumPy
того элемента массива. Это вполне допустимо в NumPy, но приводит
к возникновению ошибки в JAX.
Надо отдать должное JAX – сообщение об ошибке весьма инфор
мативно, и в нем предлагается решение возникшей проблемы. Рас
смотрим, как устранить ошибку, возникшую в коде листинга 3.21.
Функциональность обновления индекса
Для всех типовых выражений, используемых для непосредствен
ного обновления значения элемента тензора, существуют соответ
ствующие функционально чистые аналоги, которые можно приме
нять в JAX. В табл. 3.1 приведен список функциональных операций
JAX, равнозначных выражениям прямого изменения элементов
в NumPy-стиле.
Таблица 3.1 Функциональность обновления индекса в JAX
Выражение прямого изменения элементов в NumPy-стиле
Равнозначный синтаксис JAX
x[idx] = y
x = x.at[idx].set(y)
x[idx] += y
x = x.at[idx].add(y)
x[idx] *= y
x = x.at[idx].multiply(y)
x[idx] /= y
x = x.at[idx].divide(y)
x[idx] **= y
x = x.at[idx].power(y)
x[idx] = minimum(x[idx], y)
x = x.at[idx].min(y)
x[idx] = maximum(x[idx], y)
x = x.at[idx].max(y)
ufunc.at(x, idx)
x = x.at[idx].apply(ufunc)
x = x[idx]
x = x.at[idx].get()
Все приведенные в табл. 3.1 выражения x.at возвращают изме
ненную копию x, не изменяя оригинал. Такой подход может ока
заться менее эффективным, чем код прямой замены оригинала, но
благодаря JIT-компиляции низкоуровневые выражения, подобные x
= x.at[idx].set(y), гарантированно применяются непосредственно
на месте (если исходная копия больше не используется), что дела
ет вычисления эффективными. Поэтому не следует беспокоиться об
эффективности при использовании функциональности обновления
индекса.
ПРЕДУПРЕЖ ДЕНИЕ Существуют более старые функции
jax.ops.index_*, которые в настоящее время объявлены
устаревшими и не рекомендуемыми к применению, начиная
с версии JAX 0.2.22.
Теперь можно изменить код, чтобы исправить ошибку, заменив
операцию обновления элемента в стиле NumPy на функциональ
ность обновления индекса, предоставляемую JAX.
Глава 3
120
Листинг 3.22
Работа с массивами
Обновление элемента тензора в JAX
a_jnp = a_jnp.at[5].set(100)
a_jnp[5]
❶
>>> Array(100, dtype=int32)
❶ Создание копии исходного тензора с измененным заданным элементом.
Это и есть сущность неизменяемости; изменения кода обязатель
но должны быть однозначными и ясными, а если вы что-то пропус
тили, то JAX даст вам знать, как исправить ошибку.
Индексирование за границами массива
Часто встречается следующий тип ошибки: индексирование мас
сива, выходящее за его границы. В NumPy, где для обработки по
добных случаев используются исключения языка Python, ситуация
вполне очевидная. Но если код работает на акселераторе, такой под
ход применить трудно или даже невозможно. Таким образом, не
обходимо обеспечить поведение без ошибок при выходе индекса за
пределы массива. Для операций обновления индекса за границами
массива подходящим решением был бы пропуск таких обновлений,
а для операций извлечения по индексу – жесткая фиксация на гра
нице массива, поскольку необходим возврат какого-то значения.
Это похоже на обработку ошибок при вычислениях с плавающей
точкой с использованием таких значений, как NaN (Not-a-Number –
не число).
По умолчанию в JAX предполагается, что все индексы находят
ся в границах массива. Создана экспериментальная поддержка для
предоставления более точной семантики доступа к индексам вне
границ массива через параметр mode для функций обновления ин
декса. Ниже перечислены возможные варианты значений этого па
раметра:
"promise_in_bounds" (по умолчанию) – пользователь гаранти
рует, что все индексы находятся в границах массива, поэтому
дополнительная проверка не выполняется. На практике это
означает, что все индексы, выходящие за границы массива,
отсекаются в операции get() и отбрасываются (исключаются)
в set(), add() и других функциях, изменяющих элементы;
"clip" – индексы за границами массива фиксируются в допус
тимом диапазоне;
"drop" – индексы за границами массива игнорируются (отбра
сываются);
"fill" – псевдоним (алиас) для значения "drop", но для функ
ции get() будет возвращаться значение, заданное в необяза
тельном аргументе fill_value.
121
Отличия от NumPy
В листинге 3.23 показан пример использования различных вари
антов параметра mode.
Листинг 3.23 Индексация за границами массива
a_jnp = jnp.array(range(10))
a_jnp
>>> Array([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=int32)
a_jnp[42]
❶
>>> Array(9, dtype=int32)
a_jnp.at[42].get(mode='drop')
❷
>>> Array(-2147483648, dtype=int32)
a_jnp.at[42].get(mode='fill', fill_value=-1)
❸
>>> Array(-1, dtype=int32)
a_jnp = a_jnp.at[42].set(100)
a_jnp
❹
>>> Array([0, 1, 2, 3, 4, 5, 6, 7, 8, 9], dtype=int32)
a_jnp = a_jnp.at[42].set(100, mode='clip')
a_jnp
>>> Array([
0,
1,
2,
3,
4,
5,
❺
6,
7,
8, 100], dtype=int32)
❶ Поведение по умолчанию при выходе индекса за границы массива (фиксация в допустимом
диапазоне для операции 'get’).
❷ Использование поведения drop (отбрасывание, исключение).
❸ Использование поведения fill (заполнение) с указанием специального заполняющего зна
чения.
❹ Поведение по умолчанию при выходе индекса за границы массива (отбрасывание (исключе-
ние) для операции 'get’).
❺ Использование поведения clip (фиксация индекса в допустимом диапазоне).
Как видите, индексация вне границ массива не приводит к воз
никновению ошибок – всегда возвращается некоторое значение,
и в подобных случаях вы можете управлять поведением.
Вот и все о неизменяемости на текущий момент. Рассмотрим еще
одну важную тему, в которой существует множество отличий от NumPy.
3.3.2
Типы
Что касается типов данных, то в JAX существуют многочисленные
отличия от NumPy. Это поддержка формата с плавающей точкой
низкой и высокой точности, а также семантика повышения (расши
122
Глава 3
Работа с массивами
рения) типа, определяющая, какой тип будет иметь результат опера
ции, если ее операнды принадлежат к конкретным (возможно, раз
личным) типам.
Поддержка типа float64
Несмотря на то что NumPy безапелляционно расширяет операнды до
двойной точности (тип float64), JAX по умолчанию принудительно
использует числа с одиночной (обычной) точностью (тип float32).
Возможно, вы будете удивлены, когда напрямую создадите массив
типа float64, а JAX тихо и незаметно переведет его в тип float32.
Для многих рабочих процессов машинного обучения (и особенно
глубокого обучения) это идеальный подход. Для некоторых научных
вычислений с высокой точностью такое преобразование, возможно,
станет нежелательным.
Типы чисел с плавающей точкой: float64, float32, float16,
bfloat16
В научных вычислениях и в глубоком обучении используется множество
типов чисел с плавающей точкой. Международный стандарт арифметики с плавающей точкой IEEE Standard for Floating-Point Arithmetic (IEEE
754) определяет несколько форматов с различной точностью, которые
широко используются.
По умолчанию типом данных с плавающей точкой для научных вычислений является число с плавающей точкой двойной точности, или
float64 с размером 64 бита. Стандарт IEEE 754 так определяет двоичный формат числа с плавающей точкой двойной точности: 1 бит –
знак, 11 бит – показатель степени (экспонента) и 52 бита – дробная
часть. Такое представление соответствует диапазону от ≈2,23e–308 до
≈1,80e308 с полной гарантией 15–17 точных десятичных цифр.
Для некоторых случаев существуют типы с более высокой точностью,
например long double, или тип с плавающей точкой расширенной точности, который обычно представлен как 80-битовое число с плавающей
точкой на платформах x86 (но при этом возникает множество ограничений и оговорок). NumPy поддерживает тип np.longdouble для обес
печения расширенной точности, хотя JAX этот тип не поддерживает.
В приложениях глубокого обучения надежность результатов обеспечивается более низкой точностью, поэтому число с плавающей точкой
одиночной точности, или тип float32, стал типом данных по умолчанию
для таких приложений, в том числе и в фреймворке JAX. 32-битовое
число с плавающей точкой по стандарту IEEE 754 содержит 1 бит знака,
8 бит показателя степени (экспоненты) и 23 бита дробной части. Его
диапазон – от ≈1,18e–38 до ≈3,40e38 с точностью 6–9 значимых десятичных цифр.
123
Отличия от NumPy
Для многих задач глубокого обучения даже 32 бит числа с плавающей
точкой слишком много, и в течение нескольких последних лет стали все
больше распространяться методы тренировки и (логического) вывода
с еще меньшей точностью. Обычно проще выполнить вывод с меньшей
точностью, чем тренировку, поэтому существуют некоторые хитроумные
схемы, в которых объединены типы с точностью float16/32 для тренировки.
В группе форматов чисел с плавающей точкой низкой точности имеются
два 16-битовых представления: float16 и bfloat16. По стандарту IEEE
754 число с плавающей точкой половинной точности, или float16, содержит 1 бит знака, 5 бит показателя степени и 10 бит дробной части.
Его диапазон – от ≈5,96e–8 до 65 504 с четырьмя значимыми десятичными цифрами. Другой 16-битовый формат, разработанный компанией
Google (точнее, подразделением Google Brain, отсюда и первое слово
в названии формата), называется «Brain Floating Point Format», или сокращенно bfloat16. Исходный формат IEEE float16 не предназначался
специально для приложений глубокого обучения, и его динамический
диапазон слишком узок. Тип bfloat16 решает эту проблему, предоставляя динамический диапазон, равнозначный типу float32. Он содержит
1 бит знака, 8 бит показателя степени и 7 бит дробной части. Таким образом, диапазон расширяется от ≈1,18e–38 до ≈3,40e38 с тремя значимыми десятичными цифрами.
Формат bfloat16 представляет собой усеченную версию типа IEEE 754
float32 и позволяет выполнять быстрое преобразование в формат IEEE
754 float32 и обратно. При преобразовании в формат bfloat16 все
биты показателя степени сохраняются, тогда как поле значащей части
(мантиссы) может быть усечено.
Существуют и некоторые другие специализированные форматы, о которых можно узнать более подробно из статьи: https://moocaholic.
medium.com/fp64-fp32-fp16-bfloat16-tf32-and-other-members-ofthe-zoo-a1ca7897d407.
Чтобы обеспечить выполнение вычислений с типом float64, не
обходимо установить переменную конфигурации jax_enable_x64
перед запуском системы. Код в листинге 3.24 показывает, как это
делается.
Листинг 3.24 Обеспечение вычислений с использованием типа
float64
# Это работает только при первом запуске.
from jax.config import config
config.update("jax_enable_x64", True)
import jax.numpy as jnp
❶
Глава 3
124
Работа с массивами
# Возможно, это не будет работать на внутреннем устройстве TPU.
# Попробуйте использовать CPU или GPU.
x = jnp.array(range(10), dtype=jnp.float64)
❷
x.dtype
>>> dtype('float64')
❸
❶ Необходимо разрешить использование типа float64.
❷ Создание массива типа float64.
❸ Это должен быть тип 'float64’, если его использование разрешено при запуске.
Иначе это будет тип 'float32’.
Но 64-битовые типы данных поддерживаются не на каждом внут
реннем устройстве. Например, TPU не поддерживает этот тип.
Можно пойти и в другом направлении, к форматам более низкой
точности.
Поддержка типа float16/bfloat16
В глубоком обучении существует тенденция к использованию фор
матов с низкой точностью – чаще всего формата с половинной точ
ностью, или float16, а также более специализированного форма
та bfloat16, который не поддерживается базовым чистым кодом
NumPy. В JAX можно с легкостью переключаться между этими 16-би
товыми типами низкой точности.
Листинг 3.25
Использование 16-битовых типов с плавающей точкой
xb16 = jnp.array(range(10), dtype=jnp.bfloat16)
xb16.dtype
>>> dtype(bfloat16)
xb16.nbytes
❶
>>> 20
x16 = jnp.array(range(10), dtype=jnp.float16)
x16.dtype
>>> dtype('float16')
❶ Использование типа bfloat16.
❷ Использование типа float16.
❷
Следует еще раз отметить, что на конкретных внутренних устрой
ствах могут существовать некоторые ограничения.
Семантика повышения (расширения) типа
При выполнении бинарных операций правила повышения (расши
рения) типа JAX в некоторых аспектах отличаются от правил NumPy.
Отличия от NumPy
125
Разумеется, основное различие относится к типу bfloat16, по
скольку NumPy его не поддерживает. Но существуют различия
и в некоторых других случаях.
Любопытный факт: если при сложении двух 16-битовых чисел
с плавающей точкой одно из них имеет обычный тип float16, а вто
рое – тип bfloat16, то вы получите результат типа float32.
Листинг 3.26 Семантика повышения (расширения) типа
xb16+x16
>>> Array([ 0.,
xb16+xb16
2.,
4.,
❶
6.,
8., 10., 12., 14., 16., 18.], dtype=float32)
❷
>>> Array([0, 2, 4, 6, 8, 10, 12, 14, 16, 18], dtype=bfloat16)
❶ Суммирование значений типа bfloat16 со значениями типа float16.
❷ Суммирование двух значений одинакового типа bfloat16.
Таблица с описанием различий между семантикой повышения
(расширения) типа в NumPy и JAX для бинарных операций нахо
дится здесь: https://docs.jax.dev/en/latest/type_promotion.html. Кро
ме того, существует более глубокое сравнение NumPy/TensorFlow/
PyTorch / JAX NumPy и JAX lax (описание JAX lax см. в следующем
разделе): https://docs.jax.dev/en/latest/jep/9407-type-promotion.html
#appendix-example-type-promotion-tables.
Специализированные типы тензоров
До сих пор мы работали с так называемыми плотными тензорами
(dense tensors). Плотные тензоры явно содержат все значения. Но во
многих случаях приходится иметь дело с разреженными данными
(sparse data), т. е. содержащими многочисленные нулевые элементы
и некоторое количество (обычно на порядок или даже на несколько
порядков меньше) ненулевых значений.
Разреженные тензоры явно содержат только ненулевые значе
ния и позволяют сэкономить изрядный объем памяти. Тем не менее
для их эффективного использования требуются специальные под
программы линейной алгебры с поддержкой разреженности. В JAX
имеется экспериментальный API, связанный с операциями с разре
женными тензорами и матрицами, размещенный в модуле jax.experimental.sparse. Возможно, статус этого API изменится к моменту
выхода книги из печати, но сейчас мы не будем рассматривать экс
периментальный модуль.
Другой вариант использования – набор (коллекция) тензоров раз
личной формы, например речевые записи различной длины. Для
создания пакета таких тензоров в современных фреймворках суще
Глава 3
126
Работа с массивами
ствуют специальные структуры данных, такие как «неровный» тен
зор (ragged tensor) в TensorFlow.
В JAX нет специализированных структур, подобных «неровному»
тензору, хотя и не запрещено работать с такими данными. Автома
тическое формирование пакетов (auto-batching) с помощью vmap()
может помочь во многих случаях, и мы рассмотрим примеры при
менения этой трансформации в главе 6. Еще одно существенное раз
личие между JAX и «первоисточником» NumPy – в JAX существует
альтернативный API низкого уровня, который мы рассмотрим в сле
дующем разделе.
3.4
Интерфейсы высокого и низкого уровней:
jax.numpy и jax.lax
Во время разработки примера обработки изображения мы познако
мились с API jax.numpy, предназначенным для предоставления при
вычного интерфейса пользователям, хорошо знакомым с NumPy.
Для учебного примера обработки изображения NumPy API было
вполне достаточно, и нам не потребовался какой-либо другой API.
Но вы должны знать, что существует другой API низкого уровня
jax.lax, который является основой таких библиотек, как jax.numpy.
API jax.numpy представляет собой высокоуровневую обертку со все
ми операциями, выраженными через базисные элементы (primi
tives) jax.lax. Многие базисные элементы jax.lax сами по себе явля
ются тонкими обертками для равнозначных операций XLA (https://
openxla.org/xla/operation_semantics). На рис. 3.10 показана схема
уровней JAX API.
Рис. 3.10 Схема уровней JAX API
jax.numpy API
Базисные элементы jax.lax
Операции XLA
API jax.lax предлагает огромный набор математических опера
торов. Некоторые из них являются более общими и предоставля
ют возможности, отсутствующие в jax.numpy. Например, при вы
полнении одномерной свертки в jax.numpy используется функция
Интерфейсы высокого и низкого уровней: jax.numpy и jax.lax
127
jnp.convolve (https://docs.jax.dev/en/latest/_autosummary/jax.numpy.
convolve.html#jax.numpy.convolve). Эта функция, в свою очередь, об
ращается к более общей функции conv_general_dilated из jax.lax
(https://docs.jax.dev/en/latest/_autosummary/jax.lax.conv_general_di
lated.html).
Также существуют и другие группы операторов: для управления
потоком выполнения, для вычисления специализированного гра
диента, для поддержки распараллеливания и операторы линейной
алгебры. Полный список доступных операторов jax.lax находится
здесь: https://docs.jax.dev/en/latest/jax.lax.html#module-jax.lax.
Мы будем рассматривать некоторые операторы вычисления гра
диента в главе 4, а поддержку распараллеливания обсудим в главе 7.
В следующем подразделе кратко описываются базисные элементы
управления потоком выполнения.
3.4.1
Базисные элементы управления потоком выполнения
Набор структурированных базисных элементов управления пото
ком выполнения в jax.lax включает lax.cond для условного приме
нения одной из двух функций (аналог оператора if), lax.switch для
разветвления по нескольким путям, lax.while_loop и lax.fori_loop
для многократно повторяющегося вызова функции в цикле, lax.
scan и lax.associative_scan для сканирования (последовательно
го прохода) массива с помощью заданной функции с сохранением
состояния и функцию lax.map для применения заданной функции
к главным осям массива. Все эти базисные элементы являются диф
ференцируемыми и могут быть обработаны JIT-компилятором, что
помогает избежать развертывания длинных циклов.
В среде JAX вы можете продолжать использовать управляющие
структуры Python, и в большинстве случаев они работают (в разде
ле 5.2 рассматривается внутренняя структура этого процесса и его
ограничения). Но такие решения могут оказаться неоптимальными,
либо потому что существует более эффективное решение, либо по
тому что они создают менее производительный дифференцируе
мый код.
В этой книге мы будем использовать некоторые из перечислен
ных выше и другие базисные элементы jax.lax, и в последующих
главах при первом применении каждого элемента будет приводить
ся его описание. А сейчас предлагается один пример с использова
нием lax.switch.
Рассмотрим случай, прямо противоположный задаче фильтрации
изображения. В глубоком обучении для компьютерного зрения час
то требуется сделать нейронную сеть устойчивой к разнообразным
искажениям изображения. Например, необходима нейронная сеть,
устойчивая к шуму. Для ее обучения вы должны предоставить вер
Глава 3
128
Работа с массивами
сии шума в изображениях, т. е. нужно создать такие изображения
с шумом. Также может потребоваться использование аугментации
изображений для эффективного расширения тренировочного набо
ра данных вариантами исходного изображения: повороты на неко
торый угол, вокруг вертикальной оси (иногда вокруг горизонталь
ной оси), искажения цвета и т. п.
Предположим, что имеется набор функций для аугментации изо
бражений:
augmentations = [
add_noise_func,
horizontal_flip_func,
rotate_func,
adjust_colors_func
]
❶
❶
❶
❶
❶ Список из четырех функций-заглушек для обработки изображений (это только
пример, код этих функций не предоставляется).
Для этих функций код не предоставляется, их имена включены
в список исключительно в демонстрационных целях. Но вы можете
реализовать их в качестве дополнительного упражнения.
В листинге 3.27 используется lax.switch для случайного выбора
аугментации изображения из нескольких предложенных вариантов.
Значение первого параметра, здесь это значение переменной augmentation_index, определяет, какая из нескольких допустимых вет
вей будет применяться. Второй параметр, в нашем случае это спи
сок augmentations, предоставляет последовательность функций (или
ветвей) для выбора. Мы используем список аугментаций изображе
ния, каждая из которых выполняет отдельную операцию обработки
изображения: добавление шума, поворот вокруг вертикальной оси,
вращение и т. п. Если переменная augmentation_index содержит 0, то
выбирается первый элемент списка augmentations (вызывается со
ответствующая функция), если переменная augmentation_index со
держит 1, то выбирается второй элемент списка и т. д. Остальные
параметры (в нашем примере всего один – image) передаются в вы
бранную функцию, а значение, возвращаемое из выбранной функ
ции (здесь – измененное изображение), становится итоговым значе
нием оператора lax.switch.
Листинг 3.27 Пример управления потоком выполнения с использованием
lax.switch
from jax import random
❶
def random_augmentation(image, augmentations, rng_key):
'''A function that applies a random transformation to an image'''
# Функция, которая применяет случайно выбранную трансформацию к изображению.
129
Интерфейсы высокого и низкого уровней: jax.numpy и jax.lax
augmentation_index = random.randint(
key=rng_key, minval=0,
maxval=len(augmentations), shape=())
augmented_image = lax.switch(
augmentation_index, augmentations, image)
return augmented_image
❷
❷
❷
❸
new_image = random_augmentation(image, augmentations, random.PRNGKey(42))
❶ Импорт модуля генерации случайных чисел в JAX (более подробно о нем – в главе 9).
❷ Генерация случайного целого значения для получения индекса в списке аугментаций.
❸ Использование функции lax.switch для выбора одного из нескольких вариантов.
В главе 5 станет понятно, почему код с использованием базис
ных элементов jax.lax, управляющих потоком выполнения, может
обеспечить более высокую эффективность вычислений. В главе 9
мы создадим полностью завершенную версию функции из примера
в листинге 3.27 с обсуждением правильных способов работы со слу
чайными числами.
3.4.2
Повышение (расширение) типа
Правила jax.lax API строже, чем правила jax.numpy. Не выполняется
неявное повышение типа аргументов для операций со смешанными
типами данных. При использовании jax.lax в подобных случаях не
обходимо вручную проводить расширение типов.
Листинг 3.28
Повышение (расширение) типа в jax.numpy и jax.lax
import jax.numpy as jnp
jnp.add(42, 42.0)
>>> Array(84., dtype=float32, weak_type=True)
from jax import lax
try:
lax.add(42, 42.0)
except TypeError as e:
print(e)
>>> lax.add requires arguments to have the same
dtypes, got int32, float32. (Tip: jnp.add is a
similar function that does automatic type
promotion on inputs).
# lax.add требует аргументы, имеющие одинаковые
# типы dtypes, но приняты типы int32, float32.
# (Совет: jnp.add - аналогичная функция, которая
# выполняет автоматическое расширение типов
# входных данных.)
❶
❷
❸
Глава 3
130
Работа с массивами
lax.add(jnp.float32(42), 42.0)
>>> Array(84., dtype=float32)
❹
❶ Сложение целочисленного значения и значения с плавающей точкой в jax.numpy
(ОК).
❷ Тип результата расширяется до float32.
❸ Сложение целочисленного значения и значения с плавающей точкой в jax.lax
(возникает ошибка).
❹ Вручную выполняется преобразование типа, чтобы получить два входных значе-
ния типа float32, которые суммируются в jax.lax (ОК).
В примере из листинга 3.28 можно видеть свойство weak_type
в значении с типом Array. Это свойство означает, что было пере
дано значение без явно указанного пользователем типа, например
скалярные литералы Python. О значениях со слабыми типами в JAX
более подробно можно узнать здесь: https://docs.jax.dev/en/latest/
type_promotion.html#weak-types.
Здесь мы не очень подробно рассмотрели jax.lax API, посколь
ку большинство его преимуществ будут описаны в главах с 4 по 6.
А сейчас важно подчеркнуть, что jax.numpy API является не един
ственным в JAX.
Упражнение 3.1
1
2
3
Реализуйте собственный цифровой фильтр по своему выбору.
Реализуйте функции аугментации изображения, упомянутые
в разделе 3.4.
Реализуйте другие операции обработки изображения: кадриро
вание, изменение размера, поворот на некоторый угол и прочие
действия, которые вы считаете необходимыми.
Резюме
Тензор, или многомерный массив, – основная структура данных
во фреймворках глубокого обучения и научных вычислений.
Для представления тензоров NumPy предоставляет тип numpy.
ndarray, а в JAX существует тип jax.Array (ранее называвшийся
DeviceArray).
В JAX имеется NumPy-подобный API, импортируемый из модуля
jax.numpy и пытающийся соответствовать как можно точнее ис
ходному NumPy API, но некоторые различия все же имеют место.
Вы можете с высокой точностью управлять размещением на
устройствах данных типа Array, фиксируя их на конкретном вы
бранном устройстве – CPU, GPU или TPU.
Резюме
131
В JAX каждое конкретное вычисление учитывает размещение
данных.
В JAX различаются локальные и глобальные устройства. Локаль
ными являются устройства, соединенные с локальным хостом.
Глобальные устройства доступны всем процессам JAX на несколь
ких хостах.
JAX использует асинхронную диспетчеризацию для вычислений.
Если необходимо явно дождаться результата вычислений, то мож
но воспользоваться методом block_until_ready() объекта Array.
Для использования TPU необходимо развернуть собственный
хост с тензорными процессорами Cloud TPU в облаке Google Cloud
и установить соединение с блокнотом Colab, используя локальную
среду выполнения.
Массивы JAX являются неизменяемыми, поэтому необходимо
использовать функционально чистые аналоги выражений (функ
ций) NumPy.
JAX предоставляет различные режимы для управления поведени
ем без возникновения ошибок при индексировании за границами
массива.
По умолчанию в JAX типом данных с плавающей точкой является
float32. При необходимости можно использовать типы float64/
float16/bfloat16, но при этом могут существовать некоторые
ограничения, связанные с конкретными внутренними устрой
ствами.
Для бинарных операций правила повышения (расширения) типа
JAX в некоторых аспектах отличаются от правил NumPy.
API низкого уровня jax.lax более строгий и часто предоставляет
более мощные функциональные средства, нежели API высокого
уровня jax.numpy.
API jax.lax включает базисные элементы структурированного
управления потоком выполнения, которые являются дифферен
цируемыми и могут быть обработаны JIT-компилятором, что по
могает избежать развертывания длинных циклов.
API jax.lax не выполняет неявное расширение типов аргументов
для операций со смешанными типами данных.
4
Вычисление градиентов
Темы главы:
вычисление производных различными способами;
использование автоматического дифференцирования
(autodiff) в JAX для вычисления градиентов
пользовательских функций (и нейронных сетей);
использование прямого и обратного режимов autodiff.
В главе 2 мы тренировали простую нейронную сеть для классифи
кации рукописных цифр. Чрезвычайно важным фактором для тре
нировки любой нейронной сети является возможность вычисления
производной функции потерь с учетом весов нейронной сети.
Следует напомнить, что нейронные сети, как правило, трениру
ются с применением процедуры градиентного спуска (обычно это
некоторая расширенная версия подобной процедуры, такая как
Adam). Нейронная сеть – это сложная математическая функция,
определяемая по своим весам, или весовым коэффициентам (иног
да используются миллиарды весов, как для GPT-3 и других крупных
моделей). Учитывая веса, вы можете вычислить выходное значение
нейронной сети, а используя функцию потерь, вы оцениваете по
лученный результат и его приближенность к идеальному. Поэто
му необходимо обновлять веса нейронной сети, чтобы уменьшить
Резюме
133
значение функции потерь. Процедура градиентного спуска прини
мает некоторую исходную функцию (нейронную сеть с ее весами)
и функцию потерь, вычисляет производную функции потерь с уче
том весов нейросети и получает градиент, представляющий собой
направление и скорость самого быстрого роста функции. Затем,
двигаясь в направлении, противоположному градиенту, вы обнов
ляете веса нейронной сети и уменьшаете значение функции потерь.
Вполне очевидно, что здесь вычисление градиентов – самый важ
ный шаг, поэтому обязательным условием является наличие спосо
ба получения градиентов.
Существует несколько способов вычисления производных исполь
зуемых функций (или их дифференцирования), но автоматическое
дифференцирование (обозначаемое для краткости autodiff) – это ос
новной способ, применяемый в современных фреймворках глубоко
го обучения. Методика autodiff также является одним из основопо
лагающих компонентов фреймворка JAX (напомню, что JAX часто
называют «Autograd and XLA»). JAX позволяет писать код на Python
с использованием стиля NumPy, позволяя самому фреймворку вы
полнять самую трудную и хитроумную работу по вычислению про
изводных. В этой главе вы научитесь эффективно работать с autodiff.
Мы подробно рассмотрим все самые важные его свойства и главные
принципы работы в JAX.
Начнем со сравнения различных способов вычисления производ
ных: ручного, символьного, численного и автоматического диффе
ренцирования. Чрезвычайно полезно понимать различия между
методами дифференцирования, поскольку это помогает получить
намного более ясное представление о том, почему autodiff являет
ся таким великолепным инструментом. Если вы уже хорошо знаете
способы вычисления производных, то можете пропустить первый
раздел.
Во втором разделе сравнивается вычисление градиентов в Py
Torch, TensorFlow и JAX и рассматриваются все основные транс
формации градиента в JAX. Это ключевая часть главы, в которой
описаны все подробности использования JAX для вычисления про
изводных. Для тех, кто только начинает работать с JAX, это самая
важная часть.
Заключительный раздел предназначен для тех, кто хочет понять
внутренний механизм функционирования autodiff. В нем описыва
ются прямой и обратный режимы autodiff, а также две соответству
ющие трансформации JAX – jvp() и vjp(). Думаю, это самая слож
ная часть книги. Вы можете пропустить ее, так как в большинстве
случаев вполне можно использовать JAX без глубокого понимания
внутреннего механизма autodiff. Тем не менее знание этого меха
низма позволит вам лучше понять, как эффективно применять JAX
на практике.
Глава 4
134
4.1
Вычисление градиентов
Различные способы вычисления
производных
Существует много задач, в которых необходимо вычислять произ
водные. В процессе тренировки нейронной сети требуются произ
водные (точнее, градиенты) собственно для тренировки сети, т. е.
для минимизации функции потерь посредством выполнения после
довательных шагов в направлении, противоположном градиенту.
В более простых случаях, знакомых вам из курса математического
анализа, нужно находить минимумы (или максимумы) конкретной
функции. Но в физике, инженерии, биологии и прочих научных дис
циплинах существует множество других вариантов.
Начнем с простого наглядного учебного примера поиска мини
мума функции (что в действительности и является сущностью про
цедуры тренировки в глубоком обучении). Пусть имеется простая
математическая функция, например f(x) = x4 + 12x + 1/x. Реализуем ее
в форме кода. Исходный код для этой части доступен в репозитории
GitHub в соответствующем блокноте: https://github.com/che-shr-cat/
JAX-in-Action/blob/main/Chapter-4/JAX_in_Action_Chapter_4_Differ
ent_ways_of_getting_derivatives.ipynb.
Листинг 4.1 Простая функция, которую мы будем использовать
как модель
def f(x):
return x**4 + 12*x + 1/x
❶
❶ Простая математическая функция, реализованная как код Python.
Выполняем поиск локального минимума функции. Необходимо
найти точки, в которых производная этой функции равна нулю, за
тем проверить каждую найденную точку – это должна быть точка
минимума, но не максимума и не седловая точка (она же точка пере
вала, перегиба или минимакс).
Точки минимума, максимума и седловые точки
Существует три типа точек, в которых функция действительного аргумента имеет нулевую производную:
точки максимума – точки, в которых функция имеет значение, большее или равное значению в любой соседней точке. Локальные (или
относительные) максимумы – это точки, в которых функция имеет
наибольшее значение в определенном интервале. Локальный максимум дополнительно называется глобальным (или абсолютным), если
он также имеет наибольшее значение во всей области определения
(в наборе допустимых входных значений (аргументов)) функции;
Различные способы вычисления производных
135
точки минимума – точки, в которых функция имеет значение, меньшее или равное значению в любой соседней точке. Локальные минимумы – это точки, в которых функция имеет наименьшее значение
в определенном интервале. Локальный минимум дополнительно называется глобальным, если он также имеет наименьшее значение во
всей области определения функции.
Точки минимума и максимума также называются точками экстремума;
седловая точка (saddle point; точка минимакса (minimax), иногда –
точка перевала или перегиба) – точка с нулевой производной, но не
являющаяся локальным экстремумом функции.
Перечисленные выше три типа точек наглядно показаны на рис. 4.1.
4. Глобальный
максимум
1. Локальный
максимум
3. Седловая
точка
2. Локальный
минимум
5. Глобальный
минимум
Рис. 4.1 Наглядное представление точек локального и глобального минимума и максимума, а также седловой точки некоторой функции (это не та
функция, с которой мы работаем в тексте). Штриховые линии обозначают
касательную к графику функции в этих особых точках. Угол наклона касательной равен производной функции в соответствующих точках, где производная равна нулю
Точка 1 – максимум, так как она выше соседних точек, но это локальный,
а не глобальный максимум функции, поскольку существует другой максимум с более высоким значением – точка 4. Точки 2 и 5 – минимумы,
так как они ниже своих соседей, а точка 5 является глобальным минимумом (самой низкой точкой), тогда как точка 2 – локальный минимум.
Точка 3 – седловая точка с нулевой производной функции, не являющаяся ни минимумом, ни максимумом, поскольку соседняя точка слева
расположена ниже, а соседняя точка справа – выше.
Нам необходима возможность вычисления производной функции.
Мы не будем выполнять полную процедуру поиска минимума функ
ции, обсудим только часть, касающуюся вычисления производных.
Глава 4
136
Вычисление градиентов
Существуют различные способы вычисления производных. Рас
смотрим, как они работают, сравним методы и узнаем, почему auto
diff является таким полезным и удобным инструментом. Весьма
важно знать все методы дифференцирования, потому что их ино
гда некорректно ассоциируют с autodiff, тогда как только один метод
действительно является настоящим autodiff.
4.1.1
Дифференцирование вручную
Старый добрый метод, известный многим из курса математического
анализа, изучаемого в школе, заключается во взятии производных
вручную. С технической точки зрения областью определения функ
ции является множество ненулевых действительных чисел, и мы вы
числяем производную только в этой области определения.
Для рассматриваемой здесь конкретной функции производная
выглядит так: f ‘(x) = 4x3 + 12 – 1/x2 – и может быть с легкостью вычис
лена вручную (поскольку производная суммы равна сумме произ
водных, производная x4 – это 4x3, для 12x – производная 12, а для 1/x
получаем производную –1/x2), если вы помните правила дифферен
цирования и производные для некоторых простых функций. Если
результат нужен в коде, то можно реализовать его напрямую сразу
после вычисления производной вручную.
Листинг 4.2 Выражение в конечном виде (в замкнутой форме) для
производной, вычисленной вручную
def df(x):
return 4*x**3 + 12 - 1/x**2
x = 11.0
❶
print(f(x))
>>> 14773.09090909091
print(df(x))
>>> 5335.99173553719
❶ Вручную вычисленная производная для рассматриваемой функции модели.
Здесь мы получаем так называемое выражение в конечном виде
(или в замкнутой форме), которое можно вычислить в любой инте
ресующей нас точке. В примере из листинга 4.2 вычисляется произ
водная в точке x = 11.0.
Все усложняется, если функция не является настолько простой,
например для функций, представляющих собой произведение дру
гих функций или композицию других функций. Взгляните на функ
цию f(x) = (2x + 7)(17x2 – 3x)(x2 – 1/x)(x3 + 21x)/(17x – 5/x2). У меня не
Различные способы вычисления производных
137
возникает даже мысли о том, чтобы попытаться вручную вычислить
ее производную. Я воспользовался WolframAlpha (https://www.wol
framalpha.com/) для вычисления производной этой функции и полу
чил громадное выражение (x2(–6615 + 47460x + 17325x2 + 16620x3 –
122206x4 – 51282x5 – 33339x6 + 157352x7 + 58905x8 + 11526x9 + 4046x10))/
(5 – 17x3)2. Для выполнения вручную всех шагов дифференцирования
потребовалось бы несколько листов бумаги, и я не был бы уверен
в правильности конечного результата. При таких сложных вычисле
ниях легко совершить ошибку.
Необходимо понимать, как берется производная вручную, по
скольку знание основополагающих принципов полезно всегда. Но
для реальных нейронных сетей дифференцирование вручную стано
вится головной болью. Без autodiff такая процедура отняла бы очень
много времени без гарантии отсутствия ошибок (возможно, вам
приходилось делать это, если вы работали с нейросетями 10 и более
лет назад). Приходится вычислять производные всех слоев вручную
и реализовывать их как отдельный код для вычисления обратного
распространения ошибки. При таком процессе невозможно быстро
итерировать код.
4.1.2
Символьное дифференцирование
В предыдущем подразделе я воспользовался WolframAlpha для вы
числения производной функции. Этот механизм обеспечивает сим
вольное дифференцирование, в автоматизированном режиме вы
полняя все шаги, которые пришлось бы проделать вручную. Можно
запрограммировать все правила дифференцирования, и компьютер
сможет следовать этим правилам гораздо быстрее, чем человек.
Кроме того, такой подход надежнее, так как исключает возможность
возникновения случайных ошибок (при условии отсутствия ошибок
в самом механизме дифференцирования и надежности аппаратного
оборудования).
Существуют различные варианты использования символьного
дифференцирования. В Python можно воспользоваться средствами
символьного дифференцирования из библиотеки SymPy (https://
www.sympy.org/en/index.html). Эта библиотека установлена по умол
чанию в блокноте Google Colab, но ее также можно установить в лю
бой системе, выполняя инструкции с сайта https://docs.sympy.org/
latest/install.html.
Листинг 4.3 Пример использования SymPy для выполнения
символьного дифференцирования
import sympy
x = 11.0
Глава 4
138
Вычисление градиентов
x_sym = sympy.symbols('x')
f_sym = f(x_sym)
df_sym = sympy.diff(f_sym)
print(f_sym)
>>> x**4 + 12*x + 1/x
print(df_sym)
>>> 4*x**3 + 12 - 1/x**2
f = sympy.lambdify(x_sym, f_sym)
print(f(x))
>>> 14773.09090909091
df = sympy.lambdify(x_sym, df_sym)
print(df(x))
>>> 5335.99173553719
❶
❷
❸
❹
❹
❺
❻
❺
❻
❶ Определение переменной x.
❷ Создание символьного выражения, передающего переменную в конкретную
❸
❹
❺
❻
функцию.
Вычисление символьной производной.
Вывод рассматриваемой функции в символьной форме.
Преобразование выражений SymPy в выражения, которые можно вычислить.
Вычисление исходной функции и ее производной.
Как и при дифференцировании вручную, результатом символь
ного дифференцирования становится выражение в конечном виде,
в действительности являющееся отдельной функцией. Эту функцию
можно применять (вычислять) в любой интересующей нас точке.
В примере из листинга 4.3 вычисляется производная в той же точке
x = 11.0.
Если можно выполнять вычисления в символьной форме, это пре
восходный вариант. Символьное дифференцирование великолепно,
хотя в некоторых случаях и эта процедура усложняется.
Во-первых, необходимо представлять все вычисления в виде вы
ражения в конечном виде (в замкнутой форме). Могут возникнуть
трудности с формированием выражения в конечном виде, если вы
числения реализованы как алгоритм, особенно при использовании
некоторых управляющих логических конструкций с операторами if
и циклами for.
Во-вторых, размер символьных выражений увеличивается при
выполнении дальнейшего дифференцирования, например при вы
числении производных более высоких порядков. Это явление на
Различные способы вычисления производных
139
зывается «утолщением» (или «раздутием») выражения (expression
swell). Иногда размер выражения увеличивается по экспоненциаль
ному закону.
4.1.3
Численное дифференцирование
При ручном или символьном дифференцировании мы имеем вы
ражение в конечном виде с интересующей нас функцией и получа
ем выражение в конечном виде для ее производной, которое можно
использовать для вычисления производной в любой заданной точ
ке. Проблема заключается в том, что выражение в конечной форме
для исходной функции не всегда можно получить с легкостью. Кроме
того, существует множество вариантов, в которых нас интересует зна
чение производной только в одной конкретной точке, а не в любой
произвольной. Именно такой вариант соответствует целям глубокого
обучения, т. е. требуются градиенты для текущего набора весов.
Метод под названием численное дифференцирование часто ис
пользуется в науке и инженерии для количественной оценки про
изводной любой функции. Численное дифференцирование – это
способ вычисления приближенного значения математической про
изводной. Существует много методов численного дифференциро
вания. Один из наиболее широко известных подходов использует
метод конечных разностей, основанный на определении предела
производной. Метод дает оценку производной от функции f(x), вы
числяя угловой коэффициент близлежащей секущей линии, прове
денной через точки (x, f(x)) и (x + Δx), f(x + Δx)):
Этот метод использует два вычисления функции в точках, распо
ложенных близко друг к другу и обозначенных как x и (x + Δx), где
Δx, или размер шага, равен очень малому значению, например 10–6.
Листинг 4.4 Определение производной с помощью численного
дифференцирования
x = 11.0
dx = 1e-6
df_x_numeric = (f(x+dx)-f(x))/dx
print(df_x_numeric)
❶
❷
>>> 5335.992456821259
❶ Выбор размера шага.
❷ Выполнение численного дифференцирования для рассматриваемой функции
модели.
140
Глава 4
Вычисление градиентов
Результатом является приблизительное значение градиента
в конкретной точке x (здесь x = 11.0). Так как это приближенное вы
числение, результат не может быть точным, и мы видим, что он от
личается от результата ручного и символьного дифференцирования
в третьем десятичном знаке.
В отличие от ручного и символьного дифференцирования, воз
вращающего решение в конечном виде, которое можно использо
вать в любой точке, методу численного дифференцирования о таком
решении ничего неизвестно. Он просто вычисляет приближенное
значение производной в конкретной заданной точке.
У этого метода имеется несколько недостатков, связанных со ско
ростью, точностью и стабильностью вычислений. Начнем со скоро
сти. Требуются два вычисления значений функции при одном ска
лярном входном данном, что, возможно, приведет к увеличению
накладных расходов для сложных функций. Более того, для каждого
скалярного значения при передаче в функцию многих входных дан
ных необходимы два вычисления, а это приводит к вычислитель
ной сложности O(N). Почти все нейронные сети представляют собой
функции многих переменных (напомню, это могут быть миллионы,
миллиарды и даже триллионы тренируемых весов, которые мы оп
тимизируем с помощью процедуры градиентного спуска). Поэтому
такой процесс масштабируется очень плохо.
Точность не самая лучшая по определению, так как используются
приближенные вычисления. В любом случае точность ухудшают два
типа ошибок: погрешность приближения и погрешность округле
ния – обе связаны с выбранными размером шага.
К выбору размера шага необходим разумный подход. Размер шага
должен быть достаточно малым для более точного приближения
к значению производной. Вы начинаете уменьшать размер шага,
и погрешность приближения (ее иногда также называют ошибкой
усечения (отбрасывания), т. е. ошибкой, возникающей из-за того, что
размер шага не является действительно бесконечно малой величи
ной) уменьшается. Но в некоторый момент возникает ограничение
по разрядности арифметики с плавающей точкой, и с этого момента
значение ошибки (погрешности) начинает увеличиваться, если вы
продолжаете уменьшать размер шага. Это называется погрешно
стью округления, возникающей из-за того, что значимые младшие
биты конечного результата вступают в конкуренцию за простран
ство в машинном слове со старшими битами значений f(x + Δx) и f(x),
так как эти значения хранятся только до тех пор, пока не уничтожат
друг друга при вычитании в конце процедуры. Кроме того, можно
получить нестабильные результаты для функций с помехами, по
скольку малые изменения во входных данных (или в значениях ве
сов) способны привести к весьма существенным изменениям в про
изводной.
Различные способы вычисления производных
141
На практике невозможно использовать численное дифференци
рование как основной способ вычисления производных в методах
обучения на основе градиентов. Они оказываются слишком медлен
ными. В старые добрые времена, когда приходилось вручную вы
числять производные всех слоев, численное дифференцирование
использовалось как проверка правильности результатов посред
ством сравнения с производными, вычисленными вручную. Если
результаты оказывались достаточно близкими друг к другу (еще
один метапараметр, который был вынужден самостоятельно опре
делять пользователь, например результаты должны отличаться не
более чем на 10–8), то, вероятнее всего, дифференцирование вруч
ную было правильным и его можно продолжать использовать. Более
существенное расхождение указывало на наличие ошибки в вычис
лениях, и требовалось повторное выполнение ручной работы по по
лучению производной.
4.1.4
Автоматическое дифференцирование
И вот, наконец, появилось автоматическое дифференцирование –
autodiff. Основная идея autodiff проста и понятна. Дифференци
руемая функция состоит из простейших элементов (примитивов)
и операций (сложения, деления, возведения в степень и т. д.), произ
водные которых хорошо известны. Во время вычисления функции
autodiff отслеживает все операции и распространяет производные
с помощью цепного правила дифференцирования сложной функ
ции по мере развертывания вычислительной процедуры. Это чрез
вычайно важно для глубокого обучения, поскольку нейронная сеть
представляет собой композицию функций в следующей форме:
fL(fL–1...f1(f0(x))), где f i (i принадлежит отрезку [0, L]) – функция, вычис
ляющая i-й слой, а x – входные данные. Такую функцию определяют
как имеющую глубокую вложенность.
Аutodiff позволяет вычислять производные для весьма сложных
компьютерных программ, включающих управляющие структуры,
разветвления кода, циклы и рекурсию, т. е. программные компонен
ты, которые трудно поддаются представлению в форме выражения
в конечном виде. Это немалое преимущество.
Ранее при наличии функции Python и необходимости вычисления
производной этой функции вы имели в своем распоряжении два ва
рианта. Первый: представить функцию как выражение в конечной
форме, затем вычислить производную вручную или в символьном
виде. Для многих функций, применяемых на практике, с нетриви
альным потоком управления выполнением внутри такой подход
мог привести к существенным затратам времени, а сама процедура
получения выражения в конечной форме становилась трудной или
даже невозможной.
Глава 4
142
Вычисление градиентов
Второй вариант: использование численного дифференцирования,
но это медленная и плохо масштабируемая процедура при боль
шом количестве скалярных элементов входных данных. Например,
в нейронную сеть для классификации изображений MNIST, описан
ную в главе 2, передается 28 × 28 = 784 скалярных элемента входных
данных. Для каждого элемента необходимо выполнить два вычисле
ния функции, т. е. более чем 1500 операций прямого распростране
ния в нейросети для предварительной оценки всех градиентов. Это
слишком много.
Autodiff предоставляет новый вариант. (Если быть точным, auto
diff – это достаточно старое направление разработок, зародившееся
еще в 1960-х гг. К сожалению, оно долго оставалось малоизвестным
(да и сейчас остается) для многих специалистов-практиков в об
ласти глубокого обучения.) Autodiff применяется к обычному коду
с минимальными изменениями, поэтому можно сосредоточиться
на написании кода необходимых вычислений, а autodiff возьмет на
себя всю работу по вычислению производных.
Во всех основных современных фреймворках глубокого обучения,
включая TensorFlow, PyTorch и JAX, реализован механизм autodiff.
Мы уже использовали его в примере из главы 2, и вы знаете, что
трансформация grad() создает функцию, вычисляющую производ
ную исходной функции. В следующем разделе мы сравним подход,
принятый в JAX, с подходами TensorFlow и PyTorch. Пример в лис
тинге 4.5 демонстрирует применение autodiff в JAX. Полученный
здесь результат немного отличается от результатов ручного и сим
вольного дифференцирования, так как JAX по умолчанию исполь
зует числа с плавающей точкой более низкой точности – float32
вместо float64. В главе 3 вы узнали, как изменить тип числа с пла
вающей точкой.
Листинг 5.4
Вычисление производной с помощью autodiff в JAX
df = jax.grad(f)
print(df(x))
>>> 5335.9917
❶
❷
❶ Получение производной заданной функции с помощью autodiff в JAX.
❷ Вычисление производной в конкретной заданной точке.
Для autodiff существует два режима: прямой (forward mode) и об
ратный (reverse mode). Более подробно они будут рассматриваться
немного позже в этой главе.
Если вы начали заниматься нейросетями 10 и более лет назад, то
помните, как обстояли дела в те времена. В вашем распоряжении
Различные способы вычисления производных
143
находился некоторый язык программирования с поддержкой мат
ричных вычислений (тогда лучше было использовать MATLAB, чем
C++), и вы занимались реализацией архитектуры нейронной сети
как последовательности простых операций с матрицами (вы и сей
час можете делать в точности то же самое с использованием NumPy).
Затем требовалась реализация обратного распространения (ошибок
обучения). Приходилось вручную вычислять производные для слоев
нейронной сети, реализовывать их в виде другого набора матрич
ных операций, затем заниматься реализацией численного диффе
ренцирования только для того, чтобы проверить правильность пре
дыдущих вычислений (численное дифференцирование невозможно
было применять вместо механизма обратного распространения
ошибок, поскольку такой подход работал слишком медленно, поэто
му область применения численного дифференцирования ограничи
валась только проверкой результатов в процессе разработки). Весь
этот процесс требовал огромного количества времени, и в нем су
ществовала весьма высокая вероятность совершения ошибки. Ите
рации происходили очень медленно, что существенно ограничивало
скорость разработки.
Потом появились фреймворки. В некоторых основных фреймвор
ках (например, Caffe) по-прежнему требовалось дифференцирова
ние вручную для реализации конкретных слоев с использованием
прямых и обратных функций. Другие фреймворки (Theano, а за ним
TensorFlow и PyTorch) начали использовать autodiff, чтобы принять
на себя обязанности по вычислению производных и позволить поль
зователю сосредоточиться на написании только логики прямых вы
числений. Это стало весьма существенной поддержкой. Скорость
итеративных этапов разработки существенно увеличилась, и появи
лась возможность (при наличии более мощного аппаратного обеспе
чения) тестировать различные модификации буквально за минуты,
а не за часы или дни. Я считаю, что эти фреймворки стали самым
главным инструментом, обеспечившим общедоступность средств
глубокого обучения.
JAX продолжает развивать эту замечательную традицию – вычис
ление производных пользовательских функций, добавляя возмож
ность вычисления градиентов даже для специализированного кода
Python с логикой управления. Но это не первый фреймворк, пре
доставляющий подобную возможность. (Библиотека Autograd по
явилась раньше TensorFlow, Chainer и DyNet можно выделить среди
первых инструментов поддержки динамических вычислительных
графов, затем был представлен PyTorch. В TensorFlow 1 имелись не
которые весьма ограниченные средства, использующие библиотеку
Fold, но, вообще говоря, TensorFlow начал вычислять производные
для динамических вычислительных графов начиная с версии 2.)
Глава 4
144
Вычисление градиентов
JAX предоставляет гораздо бóльшую гибкость, поскольку не прос
то вычисляет градиенты для пользователя, но также выполняет
трансформацию функций в другие функции, которые вычисляют
градиенты исходной функции. На этом можно не останавливаться
и сгенерировать функцию для аналогичного дифференцирования
более высокого порядка. Это особенно важно в научных областях за
пределами глубокого обучения, именно поэтому JAX широко при
меняется не только для глубокого обучения.
Мы обсудили несколько методов дифференцирования, а теперь
более подробно рассмотрим, как использовать autodiff в современ
ных фреймворках.
4.2
Вычисление градиентов с использованием
autodiff
Сейчас мы рассмотрим другую практическую задачу – простой, но
полезный пример применения линейной регрессии. Предположим,
что имеется зашумленный (содержащий ошибки) набор данных об
измерении температуры, полученных с новейшей платы Raspberry
Pi с сенсорным датчиком температуры. Необходимо вычислить ли
нейный тренд, описывающий имеющиеся данные (более сложная
функция с трендом и периодической составляющей была бы более
подходящей, но ее реализацию я предлагаю выполнить вам в ка
честве упражнения).
Этот пример также можно рассматривать как вариант модели для
тренировки практически любой нейронной сети, например для за
дачи классификации изображений из главы 2. Здесь мы сосредото
чим внимание на том, как все работает на более низком уровне.
Для генерации зашумленных данных использовалась алгоритми
ческая процедура. Вы можете воспользоваться любыми данными по
вашему выбору – измерениями температуры или влажности, коли
чеством наблюдаемых метеоров, ценами на фондовом рынке и т. п.
Исходный код для этого раздела доступен в репозитории GitHub
в соответствующем блокноте: https://github.com/che-shr-cat/JAXin-Action/blob/main/Chapter-4/JAX_in_Action_Chapter_4_Gradients_in_
TensorFlow_PyTorch_JAX.ipynb.
Используется следующая процедура: генерация некоторых рав
номерно распределенных выборок для измерений времени (обозна
ченных как x), затем для каждого измерения времени генерируется
значение, являющееся суммой или константой (65.0), линейный воз
растающий тренд (1.8x), периодический процесс (функция косинуса)
и некоторая случайная ошибка (из нормального распределения).
Вычисление градиентов с использованием autodiff
Листинг 4.6
145
Генерация данных для задачи регрессии
import numpy as np
import matplotlib.pyplot as plt
x = np.linspace(0, 10*np.pi, num=1000)
e = np.random.normal(scale=10.0, size=x.size)
y = 65.0 + 1.8*x + 40*np.cos(x) + e
plt.scatter(x, y)
❶
❷
❸
❶ Генерация 1000 точек в диапазоне от 0 до 10π.
❷ Генерация случайного гауссова шума.
❸ Генерация данных, состоящих из смещений, линейного тренда, синусоидальной
волны и шума.
Код в листинге 4.6 генерирует и выводит в графическом виде дан
ные, как показано на рис. 4.2.
Рис. 4.2 Визуальное представление измерений температуры
по искусственно сформированным данным
Структура этой задачи очень похожа на структуру задачи класси
фикации изображений из главы 2: имеются тренировочные данные,
состоящие из значений x и y, функция y = f(x) для прогнозирования
значения y по x, функция потерь для оценки ошибки прогнозирова
ния и процедура вычисления градиента для адаптации весов модели
в направлении, противоположном градиенту функции потерь с уче
том весов модели. На более высоком уровне функция модели (здесь:
f(x)) и функция потерь связаны, как показано на рис. 4.3.
Сейчас y содержит действительные числа вместо классов. Мо
дель прогнозирования в текущий момент представлена простой
линейной функцией (вместо многослойной нейронной сети) вида
y = wx + b, где x и y – данные, а w и b – получаемые в процессе обуче
Глава 4
146
Вычисление градиентов
ния веса (параметры модели). Здесь мы обозначаем результат при
менения функции модели как ŷ, чтобы отличать его от истинного
значения y. Функцией потерь, оценивающей, насколько прогноз (ŷ)
далек от истинных данных (y), теперь становится среднеквадратиче
ская ошибка модели MSE (mean-squared error), часто используемая
для задач регрессии.
Входные
данные (x)
Параметры
модели
Функция
модели
Прогноз (ŷ)
Функция
потерь
Значение
функции
потерь
Целевые
значения (y)
Рис. 4.3
Вычисление значения функции потерь
Наконец, мы имеем тренировочный цикл, итеративно выпол
няющий шаги обновления градиента. В основном это тот же цикл,
который использовался для примера классификации изображений
в главе 2; это именно та часть, что нас больше всего интересует пря
мо сейчас.
Рассмотрим процедуру вычисления градиента с другой точки зре
ния и сравним, как это делается в TensorFlow/PyTorch и JAX. Если
вы не имеете дело с TensorFlow или PyTorch и не интересуетесь, как
эти фреймворки работают с градиентами, то можете сразу перейти
к подразделу 4.2.3.
4.2.1
Работа с градиентами в TensorFlow
В TensorFlow (и во многих других фреймворках) необходимо непре
менно дать фреймворку знать, какие именно тензоры требуются для
отслеживания вычислений, чтобы собрать градиенты. Это делается
с помощью параметра trainable=True для переменных в TensorFlow.
Затем выполняются вычисления с тензорами, а фреймворк посто
янно отслеживает производимые вычисления. В TensorFlow необхо
димо использовать ленты градиентов (gradient tapes) для их отсле
живания (более подробно о лентах градиентов вы узнаете в конце
этой главы, когда будут рассматриваться прямой и обратный режи
мы).
В ходе тренировки нейронной сети вы получаете некоторое ко
нечное значение, ошибку прогноза, или в более обобщенном смыс
ле функцию потерь. В довершение ко всему необходимо вычислить
производные функции потерь с учетом параметров вычисления
(интересующих вас тензоров). В TensorFlow используется функция
gradient(), передающая функцию потерь и требуемые параметры.
147
Вычисление градиентов с использованием autodiff
Затем autograd вычисляет градиенты для каждого параметра моде
ли, возвращая их как результат работы функции gradient().
Получив эти градиенты, вы можете выполнить шаг градиентного
спуска, если он необходим. В рассматриваемых здесь примерах реа
лизуется только один такой шаг. Возможно, вы захотите расширить
пример, чтобы получить полноценный цикл тренировки, – считай
те, что это дополнительное упражнение.
Листинг 4.7
Вычисление градиентов в TensorFlow
import tensorflow as tf
xt = tf.constant(x, dtype=tf.float32)
yt = tf.constant(y, dtype=tf.float32)
learning_rate = 1e-2
w = tf.Variable(1.0, trainable=True)
b = tf.Variable(1.0, trainable=True)
def model(x):
return w * x + b
❶
❷
❷
❸
❸
❹
def loss_fn(prediction, y):
return tf.reduce_mean(tf.square(prediction-y))
❺
with tf.GradientTape() as tape:
prediction = model(x)
loss = loss_fn(prediction, y)
❻
❻
❻
dw, db = tape.gradient(loss, [w, b])
w.assign_sub(learning_rate * dw)
b.assign_sub(learning_rate * db)
❶
❷
❸
❹
❺
❻
❼
❽
❼
❽
❽
Импорт TensorFlow.
Преобразование тренировочных данных в тензоры TensorFlow.
Тензоры с весами модели помечаются флагом отслеживания градиентов.
Функция реализует простую линейную модель.
Функция потерь MSE.
Вычисления выполняются в контексте GradientTape.
Извлечение результатов из ленты градиентов.
Выполняется один шаг обновления градиентов.
Как видите, все не так уж сложно, особенно если вы раньше уже ис
пользовали этот механизм. Но такой подход не является интуитив
но понятным при первом знакомстве с ним, поэтому необходимо
запомнить следующий набор правил: пометить тензоры специаль
ным параметром, создать ленту градиентов и вызвать специальную
функцию для получения вычисленных градиентов. Это совсем не
похоже на математическую форму записи.
Глава 4
148
Вычисление градиентов
Теперь рассмотрим, каким способом PyTorch получает градиенты.
4.2.2
Работа с градиентами в PyTorch
В PyTorch также требуется пометить тензоры, для которых необхо
димо отслеживать градиенты. Для тензоров PyTorch используется
специальный параметр requires_grad=True. PyTorch формирует на
правленный ациклический граф (directed acyclic graph – DAG) для от
слеживания операций во время вычислений, включающих помечен
ные тензоры. После завершения расчетов с тензорами градиенты
вычисляются с помощью вызова специальной функции backward()
для тензора потерь. Затем autograd вычисляет градиенты для каждо
го параметра модели, сохраняя их в специальном атрибуте тензора
с именем grad.
Листинг 4.8
Вычисление градиентов в PyTorch
import torch
xt = torch.tensor(x)
yt = torch.tensor(y)
learning_rate = 1e-2
w = torch.tensor(1.0, requires_grad=True)
b = torch.tensor(1.0, requires_grad=True)
def model(x):
return w * x + b
❶
❷
❷
❸
❸
❹
def loss_fn(prediction, y):
return ((prediction-y)**2).mean()
❺
prediction = model(xt)
loss = loss_fn(prediction, yt)
❻
❻
loss.backward()
with torch.no_grad():
w -= w.grad * learning_rate
b -= b.grad * learning_rate
w.grad.zero_()
b.grad.zero_()
❶
❷
❸
❹
❺
❻
❼
❼
❽
❾
❾
❿
❿
Импорт PyTorch.
Преобразование тренировочных данных в тензоры PyTorch.
Тензоры с весами модели помечаются флагом отслеживания градиентов.
Функция реализует простую линейную модель.
Функция потерь MSE.
Выполнение вычислений модели и потерь.
Вычисление градиентов.
Вычисление градиентов с использованием autodiff
149
❽ Использование менеджера контекста для запрещения вычислений градиентов
(недопустимо, чтобы обновления параметров воздействовали на градиенты).
❾ Выполнение одного шага обновления градиентов.
❿ Обнуление градиентов перед следующим шагом вычисления градиентов.
В целом структура та же, что и в TensorFlow. Аналогичным образом
тензоры помечаются специальным атрибутом, вы должны помнить
об особой области видимости для вычислений (или запрещения вы
числений) градиентов (теперь с внутренним представлением в фор
ме DAG вместо ленты градиентов), а кроме того, необходимо знать
специализированные методы и атрибуты для получения градиентов.
Поскольку все операции отслеживаются в ленте градиентов или
в DAG при выполнении вычислений, поток управления Python обраба
тывается естественным образом, поэтому можно использовать управ
ляющие операторы (такие как if или while) непосредственно в модели.
4.2.3
Работа с градиентами в JAX
JAX применяет autodiff способом, вполне совместимым с принци
пами функционального программирования. Специальная транс
формация jax.grad() принимает числовую функцию, написанную
на Python, и возвращает новую функцию Python, вычисляющую гра
диент исходной функции по ее первому параметру. На рис. 4.4 по
казана схема процесса вычисления градиентов в JAX.
Рис. 4.4 Получение функции
для вычисления градиентов
Параметры
модели
Входные
данные (x)
Функция
потерь
Значение
потерь
Целевые
значения (y)
Трансформация
grad()
Параметры
модели
Входные
данные (x)
Градиент
функции
потерь
Градиенты
Целевые
значения (y)
Весьма важный факт: результатом трансформации grad() явля
ется функция, а не значение градиента. Для вычисления значений
градиента необходимо передать в нее точку, в которой нужно полу
чить градиент.
Глава 4
150
Листинг 4.9
Вычисление градиентов
Вычисление градиентов в JAX
import jax
import jax.numpy as jnp
xt = jnp.array(x)
yt = jnp.array(y)
learning_rate = 1e-2
model_parameters = jnp.array([1., 1.])
def model(theta, x):
w, b = theta
return w * x + b
❶
❶
❷
❷
❸
❹
def loss_fn(model_parameters, x, y):
prediction = model(model_parameters, x)
return jnp.mean((prediction-y)**2)
❺
grads_fn = jax.grad(loss_fn)
grads = grads_fn(model_parameters, xt, yt)
model_parameters -= learning_rate * grads
❻
❼
❽
❶
❷
❸
❹
❺
❻
❼
❽
Импорт JAX и NumPy-подобного интерфейса.
Преобразование тренировочных данных в массивы JAX.
Тензоры с весами модели без каких-либо специальных пометок.
Функция реализует простую линейную модель.
Функция потерь MSE.
Создание функции для вычисления градиентов.
Вычисление градиентов.
Выполнение одного шага обновления градиентов.
Такой подход делает JAX API абсолютно непохожим на другие биб
лиотеки автоматического дифференцирования, такие как Tensor
Flow и PyTorch. В JAX вы работаете непосредственно с функциями,
более близкими к внутреннему математическому ядру и в опреде
ленном смысле более естественно воспринимаемыми: здесь функ
ция потерь – это функция параметров модели и данных, поэтому ее
градиент определяется тем же способом, который вы применили бы
в математике.
Нет необходимости в пометке тензоров каким-то особым спосо
бом, вы не обязаны знать что-либо о внутренних механизмах, свя
занных с лентами градиентов или направленными ациклическими
графами, и не требуется запоминать каждый конкретный способ по
лучения градиентов с применением специальных функций и атри
бутов. Вы просто используете функцию (трансформацию grad()) для
создания другой функции, вычисляющей градиенты. Затем начинае
те работать напрямую с полученной второй функцией.
Очевидно, что во фреймворках, подобных PyTorch или Tensor
Flow, большая часть этой «магии» обычно скрыта за высокоуровне
Вычисление градиентов с использованием autodiff
151
выми API и объектами с внутренними состояниями. Кроме того, это
поле деятельности для оптимизаторов, которые не использовались
в рассмотренных выше примерах.
Когда дифференцирования по первому параметру
недостаточно
Выше было отмечено, что трансформация jax.grad() вычисляет
градиент исходной функции по ее первому параметру. А что, если
требуется дифференцирование с учетом более одного параметра
или нужен не первый параметр? И для таких вариантов существуют
способы вычисления.
Рассмотрим вариант, когда интересующий нас параметр не являет
ся первым (приведенный ниже код и все примеры кода до конца главы
можно найти в соответствующем блокноте: https://github.com/cheshr-cat/JAX-in-Action/blob/main/Chapter-4/JAX_in_Action_Chapter_4_
Differentiating_in_JAX.ipynb). Мы создаем обобщенную функцию для
вычисления расстояния между двумя точками, называемого расстоя
нием Минковского (Minkowski distance). Эта функция принимает
дополнительный параметр, определяющий порядок. В этом случае
не требуется дифференцирование по этому параметру (указанному
первым). Необходимо дифференцирование по параметру x:
def dist(order, x, y):
❶
return jnp.power(jnp.sum(jnp.abs(x-y)**order), 1.0/order)
❶ Функция с дополнительным параметром в первой позиции. Требуется дифферен-
цирование по второму параметру.
Существует несколько способов, позволяющих переписать исход
ную функцию для использования поведения по умолчанию функции
grad() и продолжения дифференцирования по первому параметру.
Можно изменить исходную функцию и переместить параметр order
в конец списка или воспользоваться функцией-адаптером для из
менения порядка параметров.
Но предположим, что сделать это невозможно по некоторой при
чине. Для такого случая в функции grad() предусмотрен дополни
тельный параметр argnums. Этот параметр определяет, какой позици
онный аргумент необходимо учитывать при дифференцировании:
dist_d_x = jax.grad(dist, argnums=1)
❶
dist_d_x(1, jnp.array([1.0,1.0,1.0]), jnp.array([2.0,2.0,2.0]))
>>> Array([-1., -1., -1.], dtype=float32)
❶ Необходимо дифференцирование по второму параметру x.
С помощью параметра argnums можно сделать гораздо больше.
Глава 4
152
Вычисление градиентов
Дифференцирование по нескольким параметрам
Параметр argnums позволяет выполнять дифференцирование по бо
лее чем одному параметру. Если требуется дифференцирование по
обоим параметрам x и y, то передается кортеж, определяющий их
позиции:
dist_d_xy = jax.grad(dist, argnums=(1,2))
❶
dist_d_xy(1, jnp.array([1.0,1.0,1.0]), jnp.array([2.0,2.0,2.0]))
>>> (Array([-1., -1., -1.], dtype=float32), Array([1., 1., 1.],
dtype=float32))
❶ Требуется дифференцирование по параметрам x и y.
Параметр argnums может быть целым числом (если определяется
один параметр для дифференцирования по нему) или последова
тельностью целых чисел (если определяется несколько параметров).
Если argnums – целое число, то возвращаемый градиент имеет ту
же форму и тип, что и позиционный аргумент, указанный этим чис
лом. Если argnums – последовательность (например, кортеж), то гра
диент является кортежем значений с теми же формами и типами,
что и соответствующие аргументы.
В дополнение к явно заданному указанию, по каким параметрам
необходимо дифференцировать исходную функцию, и использова
нию параметра argnums также существует возможность упаковки
нескольких значений в один параметр функции, и мы уже восполь
зовались такой возможностью в листинге 4.9, правда, не уделили ей
никакого внимания. В коде из листинга 4.9 мы упаковали два значе
ния в массив и передали этот массив в функцию.
JAX позволяет пользователю выполнять дифференцирование по
различным структурам данных, а не только по массивам и корте
жам. Например, можно дифференцировать по словарям. Кроме того,
существует еще более обобщенная структура данных, называемая
pytree – древовидная структура, сформированная из контейнеро
образных объектов Python. Мы рассмотрим эту тему более подробно
в главе 10.
В листинге 4.10 показан пример дифференцирования по словарям
Python. Основная часть кода взята из листинга 4.9 с внесением из
менений для работы со словарями.
Листинг 4.10
Дифференцирование по словарям
model_parameters = {
'w': jnp.array([1.]),
'b': jnp.array([1.])
}
❶
153
Вычисление градиентов с использованием autodiff
def model(param_dict, x):
w, b = param_dict['w'], param_dict['b']
return w * x + b
❷
def loss_fn(model_parameters, x, y):
prediction = model(model_parameters, x)
return jnp.mean((prediction-y)**2)
grads_fn = jax.grad(loss_fn)
grads = grads_fn(model_parameters, xt, yt)
grads
>>> {'b': Array([-153.29868], dtype=float32),
'w': Array([-2533.0576], dtype=float32)}
❸
❶ Теперь параметры модели представлены в форме словаря, а не массива.
❷ Функция модели изменена для работы со словарями.
❸ Теперь градиенты имеют форму словаря.
Как можно видеть, дифференцирование по стандартным контей
нерам Python работает превосходно.
Возврат вспомогательных данных из функции
Функция, передаваемая в трансформацию grad(), должна возвра
щать скаляр, так как эта трансформация определена только для
скалярных функций. Иногда требуется возврат промежуточных ре
зультатов, но в этом случае функция возвращает кортеж, и grad() не
работает.
Предположим, что в примере линейной регрессии из листинга 4.9
необходим возврат результатов прогнозирования для фиксации их
в журнале. На рис. 4.5 показана схема этой процедуры.
Параметры
модели
Входные
данные (x)
Функция
потерь
Целевые
значения (y)
Значение
потерь
Прогноз
(ŷ)
grad(has_aux=True)
трансформация
со вспомогательными
данными
Параметры
модели
Входные
данные (x)
Целевые
значения (y)
Градиент
функции
потерь
Градиенты
Прогноз
(ŷ)
Рис. 4.5 Получение функции
для вычисления градиентов
с возвратом вспомогательных
данных
Глава 4
154
Вычисление градиентов
Изменение функции loss_fn() с целью возврата некоторых до
полнительных данных приводит к тому, что трансформация grad()
возвращает ошибку. Чтобы устранить ее, мы сообщаем трансформа
ции grad() о том, что функция потерь возвращает некоторые вспо
могательные данные.
Листинг 4.11
Возврат вспомогательных данных из функции
model_parameters = jnp.array([1., 1.])
def model(theta, x):
w, b = theta
return w * x + b
def loss_fn(model_parameters, x, y):
prediction = model(model_parameters, x)
return jnp.mean((prediction-y)**2), prediction
grads_fn = jax.grad(loss_fn, has_aux=True)
grads, preds = grads_fn(model_parameters, xt, yt)
model_parameters -= learning_rate * grads
❶
❷
❸
❶ Исходная функция потерь из листинга 4.9 теперь также возвращает результаты
прогнозирования.
❷ Параметр has_aux сообщает, что функция возвращает пару (out, aux).
❸ Теперь функция градиента возвращает градиенты и вспомогательные данные,
в нашем случае – прогнозы.
Параметр has_aux информирует трансформацию grad() о том,
что исходная функция возвращает пару (out, auxiliary_data). Это
позволяет grad() игнорировать дополнительный возвращаемый
параметр (auxiliary_data), передавая его напрямую пользователю,
и дифференцировать функцию loss_fn(), как если бы возвращал
ся только первый параметр (out). Если для параметра has_aux уста
новлено значение True, то возвращается пара (gradient, auxiliary_
data).
Для нейронных сетей удобно, если используемая модель имеет не
которое внутреннее состояние, которое необходимо сопровождать,
например статистику выполнения в варианте BatchNorm.
Получение градиента и значения функции
Часто возникает другая ситуация, в которой необходимо получить
значения функции потерь для отслеживания прогресса обучения.
Градиенты нужны для алгоритма градиентного спуска, но также
важно знать, насколько точен текущий набор весов и как ведут себя
значения потерь. Разумеется, можно отдельно определить качество
текущего решения, вычисляя потери и прочие метрики после обнов
ления тренировки. Но потери уже были вычислены во время обнов
155
Вычисление градиентов с использованием autodiff
ления тренировки, поэтому выполнение такого вычисления дваж
ды неоптимально (особенно для крупных нейронных сетей). Схема
процесса изображена на рис. 4.6.
Параметры
модели
Входные
данные (x)
Функция
потерь
Целевые
значения (y)
Значение
потерь
Прогноз
(ŷ)
Рис. 4.6 Получение функции
для вычисления градиентов вместе
со значением функции потерь
и вспомогательными данными
value_and_grad(has_aux=True)
трансформация
со вспомогательными
данными
Значение
потерь
Параметры
модели
Входные
данные (x)
Градиент
функции
потерь
Целевые
значения (y)
Градиенты
Прогноз
(ŷ)
Для этого варианта существует функция value_and_grad(), кото
рую мы уже использовали в главе 2.
Листинг 4.12 Возврат градиентов, значений и вспомогательных
данных
model_parameters = jnp.array([1., 1.])
def model(theta, x):
w, b = theta
return w * x + b
def loss_fn(model_parameters, x, y):
prediction = model(model_parameters, x)
return jnp.mean((prediction-y)**2), prediction
grads_fn = jax.value_and_grad(loss_fn, has_aux=True)
(loss, preds), grads = grads_fn(model_parameters, xt, yt)
model_parameters -= learning_rate * grads
❶
❷
❶ Теперь необходимы значения и градиенты.
❷ В дополнение к градиентам и вспомогательным данным возвращаются значения
потерь.
В рассматриваемом здесь примере объединены значения, гради
енты и вспомогательные данные. В этом случае возвращается кор
теж ((value, auxiliary_data), gradient). Это кортеж (value, gradient) без вспомогательных данных.
Глава 4
156
Вычисление градиентов
Поэлементные градиенты
В обычной рабочей среде машинного обучения модель тренируется
с применением процедуры градиентного спуска и с использовани
ем пакетов данных. Вы получаете градиент для всего пакета в целом
и обновляете параметры модели соответствующим образом.
В большинстве случаев такой степени детализации вполне до
статочно, но в некоторых ситуациях требуются градиенты уровня
выборки. Один из способов получения градиента такого уровня –
установка размера пакета равным 1. Но с точки зрения вычислений
такой подход неэффективен, поскольку для современного аппарат
ного оборудования, такого как GPU или TPU, наиболее предпочти
тельным является режим выполнения массовых вычислений с пол
ной загрузкой всех внутренних вычислительных компонентов.
Интересный вариант использования градиентов по отдельной
выборке возникает, когда необходимо оценить важность отдель
ных образцов данных и выбрать из них элементы с высокой ве
личиной градиента как более важные для использования в специ
альном алгоритме тренировки для выборки с приоритетами, для
реализации технологии hard sample (example) mining (пока нет
устоявшегося русскоязычного термина) или для выделения точек
данных для дальнейшего анализа. Поскольку эти числа существу
ют внутри вычислений градиента, многие библиотеки объединяют
и накапливают градиенты непосредственно по пакетам, поэтому
пользователю не предоставлен простой способ получения такой
информации.
В JAX это делается достаточно просто. Вспомним пример из гла
вы 2 – классификацию цифр MNIST. Там мы написали функцию фор
мирования одного прогноза, затем создали пакетную версию той же
функции с помощью vmap(), наконец, в функции потерь агрегирова
ли отдельные значения потерь. Если бы потребовались отдельные
значения потерь, их с легкостью можно было бы получить.
Ниже приведен набор инструкций для получения градиентов по
отдельной выборке:
1 создать
функцию для формирования прогноза по отдельной
выборке, например predict(x);
2 создать функцию для вычисления градиентов для прогноза
по отдельной выборке с помощью трансформации grad(), т. е.
grad(predict(x));
3 создать пакетную версию функции вычисления градиента с по
мощью трансформации vmap(), т. е. vmap(grad(predict))(x_
batch). Здесь применяется порядок трансформаций, обратный
описанному в главе 2 для тренировки нейронной сети. Там мы
просто выполнили grad(vmap(predict)), что позволило полу
чить градиент для пакета. Здесь мы получаем пакет градиентов;
Вычисление градиентов с использованием autodiff
157
4 дополнительно
(необязательно): компиляция полученной
функции с помощью трансформации jit() для более эффек
тивного использования внутреннего аппаратного обеспечения,
т. е. в итоге получаем jit(vmap(grad(predict)))(x_batch).
Остановка вычисления градиентов
Иногда нет необходимости в том, чтобы градиенты проходили че
рез поток некоторого подмножества графа вычислений. Для этого
могут существовать различные причины: возможно, имеются неко
торые вспомогательные вычисления, которые требуется выполнить,
но весьма нежелательно, чтобы они воздействовали на вычисления
градиентов. Например, это могут быть некоторые дополнительные
метрики или журналируемые данные, не являющиеся функциональ
но чистыми. Возможно, необходимо обновить переменные модели
другими значениями потерь, поэтому при вычислении одного из
значений потерь требуется исключить переменные, обновляемые
другими значениями потерь (такая операция также может быть вы
полнена со стороны оптимизатора; существуют способы, позволяю
щие определить, какие параметры должны обновляться, а какие не
должны). При обучении с подкреплением часто требуется подобная
функциональность.
В приведенном ниже простом примере (см. листинг 4.13) вычис
ления для функции f(x,y) = x**2 + y**2 формируют граф вычислений
(направленный граф, определяющий порядок вычислений) с двумя
ветвями: в одной вычисляется x**2, в другой – y**2. Каждая ветвь
независима от другой (см. рис. 4.7).
x
x2
x2 + y 2
y
y2
f (x, y)
Здесь мы останавливаем
вычисление градиента
Рис. 4.7 Граф вычислений для примера f(x, y) = x2 + y2. Ветви вычислений x2
и y2 не зависят друг от друга, и в данном примере желательно, чтобы поток
вычислений градиентов проходил только по одной из ветвей
Допустим, что вы не хотите, чтобы вторая ветвь воздействовала
на вычисляемые градиенты. Для подобных случаев существует спе
циальное средство. Функция jax.lax.stop_gradient(), применяемая
к некоторым входным переменным, работает как трансформация
идентификации, возвращая неизмененный аргумент. В то же вре
мя она предотвращает прохождение потока вычисления градиентов
при прямом или обратном режиме autodiff. Эта функция похожа на
метод detach() в PyTorch.
Глава 4
158
Листинг 4.13
Вычисление градиентов
Остановка потока вычисления градиентов
def f(x, y):
return x**2 + jax.lax.stop_gradient(y**2)
jax.grad(f, argnums=(0,1))(1.0, 1.0)
>>> (Array(2., dtype=float32, weak_type=True),
Array(0., dtype=float32, weak_type=True))
❶
❷
❸
❸
❶ Использование функции stop_gradient() для запрещения вычисления градиен-
та по y.
❷ Вычисление градиентов в конкретно заданной точке.
❸ Вывод полученных градиентов показывает, что они были вычислены только по
первой переменной (x).
Мы пометили часть вычислений, связанную со вторым парамет
ром исходной функции, с помощью специальной функции stop_
gradient(). Последующее вычисление градиента дает нам нулевые
градиенты по второму параметру исходной функции. Без вызова
stop_gradient() мы получили бы градиент 2.0.
4.2.4
Производные более высоких порядков
В JAX можно вычислять производные более высоких порядков. Бла
годаря функциональной природе JAX трансформация grad() выпол
няет преобразование функции в другую функцию, вычисляющую
производную исходной функции. Результатом работы grad() также
является функция, и этот процесс можно повторить несколько раз
для получения производных более высоких порядков.
Вернемся к простой функции, которую мы дифференцировали
в начале главы: f(x) = x4 + 12x + 1/x. Нам известна ее производная:
f ′(x) = 4x3 + 12 – 1/x2. Вторая производная выглядит так: f ′′(x) = 12x2 +
2/x3. Третья производная: f ′′′(x) = 24x – 6/x4 и т. д. Вычислим эти про
изводные в JAX.
Листинг 4.14
Определение производных более высоких порядков
def f(x):
return x**4 + 12*x + 1/x
f_d1 = jax.grad(f)
f_d2 = jax.grad(f_d1)
f_d3 = jax.grad(f_d2)
x = 11.0
print(f_d1(x))
print(f_d2(x))
print(f_d3(x))
❶
❷
❸
159
Вычисление градиентов с использованием autodiff
>>> 5335.9917
>>> 1452.0015
>>> 263.9996
❶ Взятие первой производной от заданной функции.
❷ Вторая производная.
❸ Третья производная.
Эти трансформации можно объединить в одной инструкции:
f_d3 = jax.grad(jax.grad(jax.grad(f)))
Мы вычислили производные в заданной точке (x = 11.0), но следует
помнить о том, что JAX возвращает функцию, которую можно при
менять везде, где она является корректной. В следующем примере
(листинг 4.15) мы берем несколько последовательных производных
от другой функции и изображаем их в виде графика (см. рис. 4.8).
Листинг 4.15 Построение графиков производных более высоких
порядков
def f(x):
return x**3 + 12*x + 7*x*jnp.sin(x)
x = np.linspace(-10, 10, num=500)
❶
fig, ax = plt.subplots(figsize=(10,10))
ax.plot(x, f(x), label = r"$y = x^3 + 12x + 7x*sin(x)$")
df = f
for d in range(3):
df = jax.grad(df)
ax.plot(x, jax.vmap(df)(x),
label=f"{['1st','2nd','3rd'][d]} derivative")
ax.legend()
❷
❸
❶ Дифференцирование другой функции.
❷ Последовательное вычисление производных (в цикле).
❸ Использование vmap() для применения производной функции к вектору значе-
ний.
Здесь vmap() используется для того, чтобы сделать производную
функцию применимой к вектору значений. Можно применить ис
ходную функцию к вектору значений с помощью средств широко
вещания и векторизации NumPy. Функции, полученные в результате
трансформации grad(), определяются только для исходных функций
со скалярным выводом, потому мы используем vmap() для вектори
зации этой функции по конкретному измерению массива. Более
подробно мы рассмотрим применение трансформации vmap() в гла
ве 6. На рис. 4.8 показаны результаты вычислений из листинга 4.15.
Глава 4
160
1000
Вычисление градиентов
y = x3 + 12x + 7x ∗ sin(x)
Первая производная
Вторая производная
Третья производная
500
0
–500
–1000
–10,0
–7,5
–5,0
–2,5
0
2,5
5,0
7,5
10,0
Рис. 4.8 Вычисление функции f(x) = x3 + 12x + 7x × sin(x) и трех ее
производных
Кроме того, в JAX можно применять оптимизацию более высо
ких порядков, такую как обучаемую оптимизацию, или метаопти
мизацию, при которой необходимо выполнять дифференцирование
по обновлениям градиентов. Любопытный ресурс, посвященный
обучаемой оптимизации, размещен в репозитории Google: https://
github.com/google/learned_optimization. Здесь вы найдете набор ин
тересных примеров, начиная с оптимизатора с обучаемыми пара
метрами.
Другой интересный пример – независимое от модели метаобуче
ние (model-agnostic meta-learning – MAML) (https://arxiv.org/abs/
1703.03400), пытающееся обучить начальный набор весов модели,
который можно быстро адаптировать к новым задачам. Качествен
ное руководство по метаобучению можно найти здесь: https://blog.
evjang.com/2019/02/maml-jax.html.
4.2.5
Вариант со многими переменными
Мы начали с простой функции для линейной регрессии с одним ска
лярным элементом входных данных (единственным числом) и од
ним скалярным элементом выходных данных. Потребовалась только
одна производная, и мы вычислили ее, используя функцию, полу
ченную с помощью трансформации grad().
Вычисление градиентов с использованием autodiff
161
Затем мы рассмотрели более обобщенный вариант с нескольки
ми входными параметрами, но выходным данным осталось един
ственное скалярное число. Мы вычислили частные производные по
нескольким входным переменным с помощью параметра argnums
трансформации grad(). Существует еще более обобщенный вариант,
когда входными и выходными данными исходной функции являют
ся векторы (несколько значений).
Матрица Якоби
Для вычисления нескольких частных производных по множеству
входных переменных, вероятнее всего, потребуется матрица Якоби.
Матрица Якоби (Jacobian matrix), или просто якобиан (Jacobian), –
это матрица, содержащая все частные производные для функции от
нескольких переменных, возвращающей результат в виде вектора
значений:
В контексте глубокого обучения якобиан для модели нейронной
сети должен содержать частные производные функции потерь по
параметрам модели. Цепочечное правило вычислений позволяет
выразить такой подход в форме произведения матриц Якоби, содер
жащих частные производные результата каждого слоя по выходным
данным предшествующего слоя.
Для вычисления полных матриц Якоби существуют две функции:
jacfwd() и jacrev(). Они вычисляют одни и те же значения, но их
реализация различна и зависит от прямого и обратного режимов
autodiff соответственно. (Мы рассмотрим эти режимы более под
робно в разделе 4.3.) Первая функция jacfwd() более эффективна
для «высоких» матриц Якоби, где количество выводимых элементов
(строк в матрице Якоби) значительно больше количества входных
переменных (столбцов в матрице Якоби). Функция jacrev() более
эффективна для «широких» матриц Якоби, где количество входных
переменных существенно превышает количество выводимых эле
ментов. Для матриц, форма которых близка к квадрату, вероятнее
всего, jacfwd() является более правильным выбором (как мы уви
дим в следующем примере).
Создадим простую векторную функцию с двумя выводимыми
элементами и тремя входными переменными.
Глава 4
162
Листинг 4.16
Вычисление градиентов
Вычисление якобиана для векторной функции
def f(x):
return [
x[0]**2 + x[1]**2 - x[1]*x[2],
x[0]**2 - x[1]**2 + 3*x[0]*x[2]
]
print(jax.jacrev(f)(jnp.array([3.0, 4.0, 5.0])))
>>> [Array([ 6.,
Array([21., -8.,
3., -4.], dtype=float32),
9.], dtype=float32)]
print(jax.jacfwd(f)(jnp.array([3.0, 4.0, 5.0])))
>>> [Array([ 6.,
Array([21., -8.,
3., -4.], dtype=float32),
9.], dtype=float32)]
❶
❷
❸
❹
%timeit -n 100 jax.jacrev(f)(jnp.array([3.0, 4.0, 5.0]))
>>> 29.9 ms ± 2.5 ms per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit -n 100 jax.jacfwd(f)(jnp.array([3.0, 4.0, 5.0]))
>>> 17.8 ms ± 336 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶
❷
❸
❹
Первый элемент вывода функции.
Второй элемент вывода функции.
Использование jacrev().
Использование jacfwd().
Здесь можно видеть, что обе функции вычисляют одинаковый
результат, но jacfwd() немного быстрее, чем jacrev(), при трех
элементах входных данных и двух элементах результата. Если ко
личество элементов входных данных существенно превышает ко
личество элементов выходных данных, то jacrev() будет работать
быстрее, но здесь эти количества сопоставимы, и победителем ста
новится jacfwd().
По умолчанию обе функции вычисления якобиана выполняют
дифференцирование по первому параметру функции. В рассматри
ваемом здесь примере все параметры передаются в одном векторе,
но вместо вектора может потребоваться использование функции
с несколькими отдельными параметрами. В такой ситуации можно
воспользоваться уже знакомым параметром argnums.
Матрица Гессе
Если требуется производная второго порядка функции с нескольки
ми элементами входных данных, то получаем матрицу Гессе (Hes
sian matrix), или просто гессиан (Hessian). Гессиан – это квадратная
матрица, содержащая все вторые (частные) производные:
Вычисление градиентов с использованием autodiff
163
Для вычисления плотных матриц Гессе используется функция
hessian().
Листинг 4.17
Вычисление гессиана функции
def f(x):
return x[0]**2 - x[1]**2 + 3*x[0]*x[2]
jax.hessian(f)(jnp.array([3.0, 4.0, 5.0]))
❶ Функция с тремя параметрами.
❷ Вычисление гессиана в конкретно заданной точке.
❶
❷
Гессиан функции f – это якобиан градиента той же функции, по
этому вычисление гессиана можно реализовать следующим спосо
бом:
def hessian(f):
return jacfwd(grad(f))
В JAX вычисление гессиана реализовано по-другому с предполо
жением о том, что трансформация grad() реализована с использова
нием обратного режима:
def hessian(f):
return jacfwd(jacrev(f))
Эта функция представляет собой обобщение обычного определе
ния гессиана, которое поддерживает вложенные контейнеры Python
(и структуры pytree) как входные и выходные данные.
Во многих случаях полный гессиан не требуется. Вместо него мож
но использовать произведение гессиана на вектор в некоторых вы
числениях без материализации всей матрицы Гессе.
Теперь вы хорошо знакомы с несколькими компонентами, необ
ходимыми для понимания прямого и обратного режимов, поэтому
наступило время поглубже заглянуть во внутренний механизм auto
diff и изучить его более подробно.
164
4.3
Глава 4
Вычисление градиентов
Прямой и обратный режимы autodiff
Этот раздел предназначен для тех, кто хочет понять основы внут
реннего механизма autodiff. Здесь будут рассматриваться прямой
и обратный режимы autodiff и две соответствующие трансформации
JAX: jvp() и vjp().
Для понимания внутреннего механизма autodiff и связанных
с ним трансформаций требуются некоторые знания из области ма
тематического анализа. Вы должны знать, что такое производная,
частная производная и производная по направлению (наклонная
производная). Также необходимо уметь применять правила диффе
ренцирования.
Как мне кажется, это самая трудная часть книги. Поэтому вам,
вероятно, придется прочитать ее несколько раз для полного пони
мания и самостоятельно решить предложенные упражнения. Не от
ступайте, если поначалу это покажется слишком сложным. Это дей
ствительно сложная тема. В большинстве случаев можно продолжать
использовать JAX без понимания внутреннего механизма autodiff.
Но наградой за знание, приобретенное с большими трудностями,
станет лучшее понимание того, как использовать JAX наиболее эф
фективно.
Дополнительные источники для изучения autodiff
Существует несколько ресурсов, которые помогут освоить эту сложную
тему.
Во-первых, это великолепное короткое видеоруководство «What Is
Automatic Differentiation?», автор Ари Сефф (Ari Seff) (https://www.you
tube.com/watch?v=wG_nF1awSSY). Ари описывает простыми словами
все нюансы автоматического дифференцирования, и это очень полезно.
Во-вторых, подробная статья «Automatic Differentiation in Machine Lear
ning: A Survey», авторы: Атилим Гунеш Байдин (Atilim Gunes Baydin), Барак А. Перлмутер (Barak A. Pearlmuter), Алексей Андреевич Радул (Ale
xey Andreyevich Radul) и Джеффри Марк Зискинд (Jeffrey Mark Siskind)
(https://arxiv.org/abs/1502.05767). Эта статья предоставляет читателю
гораздо более широкий контекст и помогает глубже понять подробности, касающиеся autodiff.
Если вы хотите понимать внутренний механизм autodiff еще глубже,
то для этого потребуются фундаментальные книги, такие как «Evalua
ting Derivatives», авторы Андреас Гриванк (Andreas Griewank) и Андреа
Вальтер (Andrea Walther) (https://epubs.siam.org/doi/book/10.1137/
1.9780898717761), и «The Art of Differentiating Computer Programs»,
автор Уве Науманн (Uwe Naumann) (https://epubs.siam.org/doi/book/
10.1137/1.9781611972078).
Прямой и обратный режимы autodiff
165
Вернемся к основополагающей идее autodiff. Выполняемые нами
вычисления в конечном счете представляют собой композиции ко
нечного набора элементарных операций – сложения, умножения,
возведения в степень, использование тригонометрических функций
и т. д. Для таких элементарных операций известны все производные.
При вычислении, состоящем из элементарных операций, можно
комбинировать производные операций, являющихся элементами
вычисления, применяя цепочечное правило дифференцирования,
чтобы получить производную всей композиции. Байдин (Baydin)
и др. (https://arxiv.org/abs/1502.05767) утверждают в своей статье:
«Автоматическое дифференцирование можно воспринимать как
выполнение нестандартной интерпретации компьютерной про
граммы, где такая интерпретация подразумевает расширение кон
кретного стандартного вычисления с помощью вычисления различ
ных производных».
4.3.1
Трассировка оценок
Вычисление можно представить как трассировку оценок (evaluation
trace) элементарных операций, также называемую списком Венгер
та (Wengert list) (в настоящее время список также называют лентой
(tape); название взято из статьи Р. Е. Венгерта (R. E. Wengert), опуб
ликованной в 1964 г.; https://dl.acm.org/doi/10.1145/355586.364791),
со специфическими входными значениями. Выполняется декомпо
зиция вычисления (или функции) в последовательность элементар
ных функциональных шагов, и вводятся промежуточные перемен
ные (intermediary variables). Список Венгерта абстрагируется от всех
условий, управляющих потоком, т. е. autodiff «маскирует» любую
операцию, включая инструкции управления потоком, которые не
изменяют напрямую числовые значения. Выбранная ветвь заменяет
все условные инструкции, все циклы развертываются, а все вызовы
функций выполняются как встроенные (inlined).
Предположим, что имеется конкретная возвращающая действи
тельные значения функция с двумя переменными f(x1, x2) = x1×x2 +
sin(x1×x2) для вычисления производной и трассировки оценок для
этой функции (см. табл. 4.1). Этот пример используется для демонст
рации прямого и обратного режимов autodiff.
Данное вычисление также можно представить в виде графа вы
числений (см. рис. 4.9).
Предположим, что необходимо вычислить частную производную
этой функции по первой переменной x1 в некоторой точке, напри
мер (x1, x2) = (7.0, 2.0). Начнем с использования прямого режима
autodiff.
Глава 4
166
Вычисление градиентов
Таблица 4.1 Трассировка оценок (список Венгерта) для функции
f(x1, x2) = x1×x2 + sin(x1×x2)
Трассировка оценок
Входные переменные
v–1 = x1
v0 = x2
Промежуточные переменные
v1 = v–1 × v0
v2 = sin(v1)
v3 = v1 + v2
Выходные переменные
y = v3
Рис. 4.9
4.3.2
Граф вычислений для примера функции f(x1, x2) = x1×x2 + sin(x1×x2)
Прямой режим и jvp()
Прямой накопительный режим autodiff (или режим касательной) яв
ляется самым простым с точки зрения концепции.
Вычисления в прямом режиме
Для вычисления производной функции по первой переменной x1 на
чинаем со связывания с каждой промежуточной переменной vi ее
производной dvi/dxi (в приведенном ниже примере производная для
краткости обозначена как v’i). Вместо одного значения для каждой
переменной vi получаем кортеж (vi, v’i). Исходные промежуточные
значения называются первичными (primals), а производные – каса
тельными (tangents). В совокупности это формирует метод двойных
чисел (dual number approach). Цепочечное правило вычисления про
изводных применяется к каждой элементарной операции в трасси
ровке прямого режима касательной (см. табл. 4.2).
Один проход по функции теперь вычисляет не только результат
самой функции (элемент выходных данных, здесь 14,99), но также ее
производную по x1 (здесь 2,274). Выполненные вручную вычисления
в табл. 4.2 можно проверить с помощью JAX.
167
Прямой и обратный режимы autodiff
Таблица 4.2 Пример прямого режима autodiff для функции
f(x1, x2) = x1×x2 + sin(x1×x2) с вычислением в точке (7.0, 2.0) по первой переменной x1
Трассировка первичных значений
в прямом режиме
Трассировка касательных (производных)
в прямом режиме
v–1 =x1
v0 =x2
v′–1 = x′1
v′0 =x′2
= 7.0
= 2.0
v1
v2
v3
=v–1 × v0 = 14.0
= sin(v1) = 0.99
=v1 + v2 = 14.99
y
= v3
= 14.99
(начало)
= 1.0
= 0.0
(начало)
v′1 =v′–1 × v0 + v–1 × v′0 = 1.0 × 2.0 + 7.0 × 0.0
= 2.0 × 0.137
v′2 =v′1 × cos(v1)
v′3 =v′1 + v′2
= 2.0 + 0.274
(конец)
y′
= v′3
= 2.274
(конец)
Листинг 4.18 Проверка выполненных вручную вычислений
в прямом режиме
def f(x1,x2):
return x1*x2 + jnp.sin(x1*x2)
❶
x = (7.0, 2.0)
❷
jax.grad(f)(*x)
❸
>>> Array(2.2734745, dtype=float32, weak_type=True)
❶ Исходная функция.
❷ Точка, в которой вычисляется производная.
❸ Вычисление градиента по первому параметру x1.
Полученные числа почти совпадают. При вычислениях вручную
мы получили 2,274, а JAX возвращает более точный ответ 2,2734745.
Меньшая точность полученного вручную значения связана с округ
лением некоторых промежуточных результатов в процессе вычис
лений.
Системы autodiff используют перегрузку операторов или метод
преобразования исходного кода для реализации таких вычислений.
В JAX применяется перегрузка операторов. Тема реализации систем
autodiff слишком объемна, чтобы рассматривать ее в этой книге.
Теперь предположим, что функция выводит несколько элементов
данных. Можно вычислить частные производные для каждого вы
ходного элемента в одном прямом проходе. Но при этом потребует
ся выполнение прямого прохода для каждой входной переменной.
В приведенном выше примере очевидно, что мы получили только
производную по x1. Для x2 необходимо выполнить отдельный пря
мой проход.
Для функции в общем виде f: Rn → Rm проход в прямом режиме
для одной входной переменной вычисляет один столбец соответ
Глава 4
168
Вычисление градиентов
ствующего якобиана функции (частные производные для каждого
элемента вывода по конкретной заданной входной переменной).
Полный якобиан вычисляется за n этапов вычислений. Вышеупомя
нутая функция jacfwd() делает именно это.
Поэтому должно быть интуитивно понятно, что прямой режим
предпочтительнее, если количество выводимых элементов значи
тельно больше количества входных переменных, т. е. m >> n, или так
называемые «высокие» якобианы.
Производная по направлению и jvp()
Производная по направлению, или наклонная производная (direc
tional derivative), обобщает форму записи частных производных.
Частные производные вычисляют наклон (угловой коэффициент)
в положительном направлении оси, представленной конкрет
ной переменной. Мы использовали входной вектор касательной
(v–1, v0) = (1.0, 0.0) для частной производной по первой переменной.
Но наклон можно вычислять в любом направлении. Для этого необ
ходимо задать конкретное направление, определив его с помощью
вектора (u1, u2), указывающего направление, в котором требуется
вычислить наклон. Производная по направлению становится рав
нозначной частной производной, когда этот вектор указывает в по
ложительном направлении x1 или x2 и выглядит как (1.0, 0.0) или
(0.0, 1.0).
Вычисление производной по направлению с использованием
autodiff выполняется просто. Нужно просто передать вектор направ
ления как начальное значение для касательных. Вот и все. В резуль
тате получаем производную по направлению в заданной точке.
В еще более обобщенном виде можно вычислить произведение
якобиана на вектор (Jacobian-vector product – JVP) без вычисления
самого якобиана просто за один проход в прямом режиме. Опреде
ляем для входного вектора касательных (v–1, v0) интересующий нас
вектор и производим вычисления в прямом режиме autodiff. Эта
операция называется jvp() по первым буквам Jacobian-vector pro
duct.
Функция jvp() принимает следующие параметры:
функцию fun, которую необходимо продифференцировать;
значения primals, в которых должен быть вычислен якобиан
функции fun;
вектор касательных tangents, для которых требуется вычислить
произведение якобиана на вектор.
Результатом является пара (primals_out, tangents_out), где primals_out – функция fun, примененная к primals (значение исходной
функции в заданной точке), а tangents_out – произведение якобиа
169
Прямой и обратный режимы autodiff
на на вектор для функции, вычисленное по primals с заданными
tangents. Например, это может быть производная по направлению
в заданном направлении или частная производная по конкретной
переменной.
Можно сказать, что для заданной функции f, вектора входных
данных x и вектора касательных v jvp() вычисляет как выходные
данные функции f(x), так и переменную по направлению ∂f(x)v:
(x, v) → ( f(x), ∂f(x)v).
Если вы знакомы с сигнатурами типов языка Haskell, то описан
ную выше функцию можно записать так:
jvp :: (a -> b) -> a -> T a -> (b, T b)
Это означает, что jvp – имя функции. Первым параметром этой
функции является другая функция с сигнатурой (a -> b), преобразу
ющая значение типа a в тип b. Второй параметр имеет тип a (здесь
это вектор первичных значений). Третий параметр, обозначенный
как T a, – тип касательный для типа a. Последний тип является ти
пом возвращаемого значения (b, T b). Он состоит из двух элемен
тов: типа выходных первичных значений b и типа соответствующих
касательных T b.
СОВЕТ Качественные вводные курсы по сигнатурам ти
пов языка Haskell можно найти здесь: https://en.wikibooks.
org/wiki/Haskell/Type_basics#Functional_types, https://learnhaskell.blog/03-html/02-type_signatures.html и https://learny
ouahaskell.com/types-and-typeclasses.
В листинге 4.19 вычисляется произведение якобиана на вектор
(JVP) для функции из листинга 4.16 (вычисление якобиана).
Листинг 4.19
Вычисление jvp()
def f2(x):
return [
x[0]**2 + x[1]**2 - x[1]*x[2],
x[0]**2 - x[1]**2 + 3*x[0]*x[2]
]
❶
x = jnp.array([3.0, 4.0, 5.0])
v = jnp.array([1.0, 1.0, 1.0])
❷
❸
p,t = jax.jvp(f2, (x,), (v,))
p
❹
❺
Глава 4
170
Вычисление градиентов
>>> [Array(5., dtype=float32), Array(38., dtype=float32)]
t
>>> [Array(5., dtype=float32), Array(22., dtype=float32)]
❻
Функция, которая использовалась в примере вычисления якобиана.
Вектор первичных значений; значение, передаваемое в функцию.
Вектор касательных, по которым вычисляются производные по направлению.
Обратите внимание: преобразование массивов типа jnp.array в кортежи с одним
элементом.
❺ Значение функции в заданной точке f(x).
❻ Производная по направлению в направлении вектора v.
❶
❷
❸
❹
Функция jvp() ожидает передачи первичных значений и каса
тельных, которые должны быть либо кортежем, либо списком аргу
ментов; в данном случае функция не работает с типом Array, поэто
му мы вручную упаковали массивы jnp.array в кортеж.
В приведенном выше примере вычислено значение функции f(x)
и производная по направлению за один проход. Можно восстано
вить все частные производные (якобиан) за три прохода, передавая
векторы [1.0, 0.0, 0.0], [0.0, 1.0, 0.0] и [0.0, 0.0, 1.0] как касательные.
Листинг 4.20
Восстановление столбцов якобиана с помощью jvp()
p,t = jax.jvp(f2, (x,), (jnp.array([1.0, 0.0, 0.0]),))
t
❶
>>> [Array(6., dtype=float32), Array(21., dtype=float32)]
p,t = jax.jvp(f2, (x,), (jnp.array([0.0, 1.0, 0.0]),))
t
❶
>>> [Array(3., dtype=float32), Array(-8., dtype=float32)]
p,t = jax.jvp(f2, (x,), (jnp.array([0.0, 0.0, 1.0]),))
t
❶
>>> [Array(-4., dtype=float32), Array(9., dtype=float32)]
❶ Передача векторов элементов для восстановления отдельных столбцов яко
биана.
Теперь мы знаем, что такое JVP, и можем использовать функцию
jvp() для проверки вручную вычисленных результатов в прямом ре
жиме.
Листинг 4.21 Проверка вычисленных в прямом режиме результатов
с помощью JVP
def f(x1,x2):
return x1*x2 + jnp.sin(x1*x2)
x = (7.0, 2.0)
❶
❷
171
Прямой и обратный режимы autodiff
p,t = jax.jvp(f, x, (1.0, 0.0))
p
>>> Array(14.990607, dtype=float32, weak_type=True)
t
>>> Array(2.2734745, dtype=float32, weak_type=True)
❸
❹
❺
❶ Функция, которую мы использовали при вычислениях вручную.
❷ Та же точка, в которой необходимо вычислить производную.
❸ Использование JVP с тем же вектором касательных, что и в вычислениях вручную
в табл. 4.2.
❹ Первичные значения (выходные данные функции).
❺ Значения касательных (производная по x1).
Мы подробно рассмотрели JVP, и наш пример вычислений вруч
ную преобразуется в код с использованием функции jvp(). В дан
ном случае первичные значения и касательные уже являются кор
тежами, поэтому нет необходимости в каком-либо преобразовании,
и они передаются напрямую в функцию jvp().
4.3.3
Обратный режим и vjp()
Прямой режим эффективен, если количество элементов входных
данных намного меньше количества элементов результата (выход
ных данных). В машинном обучении мы обычно получаем противо
положную ситуацию: количество входных данных велико, а на выхо
де наблюдается всего лишь несколько элементов. Обратный режим
autodiff решает эту проблему: распространяет производные в обрат
ном направлении от вывода в соответствии с обобщенным алгорит
мом обратного распространения (backpropagation algorithm).
Вычисления в обратном режиме
В обратном режиме процесс autodiff состоит из двух фаз. В первой
фазе исходная функция выполняется в прямом направлении. Про
межуточные переменные (значения, получаемые при формирова
нии трассировки оценок) заполняются в этой фазе, и все зависимо
сти в графе вычислений фиксируются. Эти вычисления показаны
в левом столбце табл. 4.3 сверху вниз.
Во второй фазе каждая промежуточная переменная vi дополняется
сопряженным значением (или котангенсом) v′i = dyj /dvi. Это производ
ная j-го выходного элемента yj по vi, представляющая чувствитель
ность yj к изменениям в vi. Производные вычисляются в обратном
направлении посредством распространения сопряженных значений
(котангенсов) v′i от выходных данных к входным. Этот процесс по
казан в правом столбце табл. 4.3 снизу вверх с вычислением про
Глава 4
172
Вычисление градиентов
изводных по мере продвижения по графу в обратном направлении.
Например, если вы видите, что переменная v′1 обновляется дважды
в правом столбце, то должны считать это двумя последовательными
обновлениями: первое – в нижней части ячейки таблицы, второе –
в верхней.
Таблица 4.3 Пример обратного режима autodiff для функции
f(x1, x2) = x1×x2 + sin(x1×x2) с вычислением в точке (7.0, 2.0)
Трассировка первичных значений
в прямом направлении
Трассировка сопряженных значений (производных)
в обратном направлении
v–1 =x1
v0 =x2
x′1 = v′–1
x′2 = v′0
= 7.0
= 2.0
v1
=v–1 × v0 = 14.0
v2
= sin(v1) = 0.99
v3
=v1 + v2 = 14.99
y
=
v3
= 14.99
(начало)
v′–1
v′0
v′1
v′1
v′2
(конец)
= 2.274
= 7.959
(конец)
= v′1 dv1 /dv–1 = v′1 × v0 = 1.137 × 2.0 = 2.274
=v′1 dv1 /dv0 = v′1 × v–1 = 1.137 × 7.0 = 7.959
=v′1 + v′2 dv2 /dv1 = v′1 + v′2 × cos(v1) = 1.0 + 1.0 × 0.137 = 1.137
=v′3 dv3 /dv1 = v′3 × 1 = 1.0
=v′3 dv3 /dv2 = v′3 × 1 = 1.0
v′3 =y′
= 1.0
(начало)
В рассматриваемом здесь примере после прямого прохода (он ни
чем не отличается от прохода в прямом режиме) мы выполняем об
ратный проход, начиная с v′3 = y′ = 1,0. Если некоторые переменные
воздействуют на вывод несколькими способами, то мы суммируем
их вклад по этим различным путям, как в примере для v′1. В конце
процедуры мы получаем все производные dy/dx1 = x′1 и dy/dx2 = x′2 за
один обратный проход.
Напомню, что можно вычислять производные по обеим перемен
ным функций, используя параметр argnums для проверки правиль
ности вычислений.
Листинг 4.22
Проверка вычислений вручную в обратном режиме
def f(x1,x2):
return x1*x2 + jnp.sin(x1*x2)
❶
x = (7.0, 2.0)
❷
jax.grad(f, argnums=(0,1))(*x)
❸
>>> (Array(2.2734745, dtype=float32, weak_type=True),
Array(7.9571605, dtype=float32, weak_type=True))
❶ Все та же функция.
❷ И та же точка.
❸ Но теперь мы вычисляем градиенты по параметрам x1 и x2.
И в этом случае полученные числа почти совпадают. При вычис
лениях вручную мы получили значение 2,274 для производной по x1
Прямой и обратный режимы autodiff
173
и 7,595 для производной по x2. JAX возвращает более точные значе
ния 2,2734745 и 7,9571605.
Как можно видеть, мы одновременно вычислили производные
по обеим входным переменным за один обратный проход. Поэто
му преимущество обратного режима autodiff заключается в том, что
с точки зрения вычислений он обходится дешевле, чем прямой ре
жим для функций с многочисленными элементами входных данных,
когда n >> m. Но в обратном режиме предъявляются более высокие
требования к хранению значений. В экстремальном варианте, где
используется функция с n входными переменными и единственным
выходным элементом, обратный режим вычисляет все производные
за один проход, тогда как прямой режим требует n проходов. В ма
шинном обучении многочисленные параметры и единственный ска
лярный результат – дело обычное, поэтому обратный режим (и об
ратное распространение) более предпочтителен при таких условиях.
Тем не менее существуют экспериментальные методики использо
вания прямого режима autodiff для градиентного спуска (см. «Gradi
ents Without Backpropagation», https://arxiv.org/abs/2202.08587).
Обобщение обратного режима и vjp()
По аналогии со способом получения произведения якобиана на
вектор без вычисления матриц в прямом режиме, обратный режим
можно использовать для вычисления транспонированного произ
ведения якобиана на вектор, или, что равнозначно, произведения
вектора на якобиан (vector-Jacobian product – VJP), инициализируя
обратную фазу конкретным заданным смежным значением (или
котангенсом). Такой подход применяется для формирования мат
риц Якоби последовательно по одной строке, и он эффективен для
так называемых «широких» якобианов. Такая операция называется
vjp() по первым буквам термина vector-Jacobian product.
Функция vjp() принимает следующие параметры:
функцию fun, которую необходимо продифференцировать;
значения primals, по которым должен вычисляться якобиан
функции fun.
Результатом становится пара (primals_out, vjpfun), где primals_
out – функция fun, примененная к значениям primals (значение ис
ходной функции в конкретной точке), vjpfun – функция преобразо
вания из вектора сопряженных значений (котангенсов), имеющего
ту же форму, что и primals_out, в кортеж сопряженных значений
(котангенсов), имеющий ту же форму, что и primals, и представля
ющий произведение вектора на якобиан функции fun, вычисленное
в точках primals.
Приведенное выше описание немного сложно для понимания.
Попробуем описать такой подход другим способом и рассмотрим
Глава 4
174
Вычисление градиентов
исходный код. Это означает, что для заданной функции f и входного
вектора x JAX-трансформация vjp() формирует выходные данные
самой функции f(x) и функцию для вычисления VJP с заданными со
пряженными значениями (котангенсами) для обратной фазы.
Используя сигнатуры типов языка Haskell, можно записать тип
vjp() в следующем виде:
vjp :: (a -> b) -> a -> (b, CT b -> CT a)
Здесь, как и в методе JVP, передается исходная функция с типом
(a -> b) и входное значение x типа a. В эту функцию не передают
ся сопряженные значения. Вывод также различен. Первым возвра
щаемым значением остаются выходные первичные значения, т. е.
результат вычисления f(x) с типом b. Вторым значением является
функция для обратного прохода, принимающая сопряженное зна
чение типа b и возвращающая сопряженное значение типа a.
А теперь рассмотрим код, использующий vjp() для той же функ
ции, которая вычислялась в табл. 4.3, чтобы проверить результаты
вычислений вручную в обратном режиме.
Листинг 4.23 Проверка результатов вычислений вручную
в обратном режиме с помощью VJP
def f(x1,x2):
return x1*x2 + jnp.sin(x1*x2)
❶
x = (7.0, 2.0)
❷
p,vjp_func = jax.vjp(f, *x)
p
>>> Array(14.990607, dtype=float32, weak_type=True)
vjp_func(1.0)
❸
❹
❺
>>> (Array(2.2734745, dtype=float32, weak_type=True),
Array(7.9571605, dtype=float32, weak_type=True))
❶
❷
❸
❹
❺
Функция, которая использовалась в вычислениях вручную.
Та же точка, в которой необходимо вычислить производную.
Параметры исходной функции передаются как отдельные параметры.
Первичные значения (вывод функции).
Передаются те же сопряженные значения, которые использовались при вычислениях вручную в обратной фазе.
Мы получили производные по обоим параметрам функции за
один обратный проход.
Для более сложной функции с двумя элементами выходных дан
ных воспользуемся функцией из примера с якобианом и восстано
вим полный якобиан.
175
Прямой и обратный режимы autodiff
Листинг 4.24
Восстановление строк якобиана с помощью vjp()
def f2(x):
return [
x[0]**2 + x[1]**2 - x[1]*x[2],
x[0]**2 - x[1]**2 + 3*x[0]*x[2]
]
x = jnp.array([3.0, 4.0, 5.0])
p,vjp_func = jax.vjp(f2, x)
p
>>> [Array(5., dtype=float32), Array(38., dtype=float32)]
vjp_func([1.0, 0.0])
>>> (Array([ 6.,
3., -4.], dtype=float32),)
vjp_func([0.0, 1.0])
>>> (Array([21., -8.,
9.], dtype=float32),)
❶
❶
❶ Передача векторов элементов для восстановления отдельных столбцов яко
биана.
Здесь нам пришлось дважды вызывать функцию vjp_func(), полу
ченную из трансформации vjp(), так как существуют два выходных
элемента этой функции.
4.3.4
Материалы для более глубокого изучения
Можно получить гораздо больше информации по использованию
autodiff в JAX. Если вы хотите более глубоко изучить эту тему, то ва
шему вниманию предлагается несколько источников. Во-первых,
в самой документации JAX есть превосходный материал «Autodiff
Cookbook» (https://docs.jax.dev/en/latest/notebooks/autodiff_cookbook.
html). Если нужно узнать больше о математических основах и под
робностях реализации JVP и VJP, о том, как вычислять произведения
гессиана на вектор, дифференцировать комплексные числа и т. д.,
начните с этого великолепного источника информации.
Для определения специализированных правил дифференцирова
ния в JAX прочтите «Custom Derivative Rules for JAX-Transformable
Python Functions» (https://docs.jax.dev/en/latest/notebooks/Custom_
derivative_rules_for_Python_code.html). Здесь вы узнаете об использо
вании jax.custom_jvp() и jax.custom_vjp() для определения специа
лизированных правил дифференцирования для функций Python,
которые уже являются трансформируемыми в JAX.
Еще одно полезное руководство: «How JAX Primitives Work»
(https://docs.jax.dev/en/latest/jax-primitives.html). В этом документе
Глава 4
176
Вычисление градиентов
описан интерфейс, который обязательно должен поддерживать ба
зовый элемент JAX, чтобы обеспечить выполнение всех трансфор
маций JAX. Это весьма полезно, если вы намереваетесь определять
новые экземпляры core.Primitive вместе со всеми соответствующи
ми правилами трансформации.
Руководство «Autodidax: JAX Core From Scratch» (https://docs.jax.
dev/en/latest/autodidax.html) описывает основные принципы функ
ционирования ядра системы JAX, включая описание внутренней
работы autodiff. Также существует отличное руководство по autodiff
(и Autograd, предшественнику autodiff) Мэтта Джонсона (Matt John
son): https://videolectures.net/videos/deeplearning2017_johnson_auto
matic_differentiation.
Наконец, не забывайте заглядывать в раздел по autodiff в пакете
JAX (https://docs.jax.dev/en/latest/jax.html#automatic-differentiation)
в общедоступной документации по API.
Резюме
Существуют различные способы вычисления производных: диф
ференцирование вручную, символьное дифференцирование,
численное дифференцирование и autodiff (автоматическое диф
ференцирование).
Автоматическое дифференцирование (autodiff, или просто AD) –
это продуманная до мельчайших подробностей изощренная ме
тодика определения градиентов для вычислений, выраженных
как код.
Autodiff способен вычислять градиенты даже для весьма сложных
компьютерных программ и их элементов, включая управляющие
структуры, ветвления кода, циклы и рекурсию, которые сложно
представить как конечные выражения (выражения в замкнутой
форме).
В JAX градиенты вычисляются с помощью трансформации grad().
По умолчанию grad() вычисляет производную по первому пара
метру исходной функции, но вы можете управлять этим поведе
нием с помощью параметра argnums.
Если из функции необходимо возвращать дополнительные вспо
могательные данные, то используйте параметр has_aux в grad()
и других связанных с ней трансформациях.
Можно получить градиент и значение функции с помощью транс
формации value_and_grad().
Для вычисления производных более высоких порядков восполь
зуйтесь последовательностью трансформаций grad().
Функции jacfwd() и jacrev() вычисляют матрицы Якоби, а функ
ция hessian() вычисляет матрицы Гессе.
Резюме
177
Autodiff имеет два режима: прямой и обратный.
Прямой режим вычисляет градиенты всех выходных элементов
функции по одному элементу входных данных за один проход.
Это эффективный способ, если количество элементов выходных
данных существенно больше количества входных переменных.
Обратный режим вычисляет градиент для одного элемента вы
ходных данных функции по всем элементам входных данных
функции за один проход, поэтому более эффективен для функций
с многочисленными входными данными и несколькими элемен
тами выходных данных.
Функция jvp() вычисляет произведение якобиана на вектор (JVP)
без вычисления самого якобиана за один прямой проход.
Функция vjp() вычисляет транспонированное произведение яко
биана на вектор (равнозначное произведению вектора на якоби
ан, или VJP) за один проход в обратном режиме autodiff.
5
Компиляция кода
Темы главы:
JIT-компиляция для создания кода для CPU, GPU или TPU;
внутреннее устройство JIT-компилятора: промежуточные
представления и компиляторы ускоренной линейной
алгебры;
ограничения JIT-компиляции.
В главе 1 мы сравнивали производительность простой функции JAX
на CPU и GPU с использованием и без использования JIT-компиля
ции. В главе 2 мы использовали JIT для компиляции двух функций
в тренировочном цикле для простой нейронной сети. Поэтому с JITкомпиляцией вы уже немного знакомы. Эта процедура компилирует
функцию для целевой аппаратной платформы и ускоряет ее.
Начиная с главы 4 мы приступаем к изучению трансформаций JAX
(напомню, в JAX практически все представляет собой компонуемые
трансформации функций). Из предыдущей главы мы узнали об auto
diff и о трансформации grad(). В этой главе рассматривается компи
ляция и соответствующая трансформация jit(). В следующих главах
мы узнаем о других трансформациях, связанных с автоматической
векторизацией и распараллеливанием.
Компиляция кода
179
В области числовых расчетов JAX появился как весьма внушитель
ный фреймворк, основанный на базовой мощи компилятора XLA
компании Google. Компилятор XLA – это не просто еще один ин
струмент в вычислительном комплекте, он специально спроектиро
ван для создания эффективного кода для высокопроизводительных
вычислительных задач. Его следует воспринимать как архитектора,
тщательно разрабатывающего чертежи для строительства сооруже
ний, оптимизированных для их назначения.
Значение зависимости JAX от XLA весьма существенно. У XLA
имеется хорошо зарекомендовавшее себя наследие в области под
держки фреймворков машинного обучения, в частности TensorFlow.
Универсальность XLA особенно подчеркивается его совместимо
стью с рядом устройств: CPU как основы вычислений общего назна
чения, GPU как специализированных устройств, адаптированных
для параллельных вычислений, и TPU – узкоспециализированных
ускорителей компании Google, поддерживающих рабочие процес
сы машинного обучения. Такой широкий диапазон совместимо
сти подтверждает высокую адаптируемость JAX, показывая, что он
оптимизирован для разнообразных аппаратных сред. В вычисли
тельной сфере универсальный код не всегда преобразуется в оди
наковую производительность на всех платформах, но благодаря вы
сочайшему качеству XLA JAX стремится поддерживать постоянную
эффективность.
Помимо явной интеграции с XLA, JAX предлагает упреждаю
щий подход к компиляции: введена функция трансформации jit()
и недавно добавлены средства AOT-компиляции (AOT – ahead-oftime – предварительная, досрочная (компиляция)). Представьте
себе стратега, наносящего на карту наиболее эффективные марш
руты, гарантируя, что каждое путешествие будет оптимизировано
с точки зрения времени и ресурсов. JIT-компиляция служит ана
логичной цели для графов вычислений – упорядочивает, улучшает
и объединяет операции, когда и где это возможно, выполняет пре
образование последовательностей вычислений в одиночные, более
эффективные процедуры. Такая возможность не только повышает
скорость, но также усовершенствует весь процесс вычислений в це
лом, идентифицируя и исключая избыточные вычисления. Чтобы
получить преимущества от JIT, необязательно иметь в своем рас
поряжении GPU или TPU. JIT улучшает производительность даже на
обычном CPU.
Здесь первостепенное значение имеет оптимизация производи
тельности. Вне зависимости от того, выполняются ли вычисления на
CPU или на более мощной аппаратуре, JIT гарантирует максималь
ную эффективность. Благодаря взаимовыгодной интеграции JAX,
XLA и JIT разработчик не просто пишет код, он создает инженерное
решение, ведущее к совершенству вычислений.
Глава 5
180
Компиляция кода
В этой главе рассматриваются механизмы JIT и AOT-компиляции,
описываются способы их эффективного использования и отмечают
ся ограничения.
5.1
Использование компиляции
Эта глава полностью посвящена компиляции, и в ней описывается,
как пользоваться компиляцией, чтобы ускорить работу кода. Мы бу
дем использовать простые учебные примеры, чтобы понять меха
низм компиляции в JAX и выделить его функциональные возмож
ности. Преимущества компиляции особенно заметны при массовых
и длительных вычислениях – тренировках крупной нейронной сети,
климатической модели или другой крупномасштабной вычисли
тельной задаче. Но компиляция может оказаться полезной и для не
слишком большой задачи. Например, ускорение фильтрации изо
бражений может иметь огромное значение для пользователей ва
шего приложения.
Начнем с хорошо известной функции активизации под названием
линейный блок масштабируемой экспоненциальной кривой (scaled
exponential linear units – SELU, https://arxiv.org/abs/1706.02515), ко
торая упоминалась в главе 2. Воспользуемся этой функцией для де
монстрации JIT-компиляции:
.
Код реализации этой функции приведен в листинге 5.1, где зна
чения констант alpha и scale взяты из документа SELU (см. выше).
Листинг 5.1
Функция активизации SELU
def selu(x,
alpha=1.6732632423543772848170429916717,
scale=1.0507009873554804934193349852946):
'''Scaled exponential linear unit activation function.'''
# Функция активизации линейным блоком масштабируемой
# экспоненциальной кривой.
return scale * jnp.where(x > 0, x, alpha * jnp.exp(x) - alpha)
❶ Ядро функции SELU.
❶
Функция содержит две ветви. Если аргумент x положителен, то
возвращается значение x, умноженное на коэффициент scale. Иначе
применяется более сложная формула scale * (alpha * ex – alpha). Функ
ция SELU часто используется в нейронных сетях как нелинейная ак
тивизация.
Использование компиляции
5.1.1
181
Использование JIT-компиляции
Предположим, что имеется миллион активизаций, для которых не
обходимо применить функцию SELU. Хотя при одном прямом про
ходе по обычной нейронной сети вы получаете всего лишь миллион
вычислений этой функции, в процессе тренировки такие прямые
проходы повторяются многократно, поэтому наш пример не так уж
далек от действительности. Воспользуемся JIT и сравним произво
дительность функции SELU с JIT-компиляцией и без нее.
Использование jax.jit как трансформации
или аннотации
Из глав 1 и 2 вы уже знаете, что существует трансформация jit()
и соответствующая аннотация @jit для компиляции выбранной
функции.
JIT и AOT
Существуют два типа компиляции: JIT (just-in-time) и AOT (ahead-of-time).
JIT-компиляция выполняется при работе программы во время выполнения. Код компилируется, когда он необходим. В JAX это происходит,
когда код, помеченный как компилируемый, выполняется в первый раз.
Код трассируется с использованием абстрактных значений, представляющих входные массивы, и компилируется, а затем реальные массивы
передаются в скомпилированную функцию для выполнения. Поэтому
первый проход по компилируемой функции выполняется медленнее,
чем при последующих вызовах.
AOT-компиляция, или статическая компиляция, преобразует программу
на языке высокого уровня в код на языке более низкого уровня до начала выполнения программы. Весь код компилируется заранее независимо от того, потребуется ли он в дальнейшем. Такой подход уменьшает
общую рабочую нагрузку во время выполнения и перемещает интенсивную процедуру компиляции на этап сборки программы.
JAX изначально получил широкую известность благодаря наличию собственной JIT-компиляции. Но в настоящее время JAX предоставляет некоторые опции и для AOT-компиляции (более подробно об этом в подразделе 5.2.3).
Трансформация jax.jit() принимает чистую функцию и возвра
щает версию в обертке исходной функции, подготовленную для JITкомпиляции.
Глава 5
182
Компиляция кода
Листинг 5.2 Сравнение производительности версий SELU
с JIT-компиляцией и без нее
x = jax.random.normal(jax.random.PRNGKey(42), (1_000_000,))
selu_jit = jax.jit(selu)
%timeit -n100 selu(x).block_until_ready()
>>> 944 µs ± 1.24 per loop
(mean ± std. dev. of 7 runs, 100 loops each)
# (ср.знач ± станд.откл. по 7 проходам по 100 циклов в каждом)
%timeit -n100 selu_jit(x).block_until_ready()
>>> The slowest run took 29.85 times longer than
the fastest. This could mean that an intermediate
result is being cached.
# Самый медленный проход выполнялся в 29.85 раза дольше,
# чем самый быстрый. Возможно, это означает, что
# промежуточный результат кешировался.
>>> 173 µs ± 283 µs per loop
(mean ± std. dev. of 7 runs, 100 loops each)
# (ср.знач ± станд.откл. по 7 проходам по 100 циклов в каждом)
❶
❷
❸
❹
❶
❷
❸
❹
Генерация миллиона случайных чисел.
Получение JIT-трансформированной версии исходной функции.
Измерение скорости выполнения исходной функции (без JIT).
Измерение скорости выполнения JIT-трансформированной функции.
В приведенном выше примере можно видеть, что JIT-скомпили
рованная версия исходной функции более чем в пять раз быстрее,
чем версия без компиляции. В то же время обе версии используют
GPU (в моем варианте NVIDIA A100-SXM4-40GB).
СОВЕТ Для получения максимальной производительности
от JAX применяйте jax.jit() к самым первым внешним вы
зовам функций.
Здесь необходимо особо отметить, что вызов jax.jit() не приво
дит к немедленной компиляции функции. Он выполняет всю необ
ходимую подготовку функции к JIT-компиляции, которая будет вы
полнена при первом вызове функции. Подробности этого процесса
рассматриваются в подразделе 5.1.2 и в разделе 5.2. Поэтому здесь
эталонные тесты немного некорректны, поскольку первый проход по
функции медленнее, чем последующие, так как компиляция проис
ходит во время первого вызова. Выведенное предупреждение также
подтверждает возможность возникновения такой ситуации: самый
медленный проход почти в 30 раз медленнее, чем самый быстрый.
В реальной среде скомпилированная версия функции работает даже
еще быстрее, чем сообщают эти усредненные числовые показатели.
183
Использование компиляции
Также можно воспользоваться аннотацией @jit перед функци
ей. В приведенном ниже коде (см. листинг 5.3) компилируется та
же функция selu() с помощью аннотации @jit. И в этом случае вы
полняются должным образом эталонные тесты с «разогревающим»
вызовом для инициализации компиляции, а затем тестируется уже
скомпилированная функция. Мы будем пропускать этап «разогре
ва» во многих будущих примерах, чтобы сделать код более компакт
ным, а кроме того, еще и потому, что во многих случаях выполнение
компиляции во время первого прохода по функции является вполне
обычным делом. Но во время эталонного тестирования следует пом
нить об этой особенности.
Листинг 5.3 Использование JIT в виде аннотации
@jax.jit
❶
def selu(x,
alpha=1.6732632423543772848170429916717,
scale=1.0507009873554804934193349852946):
'''Scaled exponential linear unit activation function.'''
# Функция активизации линейным блоком масштабируемой экспоненциальной кривой.
return scale * jnp.where(x > 0, x, alpha * jnp.exp(x) - alpha)
z = selu(x)
%timeit -n100 selu(x).block_until_ready()
❷
❸
>>> 60.6 µs ± 7.34 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
# (ср.знач ± станд.откл. по 7 проходам по 100 циклов в каждом)
❶ Использование аннотации для JIT-компиляции исходной функции.
❷ «Разогревающий» проход по функции.
❸ Измерение скорости выполнения JIT-скомпилированной функции.
При таком способе не нужна отдельная функция в обертке; это
чудесное преобразование скрыто от пользователя, и вы получаете
исходную функцию уже в надлежащей «упаковке». При наличии
этапа «разогрева» и корректных измерений мы видим, что реаль
ные числовые характеристики намного лучше, и функция работает
в 15 раз быстрее, чем некомпилированная версия из предыдущего
примера.
Для трансформации jit() существует несколько параметров,
и мы рассмотрим некоторые из них в следующем подразделе.
Компиляция и выполнение на специализированном
аппаратном оборудовании
С аппаратурой, которую вы намереваетесь использовать, связаны
два параметра. Внимание: оба представляют собой эксперименталь
ные функциональные характеристики, и API может измениться.
Глава 5
184
Компиляция кода
Параметр backend – строка, представляющая внутренний аппарат
ный компонент XLA: 'cpu', 'gpu' или 'tpu'. Если в системе доступен
GPU или TPU, то JAX будет использовать этот компонент вместо CPU.
Параметр device определяет устройство, на котором будет рабо
тать JIT-скомпилированная функция. Можно передать конкретное
устройство, полученное как результат вызова jax.devices(). Обыч
но по умолчанию это первый элемент возвращенного массива.
Если в системе имеется более одного GPU или TPU, то параметр device обеспечивает точное управление, позволяя конкретно указать
устройство выполнения вычислений.
В приведенном ниже примере мы преднамеренно компилируем
функцию для CPU и GPU. Здесь компилируются отдельные версии
функции для CPU и GPU, и можно наблюдать почти 10-кратное раз
личие в скорости между функциями selu_jit_cpu() и selu_jit_gpu().
Листинг 5.4 Управление, позволяющее выбирать внутренний аппаратный
компонент для использования
def selu(x,
❶
alpha=1.6732632423543772848170429916717,
❶
scale=1.0507009873554804934193349852946):
❶
'''Scaled exponential linear unit activation function.'''
❶
# Функция активизации линейным блоком масштабируемой экспоненциальной кривой.
❶
return scale * jnp.where(
x > 0, x, alpha * jnp.exp(x) - alpha)
❶
selu_jit_cpu = jax.jit(selu, backend='cpu')
selu_jit_gpu = jax.jit(selu, backend='gpu')
%timeit -n100 selu(x).block_until_ready()
❷
❸
❹
>>> 791 µs ± 66.8 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit -n100 selu_jit_cpu(x).block_until_ready()
❺
>>> 1.81 ms ± 178 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit -n100 selu_jit_gpu(x).block_until_ready()
❻
>>> The slowest run took 12.60 times longer
than the fastest. This could mean that an
intermediate result is being cached.
# Самый медленный проход выполнялся в 12.60 раза дольше,
# чем самый быстрый. Возможно, это означает, что
# промежуточный результат кешировался.
>>> 165 µs ± 247 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶ Повторное создание исходной функции, чтобы исключить ее взаимодействие с @jit-анноти
рованной версией.
❷ Версия, ориентированная на CPU.
185
Использование компиляции
❸ Версия, ориентированная на GPU.
❹ Вызов исходной версии без JIT-компиляции. Эта версия продолжает использование GPU/
TPU, если эти устройства доступны.
❺ Вызов функции на CPU.
❻ Вызов функции на GPU (помните о самом медленном первом проходе).
Измерения при вызове selu() могут оказаться искаженными,
так как JAX, возможно, будет продолжать использовать GPU или
TPU, если они доступны в системе. Единственное различие состоит
в том, что измеряемая функция не скомпилирована с помощью XLA.
В моей системе доступен GPU, и тензор с данными размещается на
GPU, поэтому вычисления также производятся на этом устройстве.
Для получения корректного эталонного теста необходимо размес
тить тензоры данных на соответствующих устройствах с помощью
метода device_put() (как это было сделано в главе 3), а затем выпол
нить исходную функцию.
Листинг 5.5 Управление используемым внутренним аппаратным компонентом
и размещением тензоров на конкретном устройстве
x_cpu = jax.device_put(x, jax.devices('cpu')[0])
x_gpu = jax.device_put(x, jax.devices('gpu')[0])
%timeit -n100 selu(x_cpu).block_until_ready()
❶
❷
❸
>>> 2.74 ms ± 95.8 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit -n100 selu(x_gpu).block_until_ready()
❹
>>> 872 µs ± 80.8 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit -n100 selu_jit_cpu(x_cpu).block_until_ready()
❺
>>> 437 µs ± 4.29 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit -n100 selu_jit_gpu(x_gpu).block_until_ready()
❻
>>> 27.1 µs ± 4.94 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶
❷
❸
❹
❺
❻
Размещение тензора данных на CPU.
Размещение тензора данных на GPU.
Измерение нескомпилированной функции на CPU.
Измерение нескомпилированной функции на GPU.
Измерение JIT-скомпилированной функции на CPU.
Измерение JIT-скомпилированной функции на GPU.
Здесь измерения выполнены на той же системе, которая исполь
зовалась ранее. Можно видеть, что данные, размещаемые на GPU,
дают в результате приблизительно то же время, что и простой вызов
нескомпилированной функции, тогда как для данных, размещенных
на CPU, получение результата замедляется более чем в три раза. Обе
186
Глава 5
Компиляция кода
JIT-скомпилированные версии работают быстрее, чем соответствую
щие нескомпилированные вызовы. Для CPU скорость выше в 6 раз,
для GPU – в 32 раза. Следует отметить, что JAX (по крайней мере те
кущая версия 0.3.17, с которой я работаю) не запрещает использо
вать GPU-скомпилированную функцию с данными, размещенными
на CPU, и наоборот.
Использование статических аргументов
В JAX механизм статических аргументов полезен в различных вари
антах использования. Об этом механизме вы узнаете более подроб
но немного позже в разделе 5.2, где рассматривается трассировка.
Простой пример, в котором может потребоваться применение ста
тических аргументов, – необходимость компиляции функции с пара
метрами, не являющимися массивами, а представленными, напри
мер, экземпляром некоторого класса или функции. Скажем, нужно
скомпилировать слой нейронной сети с функцией активации, пере
даваемой как параметр. В этом случае JIT выдаст сообщение об ошиб
ке, поскольку аргументы и возвращаемое значение компилируемой
функции должны быть массивами, скалярами или стандартными
контейнерами Python (tuple, list или dict, возможно, вложенными).
Статические аргументы могут устранить эту проблему, так как аргу
менты, объявленные как статические, могут быть чем угодно, при ус
ловии что они хешируемые и для них определена операция равенства.
Поэтому можно помечать некоторые аргументы как статические,
или константы времени компиляции, с помощью параметров static_argnums или static_argnames. Оба параметра являются необяза
тельными. Параметр static_argnums – это целое число или набор
целых чисел, определяющих, какие позиционные аргументы интер
претируются как статические. При отсутствии параметров static_
argnums и static_argnames никакие аргументы не воспринимаются
как статические.
С технической точки зрения это означает, что операции, завися
щие только от статических аргументов, будут выполнять свертыва
ние констант в Python во время трассировки. Свертывание констант
(constant folding) – это методика оптимизации, позволяющая исклю
чать выражения, вычисляющие значения, которые можно опреде
лить еще до начала выполнения кода (т. е. во время компиляции).
В листинге 5.6 показана реализация такого примера.
Листинг 5.6 Исправление ошибки JIT-компиляции
с помощью статических аргументов
def dense_layer(x, w, b, activation_func):
return activation_func(x*w+b)
❶
x = jnp.array([1.0, 2.0, 3.0])
❷
187
Использование компиляции
w = jnp.ones((3,3))
b = jnp.ones(3)
dense_layer_jit = jax.jit(dense_layer)
dense_layer_jit(x, w, b, selu)
>>> ...
>>> ----> 3 dense_layer_jit(x, w, b, selu)
>>> ...
>>> TypeError: Cannot interpret value of type
<class 'function'> as an abstract array; it does
not have a dtype attribute
# TypeError: Невозможно интерпретировать значение типа
# <class 'function'> как абстрактный массив; это
# значение не имеет атрибута dtype.
dense_layer_jit = jax.jit(dense_layer, static_argnums=3)
dense_layer_jit(x, w, b, selu)
>>> Array([[2.101402, 3.152103, 4.202804],
>>>
[2.101402, 3.152103, 4.202804],
>>>
[2.101402, 3.152103, 4.202804]], dtype=float32)
❷
❷
❸
❹
❺
❻
Функция с параметром в виде другой функции.
Некоторые тестовые данные.
Создание JIT-скомпилированной версии исходной функции.
Компиляция завершается ошибкой, потому что четвертым аргументом является
функция, т. е. тип, не допускаемый в JAX.
❺ Пометка четвертого аргумента как статического.
❻ Теперь JIT-компиляция завершается успешно.
❶
❷
❸
❹
Здесь мы передали функцию активации как параметр для функ
ции dense_layer(). Так как функция не является допустимым типом
для входных и выходных значений JIT-компилируемой функции,
мы получили ошибку. А поскольку действительная компиляция про
исходит во время первого вызова исходной функции, ошибка появ
ляется только при первом ее вызове, а не во время вызова jit(). Для
устранения этой ошибки конкретный параметр (ставший причиной
ошибки) помечается как статический. После определения функции
активации как статической JAX узнает о необходимости компиля
ции одной версии исходной функции для различных функций ак
тивации и о необходимости создания скомпилированных функций,
принимающих переменные x, w и b.
В другом примере (см. листинг 5.7) создается функция для вычис
ления расстояния Минковского (Minkowski distance) между двумя
векторами. Одним из параметров, принимаемых этой функцией,
является order. При значении order=1 вычисляется манхэттенское
расстояние (Manhatten distance), а если задано значение order=2,
Глава 5
188
Компиляция кода
то вычисляется евклидово расстояние (Euclidian distance). Функция
превосходно компилируется без каких-либо статических парамет
ров, но если параметр order является статическим, то обеспечивает
создание специализированных скомпилированных версий функ
ции. Поэтому при вызове функции с order=1 ее версия компилиру
ется с постоянно установленным значением 1. Если функция вы
зывается с order=2, то компилируется другая версия, для которой
постоянно устанавливается значение 2.
Листинг 5.7
Пометка аргумента как статического
def dist(order, x, y):
print("Compiling")
return jnp.power(jnp.sum(jnp.abs(x-y)**order), 1.0/order)
❶
❷
dist_jit = jax.jit(dist, static_argnums=0)
❸
dist_jit(1, jnp.array([0.0, 0.0]), jnp.array([2.0, 2.0]))
>>> Compiling
>>> Array(4., dtype=float32)
dist_jit(2, jnp.array([0.0, 0.0]), jnp.array([2.0, 2.0]))
>>> Compiling
>>> Array(2.828427, dtype=float32)
dist_jit(1, jnp.array([10.0, 10.0]),
jnp.array([2.0, 2.0]))
>>> Array(16., dtype=float32)
❹
❺
❻
❶ Функция с тремя параметрами. Первый параметр должен иметь только ограни-
ченное количество значений.
❷ Побочный эффект, который должен становиться видимым при каждой компи
ляции.
❸ Объявление JIT-компиляции исходной функции и объявление ее первого парамет
ра статическим.
❹ Компиляция функции для заданного значения параметра и ее выполнение.
❺ Компиляция функции для другого заданного значения параметра и ее выпол
нение.
❻ Для этого значения параметра функция уже скомпилирована.
Вызов функции, объявленной для JIT-компиляции, с различными
значениями статических параметров приводит к повторной компи
ляции. В рассматриваемом здесь примере функция скомпилирова
лась дважды: для значения первого параметра 1 и для значения 2
того же параметра. Мы преднамеренно использовали побочный эф
фект, который становился видимым только во время первого про
хода по функции. В разделе 5.2 объясняется, почему это происходит.
Использование компиляции
189
Такие специализированные скомпилированные версии могут
упрощать полученный в итоге скомпилированный код. Например,
для order=1 не требуется возведение чисел в квадрат и извлечение
квадратного корня, и получаемый в итоге код упрощается. Кроме
того, этот код может выполняться быстрее (поскольку он проще и бо
лее специализирован). Подобный подход имеет смысл при ограни
ченном количестве возможных значений статического параметра.
Далее в книге мы будем встречаться и с другими вариантами,
в которых имеет смысл использование статических параметров, на
пример в листинге 5.15.
Если необходимо определить статические аргументы при исполь
зовании jit как декоратора (или аннотации), то можно применить
средство Python functools.partial.
Листинг 5.8 Использование functools.partial с jit в качестве
декоратора
from functools import partial
❶
@partial(jax.jit, static_argnums=0)
def dist(order, x, y):
return jnp.power(jnp.sum(jnp.abs(x-y)**order), 1.0/order)
❶ Создание частично применимой функции jax.jit() с аргументом static_arg-
nums=0.
Функция partial() создает функциональный объект, который
при вызове будет вести себя как исходная функция с заданными па
раметрами. Здесь создается частично применяемая функция jax.
jit() с фиксированным параметром static_argnums=0. Эта частично
определенная функция применяется к функции dist() как аннота
ция, компилируя ее с первым аргументом (order), который является
статическим.
Аргументы, связанные с оптимизацией
Для трансформации jit() существуют два аргумента, которые мож
но использовать для различных оптимизаций. Здесь эти аргументы
не будут описываться во всех подробностях, я просто кратко отмечу
их наличие, чтобы вы не были обескуражены при встрече с ними.
Первый аргумент для функции jit() называется donate_argnums.
Он помечает конкретные позиционные аргументы, которые могут
быть переданы («безвозмездно предоставлены» – donated) в вычис
ление. Обоснование следующее: входные аргументы используют
некоторые буферы памяти для хранения значений. Во время вы
полнения функции они могут перестать быть необходимыми в не
который момент после использования этих значений в вычислении.
190
Глава 5
Компиляция кода
Вы можете «просто отдать» эти буферы памяти, и XLA воспользуется
ими для сокращения объема памяти, требуемого для компиляции,
например для сохранения результата. Разумеется, после передачи
вы не должны снова использовать «пожертвованные» буферы памя
ти. Если попытаться сделать это, то JAX выдаст сообщение об ошиб
ке. По умолчанию буферы аргументов никуда не передаются. Для
именованных аргументов существует аналогичный параметр donate_argnames. Более подробно о параметрах donate_* можно узнать
здесь: https://docs.jax.dev/en/latest/faq.html#buffer-donation.
Второй аргумент keep_unused (для него по умолчанию установле
но значение False) управляет тем, как JAX интерпретирует аргумен
ты, определяемые как неиспользуемые. По умолчанию JAX исклю
чает такие аргументы из полученного в результате XLA-компиляции
кода. Они не будут передаваться на устройство и не предоставляются
во внутренний механизм выполнения. Эту логику можно отключить,
если вы не хотите отсекать такие параметры.
Также существует опция для встраивания (inline) функции в объ
емлющее представление Jaxpr (промежуточное представление, в ко
торое преобразуется код Python; более подробно о Jaxpr см. в раз
деле 5.2), позволяющее исключить накладные расходы на вызовы
функций. Для этого используется параметр inline=True. По умол
чанию для этого параметра установлено значение False, и вызов
функции представлен как применение базисного элемента xla_call
с собственным представлением sub-Jaxpr. При встраивании отдель
ные вызовы функции не создаются. Код функции вставляется в том
месте, где она вызывается, что позволяет исключить накладные рас
ходы на вызов функции и возврат значений из нее.
5.1.2
Чистые функции и процесс компиляции
JIT-компиляция работает с JIT-совместимыми функциями. JAX
предназначен для работы с функционально чистым кодом без гло
бального состояния и побочных эффектов. Вы можете продолжать
писать и выполнять функции, не являющиеся чистыми, но JAX не
гарантирует их корректную работу.
JAX выполняет преобразование функций Python, сначала транс
лируя их код в простой промежуточный язык под названием Jaxpr,
используя процесс трассировки (tracing) (мы подробно рассмотрим
Jaxpr и процесс трассировки в разделе 5.2). Затем трансформации
работают с представлением Jaxpr. При необходимости компиляции
представление Jaxpr продолжает процедуру компиляции с помощью
XLA (на этом этапе выполняется больше шагов, которые мы также
рассмотрим в разделе 5.2). Общая схема высокого уровня процесса
представлена на рис. 5.1.
191
Использование компиляции
Код Python
Jaxpr
Компиляция
с помощью XLA
Рис. 5.1 Схема высокого уровня, показывающая, как JAX компилирует код
Компиляция выполняется во время первого вызова функции
(отсюда название just-in-time компиляция). JAX производит трас
сировку функции, выполняя ее код, затем создает представление
Jaxpr, компилирует его с помощью XLA и кеширует скомпилиро
ванный код.
У этого процесса имеется несколько последствий. Первый проход
по функции происходит медленнее, чем следующие. Поэтому для
правильного измерения производительности необходимо «разо
греть» ее с помощью первого выполнения, чтобы произошла компи
ляция. После этого можно измерять производительность функции
без учета издержек на компиляцию.
Все, что происходит во время первого прохода по функции, со
храняется в ее представлении. Если поведение функции отличается
при следующих проходах (такое может случиться, если функция не
является чистой), вы не увидите этого в скомпилированной версии.
Кроме того, побочные эффекты не фиксируются в Jaxpr, поэтому
результаты их воздействия вы наблюдаете только во время первого
прохода по функции. В листинге 5.9 показан пример того, что про
исходит, если функция не является чистой. В этой функции имеется
побочный эффект и используется глобальное состояние. Мы компи
лируем функцию и наблюдаем за ее поведением.
Листинг 5.9
Компиляция функции, не являющейся чистой
global_state = 1
def impure_function(x):
print(f'Side-effect: printing x={x}')
y = x*global_state
return y
impure_function_jit = jax.jit(impure_function)
impure_function_jit(10)
>>> Side-effect: printing x=Traced<ShapedArray(int32[],
weak_type=True)>with<DynamicJaxprTrace(level=1/0)>
>>> Array(10, dtype=int32, weak_type=True)
impure_function_jit(10)
>>> Array(10, dtype=int32, weak_type=True)
global_state = 2
❶
❷
❸
❹
❺
❺
❺
❻
❼
Глава 5
192
Компиляция кода
impure_function_jit(10)
>>> Array(10, dtype=int32, weak_type=True)
impure_function(10)
>>> Side-effect: printing x=10
>>> 20
❽
❾
Глобальное состояние, используемое в функции, не являющейся чистой.
Побочный эффект, создаваемый функцией, не являющейся чистой.
Использование глобального состояния.
Создание JIT-скомпилированной версии функции.
Наблюдение побочных эффектов во время первого прохода по функции.
Отсутствие побочных эффектов во время второго прохода.
Изменение глобального состояния.
Измененное глобальное состояние не влияет на скомпилированную версию
функции.
❾ Некомпилированная версия функции продолжает демонстрировать побочные
эффекты и воздействие глобального состояния.
❶
❷
❸
❹
❺
❻
❼
❽
В листинге 5.9 создана функция, не являющаяся чистой, с побоч
ным эффектом (инструкция print) и глобальным состоянием (пере
менная global_state). Затем функция компилируется, но при этом
в действительности создается обертка для нее, а реальная компи
ляция происходит позже, при первом вызове функции. Во время
первого вызова JAX выполняет трассировку кода функции, исполь
зуя абстрактные значения, представляющие входные массивы, и мы
наблюдаем побочный эффект. Но этот побочный эффект не вклю
чается в итоговое представление Jaxpr и в скомпилированный код
(ниже мы увидим его внутреннее представление), поэтому во время
второго вызова функции побочный эффект не наблюдается. Кроме
того, глобальное состояние также отсутствует в скомпилированном
коде, так что после его изменения функция ведет себя так, как если
бы глобальное значение оставалось прежним. Исходная некомпили
рованная версия функции ведет себя, как и ожидалось, демонстри
руя побочный эффект и воздействие глобального состояния.
Здесь мы использовали инструкции print(), чтобы проанализи
ровать, что происходит и когда. Но это всего лишь одна из деталей
реализации, заключающаяся в том, что код Python запускается как
минимум один раз. Не следует полагаться на это. Внимательно про
веряйте свой код и помните, что единственный правильный способ
использования JAX – только функционально чистые функции Python.
Чтобы лучше понять, почему JIT-компиляция работает именно
так, следует поглубже заглянуть в ее внутренний механизм. Это со
вершенно необязательно, если вы хотите просто использовать JITкомпиляцию, но понимание того, что происходит внутри, поможет
эффективно справляться с ограничениями JIT (это тема раздела 5.3)
и, возможно, окажется полезным при отладке в сложных случаях,
193
Внутренний механизм JIT
таких как слишком медленная компиляция и неэффективный код.
Я уверен, что такой подход также принесет немалую пользу тем про
граммистам, которые (как я) любят заглядывать внутрь и разбирать
ся в том, как работают программные механизмы.
5.2
Внутренний механизм JIT
Ранее уже неоднократно упоминались представление Jaxpr и ра
бочий процесс компиляции в JAX. Теперь мы более подробно рас
смотрим детали этих механизмов и познакомимся с описанием
отдельных шагов процесса компиляции. Сначала представлено пре
образование кода Python в Jaxpr, затем – преобразование представ
ления Jaxpr в код целевой платформы.
5.2.1
Jaxpr – промежуточное представление
для программ JAX
Сначала мы сосредоточимся на первом этапе компиляции, т. е. на
преобразовании кода Python в представление Jaxpr (см. рис. 5.2).
Код Python
Компиляция
с помощью XLA
Jaxpr
Рис. 5.2 Этот подраздел полностью посвящен преобразованию Python в Jaxpr
JAX выполняет преобразование исходного кода в промежуточное
представление вычисления перед какими-либо трансформациями
и передачей кода в XLA. Это промежуточное представление называ
ется Jaxpr – сокращение от JAX Expression. Затем в полученном пред
ставлении Jaxpr начинают работать трансформации.
Язык Jaxpr
По существу, Jaxpr представляет собой простой функциональный
язык с весьма ограниченными возможностями высшего порядка
(базисный элемент (primitive) имеет высший порядок, если он пара
метризован функцией). Код Jaxpr можно увидеть, воспользовавшись
трансформацией jax.make_jaxpr().
Листинг 5.10 Использование трансформации jax.make_jaxpr()
def f1(x, y, z):
return jnp.sum(x + y * z)
❶
x = jnp.array([1.0, 1.0, 1.0])
❷
Глава 5
194
Компиляция кода
y = jnp.ones((3,3))*2.0
z = jnp.array([2.0, 1.0, 0.0]).T
jax.make_jaxpr(f1)(x,y,z)
❸
>>> { lambda ; a:f32[3] b:f32[3,3] c:f32[3]. let
>>>
d:f32[1,3] = broadcast_in_dim[broadcast_dimensions=(1,) shape=(1, 3)]
c
>>>
e:f32[3,3] = mul b d
>>>
f:f32[1,3] = broadcast_in_dim[broadcast_dimensions=(1,) shape=(1, 3)]
a
>>>
g:f32[3,3] = add f e
>>>
h:f32[] = reduce_sum[axes=(0, 1)] g
>>> in (h,) }
❶ Простая функция с несколькими операциями.
❷ Некоторые тестовые данные.
❸ Генерация кода Jaxpr для исходной функции.
Здесь можно видеть код исходной функции, преобразованный
в представление Jaxpr. Код Jaxpr выводится с использованием следу
ющей грамматической формы:
jaxpr ::= { lambda Var* ; Var+.
let Eqn*
in [Expr+] }
Представление Jaxpr содержит один или несколько входных пара
метров, составляющих два списка, размещенных после ключевого
слова lambda и разделенных точкой с запятой: список констант (Var*
в грамматической форме; в листинге 5.10 этот список пуст) и список
входных переменных для функции Python (Var+ в грамматической
форме; в листинге 5.10 это список a:f32[3] b:f32[3,3] c:f32[3]). Лю
бые переменные, используемые функцией, но не являющиеся вход
ными параметрами, будут интерпретироваться как константные
значения и содержаться в списке констант. Существует список вы
ходных атомарных выражений (представленный в грамматической
форме как [Expr+], а в листинге 5.10 как (h,) – в действительности
это кортеж). Список равенств (выражений) (let Eqn*), определяющий
промежуточные переменные, ссылается на промежуточные выраже
ния. Каждое равенство определяет одну или несколько переменных
как результат применения базисного элемента (primitive) к некото
рым атомарным выражениям (например, g:f32[3,3] = add f e для вы
полнения операции сложения). Каждое равенство использует только
входные переменные и промежуточные переменные, определенные
предыдущими равенствами. В Jaxpr применяется явная типизация,
поэтому в рассматриваемом здесь примере можно видеть указание
типов и форм для каждой переменной.
Внутренний механизм JIT
195
Большинство базисных элементов Jaxpr принимают только одно
или несколько выражений Expr как аргументы (описание см. в до
кументации модуля jax.lax: https://docs.jax.dev/en/latest/jax.lax.
html#module-jax.lax). К таким базисным элементам относятся add,
sub, sin, mul и reduce_sum. Jaxpr также содержит несколько базисных
элементов высшего порядка, которые включают sub-Jaxprs. Среди
них можно выделить проверки условий switch и cond, циклы while
и fori, scan (упоминаемый в главе 3) и специальный базисный эле
мент xla_call, где инкапсулирован sub-Jaxpr вместе с параметрами,
определяющими внутренний компонент и устройство, на котором
должна выполняться компиляция (базисный элемент xla_call упо
минался при обсуждении встраивания в подразделе «Аргументы,
связанные с оптимизацией» раздела 5.1). Если вы интересуетесь бо
лее подробным описанием языка Jaxpr, то рекомендуется следую
щий документ: https://jax.readthedocs.io/en/latest/jaxpr.html.
Трансформация make_jaxpr() создает другую функцию, прини
мающую параметры исходной функции и возвращающую код Jaxpr
в значении типа jax.core.ClosedJaxpr (https://docs.jax.dev/en/latest/
jaxpr.html#jax-core-closedjaxpr). Тип jax.core.ClosedJaxpr содержит
код Jaxpr и константы. Jaxpr размещается в атрибуте jaxpr и имеет
тип jax.core.Jaxpr (https://docs.jax.dev/en/latest/_autosummary/jax.
extend.core.Jaxpr.html#jax.extend.core.Jaxpr). Константы находятся
в атрибуте consts, являющемся списком.
Трассировка
JAX выполняет преобразование в Jaxpr, используя трассировку. Во
время трассировки JAX обертывает каждый аргумент в специальный
объект-трассировщик класса jax.core.Tracer и использует его как
замену массива JAX Array для определения последовательности опе
раций, выполняемых функцией Python.
Трассировщик записывает все операции JAX, выполняемые с со
ответствующим аргументом во время вызова функции. Затем JAX
реконструирует функцию, используя записи трассировщика. Резуль
татом этой реконструкции становится представление Jaxpr. Побоч
ные эффекты Python продолжают действовать во время трассировки,
но трассировщики их не записывают, поэтому побочные эффекты не
проявляются в Jaxpr. В листинге 5.11 исходная функция изменена:
она содержит побочный эффект (инструкцию print()) и использует
глобальную переменную z.
Листинг 5.11 Изучение представления Jaxpr для функции с побочными
эффектами
x = jnp.array([1.0, 1.0, 1.0])
y = jnp.ones((3,3))*2.0
Глава 5
196
Компиляция кода
z = jnp.array([2.0, 1.0, 0.0]).T
def f2(x, y):
print(f'x={x}, y={y}, z={z}')
return jnp.sum(x + y * z)
f2_jaxpr = jax.make_jaxpr(f2)(x,y)
>>> x=Traced<ShapedArray(float32[3])>with<
DynamicJaxprTrace(level=1/0)>,
y=Traced<ShapedArray(float32[3,3])>with<
DynamicJaxprTrace(level=1/0)>, z=[2. 1. 0.]
f2_jaxpr.jaxpr
❶
❷
❸
❹
❺
>>> { lambda a:f32[3]; b:f32[3] c:f32[3,3]. let
>>>
d:f32[1,3] = broadcast_in_dim[broadcast_dimensions=(1,) shape=(1, 3)]
a
>>>
e:f32[3,3] = mul c d
>>>
f:f32[1,3] = broadcast_in_dim[broadcast_dimensions=(1,) shape=(1, 3)]
b
>>>
g:f32[3,3] = add f e
>>>
h:f32[] = reduce_sum[axes=(0, 1)] g
>>> in (h,) }
f2_jaxpr.consts
>>> [Array([2., 1., 0.], dtype=float32)]
❶
❷
❸
❹
❺
❻
❼
❻
❼
Побочный эффект.
Использование глобальной переменной z.
Генерация кода Jaxpr для исходной функции.
Результат побочного эффекта.
Проверка полученного представления Jaxpr.
Проверка списка констант.
Используемая глобальная переменная теперь стала константой.
Побочный эффект возникает во время трассировки, поэтому мож
но видеть его результат во время вызова функции, полученной с по
мощью трансформации make_jaxpr(). Полученное представление
jaxpr не содержит никаких элементов, каким-либо образом связан
ных с инструкцией print. Глобальная переменная теперь сохранена
в списке констант, связанных с полученным представлением jaxpr,
и если в дальнейшем изменить глобальную переменную, то скомпи
лированная версия будет продолжать использовать старое значение,
сохраненное как константа.
По умолчанию jax.jit для трассировки использует объект-трас
сировщик ShapedArray. Этот объект имеет конкретную форму, но не
содержит конкретное значение. Поэтому компилируемая функция
работает со всеми возможными входными данными, имеющими ту
же форму, что является обычным вариантом использования в ма
197
Внутренний механизм JIT
шинном обучении. Например, предположим, что мы создаем функ
цию для вычисления скалярного произведения двух векторов, и тре
буется выполнить JIT-компиляцию этой функции. При первом ее
вызове передаются два вектора формы [100,1], состоящие из чисел
с плавающей точкой. Формы и типы данных будут сохранены в кеше
во время первого вызова. Теперь скомпилированная версия исход
ной функции доступна для многократного использования с различ
ными экземплярами векторов того же размера и типа.
Основная идея заключается в том, что скомпилированный (или
в более обобщенном смысле – трансформированный) код должен
работать с различными входными значениями, поэтому JAX вы
полняет трассировку по абстрактным значениям, представляющим
наборы возможных входных данных. Существует несколько уров
ней абстракции (https://github.com/jax-ml/jax/blob/main/jax/_src/ab
stract_arrays.py), и различные трансформации используют разные
уровни. Более высокие уровни абстракции дают более обобщенное
представление кода Python и сокращают количество повторных ком
пиляций, но при этом накладывают больше ограничений на код Py
thon, чтобы обеспечить возможность его трассировки. Трассировка
на самом высоком уровне абстракции, использующая трассировщик
UnshapedArray, в настоящее время не устанавливается по умолчанию
для любой трансформации.
При использовании JIT это означает, что при трассировке не воз
никает проблем с инструкциями управления потоком выполнения,
если они не зависят от какого-либо значения входного параметра
(но для инструкций управления потоком допускается использова
ние форм входных параметров). Код в листинге 5.12 использует фик
сированное значение, независимое от какого-либо входного пара
метра, и форму входного параметра для формирования цикла.
Листинг 5.12 Трассировка при наличии управляющих структур
def f3(x):
y = x
for i in range(5):
y += i
return y
❶
jax.make_jaxpr(f3)(0)
>>> { lambda ; a:i32[].
>>>
b:i32[] = add a
>>>
c:i32[] = add b
>>>
d:i32[] = add c
>>>
e:i32[] = add d
>>>
f:i32[] = add e
>>>
in (f,) }
let
0
1
2
3
4
❷
❷
❷
❷
❷
Глава 5
198
Компиляция кода
jax.jit(f3)(0)
>>> Array(10, dtype=int32, weak_type=True)
def f4(x):
y = 0
for i in range(x.shape[0]):
y += x[i]
return y
❸
❹
jax.make_jaxpr(f4)(jnp.array([1.0, 2.0, 3.0]))
❺
>>> { lambda ; a:f32[3]. let
>>>
b:f32[1] = slice[limit_indices=(1,) start_indices=(0,) strides=(1,)]
a
>>>
c:f32[] = squeeze[dimensions=(0,)] b
>>>
d:f32[] = add 0.0 c
>>>
e:f32[1] = slice[limit_indices=(2,) start_indices=(1,) strides=(1,)]
a
>>>
f:f32[] = squeeze[dimensions=(0,)] e
>>>
g:f32[] = add d f
>>>
h:f32[1] = slice[limit_indices=(3,) start_indices=(2,) strides=(1,)]
a
>>>
i:f32[] = squeeze[dimensions=(0,)] h
>>>
j:f32[] = add g i
>>> in (j,) }
jax.jit(f4)(jnp.array([1.0, 2.0, 3.0]))
>>> Array(6., dtype=float32)
❶
❷
❸
❹
❺
❻
❻
Этот цикл не зависит от входного параметра.
Цикл развертывается.
JIT-компиляция завершилась успешно.
Этот цикл зависит от формы входного параметра.
Цикл развертывается.
JIT-компиляция завершилась успешно.
JAX успешно осуществил трассировку обоих вариантов, развернув
циклы и выполнив JIT-компиляцию исходных функций. Но если вы
попытаетесь использовать значение входного параметра в управ
ляющей структуре, то возникнет ошибка.
Листинг 5.13 Трассировка цикла for, зависящего от значения входного
параметра
def f5(x):
y = 0
for i in range(x):
y += i
return y
❶
199
Внутренний механизм JIT
f5(5)
❷
>>> 10
jax.make_jaxpr(f5)(5)
❸
>>> ...
>>> The above exception was the direct cause of the following exception:
# Сгенерированное выше исключение стало прямой причиной следующего исключения:
>>>
>>> TracerIntegerConversionError
Traceback (most recent call last)
#
Обратная трассировка (последним выводится самый недавний вызов)
>>> <ipython-input-54-626705d2393b> in f5(x)
>>>
1 def f5(x):
>>>
2
y = 0
>>> ----> 3
for i in range(x):
>>>
4
y += i
>>>
5
return y
>>>
>>> /usr/local/lib/python3.10/dist-packages/jax/_src/core.py in __index__
(self)
>>>
617
>>>
618
def __index__(self):
❹
>>> --> 619
raise TracerIntegerConversionError(self)
>>>
620
>>>
621
def tolist(self):
>>>
>>> TracerIntegerConversionError: The __index__() method was called on traced
array with shape int32[].
# Метод __index__() был вызван из трассируемого массива с формой int32[].
>>> The error occurred while tracing the function
f5 at <ipython-input-77-626705d2393b>:1 for make_jaxpr.
This concrete value was not available in Python because
it depends on the value of the argument x.
# Ошибка возникла при трассировке функции f5 в точке
# <ipython-input-77-626705d2393b>:1 для make_jaxpr.
# Это конкретное значение было недоступным в среде Python,
# потому что оно зависит от значения аргумента x.
>>> See https://jax.readthedocs.io/en/latest/errors.html#jax.errors.
TracerIntegerConversionError
❶
❷
❸
❹
Цикл зависит от (значения) входного параметра.
Это абсолютно обычная функция Python, которая работает.
У JAX возникают проблемы при трассировке этой функции.
Ошибка JIT-компиляции из-за проблем при трассировке.
Здесь мы использовали значение входного параметра как пара
метр цикла, и при трассировке возникли проблемы, потому что не
допустима какая-либо работа со значением параметра. То же самое
происходит при использовании оператора if (см. листинг 5.14).
200
Глава 5
Компиляция кода
Листинг 5.14 Трассировка кода, использующего оператор if, зависящий от
значения входного параметра
def relu(x):
if x > 0:
return x
return 0.0
relu(10.0)
>>> 10.0
jax.make_jaxpr(relu)(10.0)
❶
❷
❸
>>> The above exception was the direct cause of the following exception:
# Сгенерированное выше исключение стало прямой причиной следующего исключения:
>>>
>>> ConcretizationTypeError
Traceback (most recent call last)
#
Обратная трассировка (последним выводится самый недавний вызов)
>>> <ipython-input-58-5bd36ce502aa> in relu(x)
>>>
1 def relu(x):
>>> ----> 2
if x > 0:
>>>
3
return x
>>>
4
return 0.0
>>>
>>>
>>> /usr/local/lib/python3.10/dist-packages/jax/_src/core.py in error(self, arg)
>>>
1394
if fun is bool:
>>>
1395
def error(self, arg):
>>> -> 1396
raise TracerBoolConversionError(arg)
>>>
1397
else:
>>>
1398
def error(self, arg):
>>>
>>> TracerBoolConversionError: Attempted boolean conversion of traced array
with shape bool[]..
# Попытка преобразования в логический тип трассируемого массива
# с формой bool[]..
>>> The error occurred while tracing the function
relu at <ipython-input-83-5064d3aeee39>:1 for make_jaxpr.
This concrete value was not available in Python because
it depends on the value of the argument x.
# Ошибка возникла при трассировке функции relu в точке
# <ipython-input-83-5064d3aeee39>:1 для make_jaxpr.
# Это конкретное значение было недоступным в среде Python,
# потому что оно зависит от значения аргумента x.
>>> See https://jax.readthedocs.io/en/latest/errors.html#jax.errors.
TracerBoolConversionError
❶ Оператор if зависит от входного параметра.
❷ Это абсолютно обычная функция Python, которая работает.
❸ Ошибка при трассировке.
201
Внутренний механизм JIT
Здесь мы использовали значение входного параметра в операто
ре if, и при трассировке возникла ошибка. Но трассировка выпол
нялась для обычной функции Python, которая нормально работает.
Процессом трассировки можно управлять, используя механизм
статических параметров, с которым вы уже знакомы (см. подраз
дел 5.1.1). Объявление некоторого параметра как статического (обе
функции jax.jit() и jax.make_jaxpr() поддерживают такой подход)
позволяет при трассировке использовать конкретные значения,
и компиляция работает, как полагается.
Листинг 5.15 Использование статических параметров
для трассировки с конкретными значениями
def f5(x):
y = 0
for i in range(x):
y += i
return y
❶
def relu(x):
if x > 0:
return x
return 0.0
❶
jax.make_jaxpr(f5, static_argnums=0)(5)
>>> { lambda ; . let
in (10,) }
jax.jit(f5, static_argnums=0)(5)
>>> Array(10, dtype=int32, weak_type=True)
jax.make_jaxpr(relu, static_argnums=0)(12.3)
>>> { lambda ; . let
in (12.3,) }
jax.jit(relu, static_argnums=0)(12.3)
>>> Array(12.3, dtype=float32, weak_type=True)
❷
❸
❷
❹
❷
❸
❷
❹
❶ Зависимость от входного параметра.
❷ Пометка первого параметра как статического.
❸ Полученное в итоге выражение фактически является предварительно вычисляе-
мой константой.
❹ Теперь компиляция завершается успешно.
Это компромисс. Теперь исходная функция компилируется при
каждом вызове с новым значением входного параметра. Такой под
ход, возможно, будет приемлемым для функции с небольшим набо
ром допустимых входных значений (как для показанной в приме
ре функции f5()). Но подобное решение не подходит для функций
Глава 5
202
Компиляция кода
с многочисленными допустимыми входными значениями (явным
примером служит функция активации relu(), которая может при
нимать практически любое значение во время тренировки нейрон
ной сети).
Другой, более эффективный способ устранения этой проблемы –
использование структурированных базисных элементов управления
потоком выполнения, которые упоминались в разделе 3.4. Для нача
ла заменим цикл for Python из листинга 5.13 на базисный элемент
jax.lax.fori_loop() (https://jax.readthedocs.io/en/latest/_autosummary/
jax.lax.fori_loop.html). Цикл fori_loop(lower, upper, body_fun, init_
val) равнозначен следующему коду Python:
def fori_loop(lower, upper, body_fun, init_val):
val = init_val
for i in range(lower, upper):
val = body_fun(i, val)
return val
Измененный код показан в листинге 5.16.
Листинг 5.16 Замена цикла for на структурированный базисный
элемент управления потоком выполнения
def f5(x):
return jax.lax.fori_loop(0, x, lambda i,v: v+i, 0)
❶
f5(5)
>>> Array(10, dtype=int32, weak_type=True)
❷
jax.make_jaxpr(f5)(5)
>>> { lambda ; a:i32[]. let
>>>
_:i32[] _:i32[] b:i32[] = while[
>>>
body_jaxpr={ lambda ; c:i32[] d:i32[] e:i32[]. let
>>>
f:i32[] = add c 1
>>>
g:i32[] = add e c
>>>
in (f, d, g) }
>>>
body_nconsts=0
>>>
cond_jaxpr={ lambda ; h:i32[] i:i32[] j:i32[]. let
>>>
k:bool[] = lt h i
>>>
in (k,) }
>>>
cond_nconsts=0
>>>
] 0 a 0
>>>
in (b,) }
jax.jit(f5)(5)
>>> Array(10, dtype=int32, weak_type=True)
❸
❹
❶ Использование jax.lax.fori_loop для замены цикла for из листинга 5.13.
❷ Результат остается тем же самым.
Внутренний механизм JIT
203
❸ Код Jaxpr становится более сложным.
❹ Теперь функция компилируется.
Здесь мы заменили цикл for, зависимый от входного параметра,
из листинга 5.13 на базисный элемент jax.lax.fori_loop(), и трас
сировка прошла успешно.
Теперь заменим оператор if из листинга 5.14 на базисный эле
мент jax.lax.cond() (https://docs.jax.dev/en/latest/_autosummary/jax.
lax.cond.html), равнозначный следующей реализации на языке Py
thon:
def cond(pred, true_fun, false_fun, *operands):
if pred:
return true_fun(*operands)
else:
return false_fun(*operands)
Оба элемента true_fun() и false_fun() обязательно должны быть
вызываемыми объектами и возвращать в точности те же самые типы.
Применение базисного элемента cond() в исходном коде не вы
зывает никаких затруднений (см. листинг 5.17).
Листинг 5.17 Замена оператора if на структурированный базисный
элемент управления потоком выполнения
def relu(x):
return jax.lax.cond(x>0, lambda x: x, lambda x: 0.0, x)
relu(12.3)
>>> Array(12.3, dtype=float32, weak_type=True)
jax.make_jaxpr(relu)(12.3)
❶
❷
❸
>>> { lambda ; a:f32[]. let
>>>
b:bool[] = gt a 0.0
>>>
c:i32[] = convert_element_type[new_dtype=int32 weak_type=False] b
>>>
d:f32[] = cond[
>>>
branches=(
>>>
{ lambda ; e:f32[]. let in (0.0,) }
>>>
{ lambda ; f:f32[]. let in (f,) }
>>>
)
>>>
linear=(False,)
>>>
] c a
>>> in (d,) }
jax.jit(relu)(12.3)
>>> Array(12.3, dtype=float32, weak_type=True)
❶ Использование jax.lax.cond для замены оператора if из листинга 5.14.
❷ Результат остается тем же самым.
❹
Глава 5
204
Компиляция кода
❸ Код Jaxpr становится более сложным.
❹ Теперь функция компилируется.
Мы заменили оператор языка Python if на jax.lax.cond(), и те
перь трассировка проходит успешно.
В действительности не каждая трансформация JAX приводит
к материализации кода Jaxpr, как показано выше. Некоторые транс
формации, например взятие градиентов или пакетирование, при
меняют трансформационные операции постепенно во время трас
сировки (не принимая Jaxpr сначала, а затем обрабатывая его для
получения измененного кода Jaxpr).
Мы завершили рассмотрение первого этапа компиляции – преоб
разования кода Python в промежуточное представление Jaxpr. После
этого начинается второй этап: компиляция кода Jaxpr в машинный
код целевой платформы с помощью XLA.
5.2.2
XLA
XLA – это сокращение от Accelerated Linear Algebra (https://www.
tensorfow.org/xla), названия предметно-ориентированного компи
лятора для задач линейной алгебры, изначально разработанного для
ускорения моделей TensorFlow, предположительно без изменений
в исходном коде (см. рис. 5.3).
Код Python
Jaxpr
Компиляция
с помощью XLA
Рис. 5.3 Преобразование представления Jaxpr в машинный код целевой
платформы
Происхождение и архитектура XLA
Причиной появления компилятора XLA является то обстоятельство,
что хотя каждая отдельная операция в графе вычислений может быть
высокооптимизированной, у пользователя имеется возможность
формирования более сложных операций из простых или создания
крупной композиции, не гарантирующей эффективное выполнение.
Поэтому компания Google разработала XLA, использующий мето
дики JIT-компиляции для анализа графа TensorFlow со специали
зацией его для реальных измерений и типов времени выполнения,
а также, что более важно, для объединения нескольких операций
в одну. Компилятор может использовать информацию о конкретной
модели для оптимизации и компилировать граф TensorFlow в после
довательность ядер вычислений, сгенерированных специально для
конкретной модели.
Внутренний механизм JIT
205
Например, рассмотрим операцию из листинга 5.10, принимаю
щую три тензора как входные данные и вычисляющую выводимый
результат:
def f(x, y, z):
return jnp.sum(x + y * z)
❶
❶ Выполняются три операции: умножение, сложение и суммирование.
Без использования XLA такое вычисление, вероятнее всего, пред
полагало бы выполнение трех различных операций, реализованных
с помощью различных вычислительных ядер, например, на GPU:
одно ядро для умножения, одно – для сложения и одно – для итого
вого суммирования. XLA может оптимизировать это вычисление для
получения результата при запуске единственного ядра. Компилятор
способен объединять операции умножения, сложения и суммиро
вания в одном ядре GPU. Кроме того, такая объединенная опера
ция не создает промежуточные переменные, хранящие результаты
y × z и x + y × z в памяти. Есть возможность передавать напрямую
эти результаты в последующие вычисления, сохраняя данные в той
же локации памяти или в регистрах GPU. Исключение избыточных
операций передачи данных между локациями памяти – это большое
преимущество, так как пропускная способность памяти может стать
узким местом для производимых вычислений.
XLA генерирует эффективный машинный код целевой платфор
мы для таких устройств, как CPU, GPU, и специализированных аксе
лераторов, например TPU компании Google. Подсистема XLA, кото
рая выполняет целенаправленную, ориентированную на конкретное
устройство оптимизацию и генерацию машинного кода, называется
внутренним компонентом (backend).
Система, передающая данные в XLA, называется внешним ком
понентом (frontend). Изначально внешним компонентом для XLA
был фреймворк TensorFlow, но сейчас программы XLA могут также
генерироваться фреймворками PyTorch, JAX, Julia и Nx (библиотека
вычислительных методов для языка программирования Elixir).
OpenXLA
OpenXLA (https://github.com/openxla) – экосистема с открытым исходным кодом компилятора ML (машинного обучения), совместно разработанного лидерами индустрии искусственного интеллекта / машинного обучения: Alibaba, Amazon Web Services, AMD, Anyscale, Apple,
Arm, Cerebras, Google, Graphcore, Hugging Face, Intel, Meta, и NVIDIA. Эта
разработка позволила предоставить место постоянной дислокации для
управляемой сообществом экосистемы ML-компилятора с открытым исходным кодом.
206
Глава 5
Компиляция кода
Компилятор XLA был отделен от TensorFlow и передан в проект OpenXLA, и сообщество начало развивать его совместными усилиями. Проект
также включает StableHLO (более подробно об этом будет сказано немного позже) и репозитории IREE. Все это направлено на расширение
MLIR: инфраструктуру компилятора, обеспечивающую для моделей машинного обучения согласованное представление, оптимизацию и выполнение на аппаратных устройствах (подробнее об этом также будет
сказано немного позже).
OpenXLA предоставляет модульный комплект инструментальных
средств, поддерживаемый всеми ведущими фреймворками через общий интерфейс компилятора, и расширяет стандартизированные представления моделей, являющиеся переносимыми, а также предоставляет
предметно-ориентированный компилятор с мощными средствами оптимизации, специализированными для аппаратуры и независимыми от
целевой платформы.
Более подробную информацию об OpenXLA можно получить в Google
Open Source Blog: https://opensource.googleblog.com/2023/03/openx
la-is-ready-to-accelerate-and-simplify-ml-development.html.
XLA использует специализированный входной язык промежуточ
ного представления операций высокого уровня HLO IR, или просто
HLO (high-level operations intermediate representation). XLA прини
мает вычисления, определенные на языке HLO, и компилирует их
в специализированные машинные инструкции для целевой аппа
ратной архитектуры.
В языке HLO существует множество простых базовых операций.
Полный их список можно найти здесь: https://www.tensorfow.org/
xla/operation_semantics.
На рис. 5.4 показаны два этапа оптимизации: независимой от це
левой платформы и ориентированной на целевую платформу. На
этапе независимой от целевой платформы оптимизации XLA соз
дает оптимизации, не зависящие от аппаратуры, на которой будут
производиться вычисления. В этот этап включено устранение под
выражений, объединение операций, независимых от целевой плат
формы, и анализ буферов для распределения памяти для вычисле
ний во время выполнения.
На этапе ориентации на целевую платформу внутренний компо
нент XLA может выполнять дальнейшие оптимизации уровня HLO.
Например, он определяет, как лучше распределить вычисления по
потокам GPU, и осуществляет другие объединения операций, спе
циализированные для конкретного GPU. Также возможно выполне
ние специализированного алгоритма сопоставления с паттерном
для замены некоторой комбинации операций на оптимизирован
ные библиотечные вызовы.
Внутренний механизм JIT
XLA HLO
207
Рис. 5.4 Процесс
компиляции в XLA
Независимые
от целевой платформы
оптимизации и анализ
XLA HLO
Зависимые
от целевой платформы
оптимизации и анализ
Генерация кода,
специализированного
для целевой платформы
Внутренний компонент XLA
После оптимизации, ориентированной на целевую платфор
му, и анализа внутренний компонент XLA генерирует код, специа
лизированный для целевой платформы. Внутренние компоненты
CPU и GPU используют LLVM (https://llvm.org/) для промежуточного
представления низкого уровня, оптимизации и генерации кода. Эти
внутренние компоненты выводят LLVM IR (IR – intermediate repre
sentation – промежуточное представление), представляющее XLA
HLO IR, а затем вызывают LLVM для генерации машинного кода це
левой платформы из LLVM IR.
Внутренний компонент CPU поддерживает архитектуры x64
и ARM64, внутренний компонент GPU поддерживает NVIDIA GPU
и в определенной степени графические процессоры компаний AMD
(https://docs.jax.dev/en/latest/developer.html#additional-notes-forbuilding-a-rocm-jaxlib-for-amd-gpus), Apple и Intel (https://github.
com/jax-ml/jax/issues/2012#issuecomment-1603312724). Разумеется,
существует и внутренний компонент TPU, поддерживающий Google
Cloud TPU. Кроме того, предоставляется способ разработки пользо
вателем собственного внутреннего компонента XLA (https://openxla.
org/xla/developing_new_backend), что создает реальную возможность
для поддержки всего новейшего оборудования для глубокого обуче
ния в ближайшем будущем.
XLA и JAX
JAX использует XLA для генерации эффективного кода для конкрет
ных внутренних компонентов. Но это еще не все. С января 2022 г. JAX
начал применять диалект MHLO MLIR как свой основной целевой
компилятор промежуточного представления (IR) по умолчанию, т. е.
внутренний компонент переключился с XLA/HLO на MLIR/MHLO.
Таким образом, рабочий процесс стал выглядеть так (https://github.
com/google/jax/issues/10715):
Глава 5
208
Компиляция кода
пользователь пишет функцию Python;
JAX выполняет преобразование функции Python в Jaxpr;
3 JAX выполняет преобразование Jaxpr в MHLO для MLIR;
4 MLIR выполняет преобразование MHLO в оптимизированный
MHLO;
5 JAX выполняет преобразование оптимизированного MHLO
в HLO;
6 XLA выполняет преобразование HLO в оптимизированный HLO;
7 XLA выполняет преобразование оптимизированного HLO в ма
шинный код целевой платформы для CPU/GPU/TPU (для CPU
и GPU используется LLVM).
1
2
Затем в конце 2022 г. JAX перешел на StableHLO вместо MHLO
(https://github.com/jax-ml/jax/commit/a1480c454e69a2631dbd51cb2f4
fccc2752c18a0#diff-002064eb864e3ba1dc60e6e9ac5e06c7ca0890f32297
84281f07e923218a18d3). При необходимости можно продолжать ге
нерировать диалект MHLO; пока обеспечивается поддержка функ
ций MHLO и StableHLO, хотя совместимость гарантируется только
для StableHLO, но не для MHLO.
StableHLO (https://github.com/openxla/stablehlo) был создан для
обеспечения надежного уровня переносимости между различными
фреймворками машинного обучения и ML-компиляторами: фрейм
ворки машинного обучения, генерирующие программы StableHLO,
совместимы с ML-компиляторами, потребляющими программы
StableHLO. StableHLO основан на диалекте MHLO и расширяет его,
добавляя функциональность, включая сериализацию и управление
версиями. StableHLO является частью проекта OpenXLA.
MLIR
MLIR (https://mlir.llvm.org/), или Multi-Level Intermediate Representation (многоуровневое промежуточное представление), является преемником LLVM, а идею создания MLIR подал Крис Латтнер (Chris Lattner)
(https://www.youtube.com/watch?v=qzljG6DKgic).
Представление MLIR ориентировано на создание многократно используемой и расширяемой инфраструктуры компилятора, удовлетворяющей требованиям сферы разнообразного аппаратного обеспечения.
В частности, MLIR имеет великолепные перспективы в области создания инфраструктуры оптимизирующего компилятора для приложений
глубокого обучения.
MLIR позиционировано как гибридное промежуточное представление
(IR), способное поддерживать многочисленные разнообразные требования в унифицированной инфраструктуре, включая следующие (но
этот список не является полным):
Внутренний механизм JIT
209
представление графов потоков данных, например таких, как в TensorFlow;
оптимизации и трансформации обычно выполняются на подобных
графах;
способность управлять оптимизациями вычислительных циклов с высокой производительностью для всех ядер (слиянием, перестановкой
порядка вложенных циклов, тайлинга и т. п.) и выполнять трансформации схем размещений данных в памяти;
генерация кода «понижающих» трансформаций, таких как DMAвставка, явное управление кешем, тайлинг памяти и векторизация для
архитектур регистров 1D и 2D;
возможность представления операций, специализированных для
конкретной платформы, например операций высокого уровня, спе
циализированных для конкретного акселератора;
квантование (дискретизация) и другие трансформации графа глубокого обучения;
многогранные (полиэдрические) базисные элементы (primitives);
инструменты синтеза (построения) аппаратного оборудования / синтез высокого уровня.
MLIR – это мощное представление, но при этом не имеющее конкретно
определенных целей. MLIR не пытается поддерживать алгоритмы генерации машинного кода низкого уровня (например, распределение регистров и планирование инструкций), так как с этим лучше справляются
оптимизаторы низкого уровня (такие как LLVM). Кроме того, MLIR не позиционируется как язык исходного кода, на котором конечные пользователи могли бы самостоятельно писать ядра (по аналогии с CUDA C++).
MLIR поддерживает различные диалекты (https://www.tensorfow.org/
mlir/dialects) для своего промежуточного представления. JAX использует диалект MLIR MHLO («мета»-HLO) (https://github.com/tensorflow/
mlir-hlo#meta-hlo-dialect-mhlo). Со списком операций этого диалекта
можно ознакомиться здесь: https://www.tensorfow.org/mlir/hlo_ops.
Предпринимались попытки создания отдельного независимого компилятора «HLO» на основе MLIR под названием MLIR-HLO (https://github.
com/tensorfow/mlir-hlo). В настоящее время группа разработки останавливает поддержку репозитория MLIR-HLO (https://discourse.llvm.
org/t/sunsetting-the-mlir-hlo-repository/70536) и переключается на
проект StableHLO (https://github.com/openxla/stablehlo).
В коде листинга 5.18 можно видеть, как функция Python преоб
разуется сначала в StableHLO (или в MHLO, если вы используете
более старую версию JAX), а затем в HLO низкого уровня. Мы рас
сматриваем ту же функцию, которая использовалась в разделах,
описывающих Jaxpr и XLA. Функция содержит операции умноже
210
Глава 5
Компиляция кода
ния, сложения и суммирования. Для вычисления предоставляются
те же тестовые данные.
Далее начинаются чрезвычайно интересные вещи. Мы выполня
ем JIT-компиляцию исходной функции. Для JIT-скомпилированной
функции можно создать версию более низкого уровня. Снижение
уровня (lowering) – это процесс преобразования представления вы
сокого уровня в представление низкого уровня. Здесь мы создаем
код StableHLO IR, состоящий из простейших базисных операций
StableHLO. Этот код относительно прост и очень похож на исходные
вычисления. Мы двигаемся дальше и компилируем полученный код
StableHLO в представление HLO, оптимизированное для целевого
внутреннего компонента. Такой код HLO содержит объединенные
вычисления. Полный код доступен в репозитории книги, а здесь
приводятся только части кода StableHLO и HLO.
Листинг 5.18 Компиляция кода Python в представление StableHLO
и HLO
def f(x, y, z):
return jnp.sum(x + y * z)
❶
x = jnp.array([1.0, 1.0, 1.0])
y = jnp.ones((3,3))*2.0
z = jnp.array([2.0, 1.0, 0.0]).T
❷
❷
❷
f_jitted = jax.jit(f)
f_lowered = f_jitted.lower(x,y,z)
print(f_lowered.as_text())
>>> ...
>>>
%5 = stablehlo.add %4, %2 : tensor<3x3xf32>
>>>
%6 = stablehlo.constant dense<0.000000e+00> : tensor<f32>
>>> ...
f_compiled = f_jitted.lower(x,y,z).compile()
print(f_compiled.as_text())
>>> ...
>>>
%multiply.1 = f32[3,3]{1,0} multiply(f32[3,3]{1,0}
➥%param_0.4, f32[3,3]{1,0} %broadcast.4),
➥metadata={op_name="jit(f)/jit(main)/mul"
➥source_file="<ipython-input-99-95d48614b981>"
➥source_line=2}
>>>
%add.2 = f32[3,3]{1,0} add(f32[3,3]{1,0}
➥%broadcast.5, f32[3,3]{1,0} %multiply.1),
➥metadata={op_name="jit(f)/jit(main)/add"
➥source_file="<ipython-input-99-95d48614b981>" source_line=2}
>>> ...
❶ Простая функция с несколькими операциями.
❸
❹
❺
❺
❺
❺
❻
❼
❼
❼
❼
Внутренний механизм JIT
211
Некоторые тестовые данные.
JIT-компиляция исходной функции.
Снижение уровня кода исходной функции (генерация кода StableHLO).
Представление StableHLO.
Компиляция кода пониженного уровня функции для конкретного внутреннего
компонента (генерация кода HLO).
❼ Представление HLO.
❷
❸
❹
❺
❻
Анализ кода StableHLO и HLO, а также более подробное рассмот
рение внутренних механизмов XLA и MLIR не связаны с тематикой
этой книги, тем не менее теперь вы имеете представление об этапах
процесса компиляции и о том, что на них происходит.
В рассматриваемом здесь примере для получения представлений
StableHLO и HLO мы использовали специальные функции для сни
жения уровня и компиляции кода. В действительности эти функции
формируют новый API AOT-компиляции.
5.2.3
Использование AOT-компиляции
С сентября 2022 г. и выпуска версии JAX 0.3.18 была введена AOTкомпиляция (https://github.com/jax-ml/jax/blob/main/docs/aot.md).
AOT-компиляция (ahead-of-time compilation) может оказаться по
лезной, если необходимо взять на себя управление во время вы
полнения различных частей процесса компиляции или если требу
ется полная компиляция до времени выполнения. Например, такой
подход позволяет сократить время запуска функции, если заранее
выполнить этап компиляции и исключить первый проход по коду
функции, необходимый при JIT-компиляции.
Предположим, что имеется некоторая Python-функция f(x), где
x – массив. Вы применяете трансформацию jit() к этой функции
и получаете версию в обертке f_jit = jit(f). Затем в некоторый мо
мент вы вызываете jit-трансформированную версию функции с ар
гументами f_jit(x). Именно здесь выполняется компиляция. Схема
этого процесса показана на рис. 5.5.
Процесс компиляции содержит несколько этапов:
этап преобразования исходной функции f во внутреннее пред
ставление. Эта специализированная версия функции отобра
жает ограничения для входных типов, выведенные из свойств
аргументов (здесь: только x);
снижение уровня специализированного преобразованного пред
ставления до входного языка MLIR и XLA (здесь: StableHLO);
компиляция HLO-программы со сниженным уровнем в специа
лизированный для устройства оптимизированный код для CPU,
GPU или TPU;
выполнение скомпилированного кода с заданными аргумента
ми (здесь: x).
Глава 5
212
Входные
данные
Компиляция кода
jit()-трансформированная
функция
Рис. 5.5 Этапы компиляции
при вызове JIT-трансформированной
версии функции
Этап
преобразования
Внутреннее
представление
Понижение уровня
StableHLO
Компиляция
Специализированный
для устройства
оптимизированный код
Выполнение
Вывод результатов
AOT API обеспечивает управление последними тремя этапами.
Существует специализированный API jax.stages (https://docs.jax.
dev/en/latest/jax.stages.html) для представления этапов процесса вы
полнения скомпилированного кода.
В листинге 5.19 используется та же функция активации selu(), что
и в начале текущей главы. Мы по-прежнему преднамеренно добав
ляем побочный эффект, чтобы понять, где происходит трассиров
ка и компиляция. Создается JIT-скомпилированная версия и AOTскомпилированная версия.
Листинг 5.19
JIT-компиляция и AOT-компиляция
def selu(x,
❶
alpha=1.6732632423543772848170429916717,
scale=1.0507009873554804934193349852946):
'''Scaled exponential linear unit activation function.'''
# Функция активизации линейным блоком масштабируемой
# экспоненциальной кривой.
print('Function run')
❷
return scale * jnp.where(x > 0, x, alpha * jnp.exp(x) - alpha)
213
Внутренний механизм JIT
selu_jit = jax.jit(selu)
selu_aot = jax.jit(selu).lower(1.0).compile()
>>> Function run
selu_jit(17.8)
>>> Function run
>>> Array(18.702477, dtype=float32, weak_type=True)
selu_aot(17.8)
>>> Array(18.702477, dtype=float32, weak_type=True)
❶
❷
❸
❹
❺
❻
❼
❽
❸
❹
❺
❻
❼
❽
Некоторая функция для тестирования компиляции.
Побочный эффект для отладочных целей.
JIT-скомпилированная функция.
AOT-скомпилированная функция (обратите внимания на фиктивный аргумент 1.0,
необходимый при AOT-компиляции для вывода типов).
Побочный эффект работает во время AOT-компиляции.
Вызов JIT-скомпилированной функции.
Побочный эффект работает во время JIT-компиляции (он будет проявляться, только если ранее не была выполнена AOT-компиляция).
Вызов AOT-скомпилированной функции без побочных эффектов.
Обе функции производят одинаковые результаты. Но AOT-ском
пилированная версия не требует разогрева, и в момент ее вызова
мы уже имеем скомпилированный код. Мы видим это в тот момент,
когда работает побочный эффект. AOT-скомпилированная функция
работает во время прямой компиляции; JIT-компиляция функции
происходит во время ее первого вызова, когда выполняется дей
ствительный процесс компиляции.
Снижение уровня и компиляция выполняются для фиксирован
ной сигнатуры типа, и AOT-скомпилированная функция может вы
зываться только с аргументами этой фиксированной сигнатуры
типа. Если вызвать AOT-скомпилированную функцию с аргумента
ми, несовместимыми с представлением со сниженным уровнем (на
пример, float32 вместо int32), то возникнет ошибка.
AOT-скомпилированные функции нельзя преобразовать с по
мощью таких трансформаций, как jit, grad, jvp, vmap. Это запре
щение установлено потому, что внутри многих трансформаций из
меняется сигнатура типа функций. Хотя jit не изменяет сигнатуру
типов своих аргументов, эта трансформация также запрещена.
Код в листинге 5.20 применяет обе скомпилированные функ
ции из предыдущего примера (листинг 5.19) к различным типам
данных – теперь для 32-битовых целых чисел вместо 32-битовых
чисел с плавающей точкой. Для JIT-скомпилированной функции
проблема не возникает – функция просто перекомпилируется. AOT-
Глава 5
214
Компиляция кода
скомпилированную функцию невозможно перекомпилировать, по
этому возвращается ошибка.
Листинг 5.20
Различия в поведении между JIT и AOT версиями
selu_jit(17)
>>> Function run
>>> Array(17.861917, dtype=float32, weak_type=True)
selu_aot(17)
❶
❷
>>> ...
>>> TypeError: Argument types differ from the types
for which this computation was compiled.
The mismatches are:
>>> Argument 'x' compiled with float32[] and called with int32[]
# TypeError: типы аргументов отличаются от типов,
# для которых было скомпилировано это вычисление.
# Различия:
# Аргумент 'x' скомпилирован для типа float32[],
# а вызывается с типом int32[]
selu_jit_batched = jax.vmap(selu_jit)
selu_aot_batched = jax.vmap(selu_aot)
selu_jit_batched(jnp.array([42.0, 78.0, -12.3]))
>>> Function run
>>> Array([44.129444 , 81.95468
❸
❸
❹
, -1.7580913], dtype=float32)
selu_aot_batched(jnp.array([42.0, 78.0, -12.3]))
>>> ...
>>> TypeError: Cannot apply JAX transformations to
a function lowered and compiled for a particular
signature. Detected argument of Tracer type <class
'jax._src.interpreters.batching.BatchTracer'>.
# TypeError: невозможно применить трансформации JAX
# к функции со сниженным уровнем и скомпилированной
# для конкретной сигнатуры. Обнаружен аргумент для
# типа трассировщика (Tracer)<class
# 'jax._src.interpreters.batching.BatchTracer'>.
❺
❶ Вызов JIT-скомпилированной функции с целочисленными данными. Выполняется
перекомпиляция.
❷ Вызов AOT-скомпилированной функции с целочисленными данными. Невозмож-
но выполнить перекомпиляцию, поэтому возвращается ошибка.
❸ Создание пакетных версий исходной функции (с помощью vmap; см. следующую
главу).
Внутренний механизм JIT
215
❹ JIT-скомпилированную функцию можно трансформировать, и снова происходит
перекомпиляция.
❺ AOT-скомпилированную функцию трансформировать невозможно, поэтому воз-
вращается ошибка.
Этапы AOT-компиляции предоставляют некоторые дополнитель
ные функциональные возможности для отладочных целей. Мы ис
пользовали эти возможности в листинге 5.18 для получения пред
ставлений StableHLO и HLO. Для скомпилированных функций также
можно получить анализ стоимости (накладных расходов) и исполь
зования памяти с помощью функций cost_analysis() и memory_ana
lysis(). Обе функции предназначены для визуального представле
ния и отладки. Они предоставляют некоторые простые структуры
данных, которые можно с легкостью вывести (на экран или на пе
чать) или сериализовать, но возможна их несовместимость между
различными версиями JAX и jaxlib или даже между вызовами.
JAX и jaxlib
JAX выпускается в виде двух отдельных пакетов Python:
jax – чистый пакет Python;
jaxlib – пакет, в основном написанный на языке C++, содержащий такие библиотеки, как XLA, части LLVM, используемые XLA, инфраструктура MLIR с привязками MHLO Python, а также специализированные
для JAX библиотеки C++ для быстрой обработки JIT и pytree.
Дистрибутив JAX имеет такую структуру, потому что большинство изменений в JAX относятся только к коду Python. Такой подход упрощает
работу с частью JAX, написанной на Python, без необходимости сборки
кода C++. Поэтому компоненты на Python можно обновлять независимо
от компонентов на C++, что повышает скорость разработки.
Для jax и jaxlib совместно используется одинаковый номер версии,
но выпускаются они раздельно. При установке номер версии пакета
jax обязательно должен быть больше или равен номеру версии jaxlib,
а номер версии jaxlib должен быть больше или равен минимальному
номеру версии jaxlib, определенному пакетом jax.
Об управлении версиями JAX и jaxlib более подробно можно узнать
здесь: https://docs.jax.dev/en/latest/jep/9419-jax-versioning.html#why
-are-jax-and-jaxlib-separate-packages.
Мы подробно рассмотрели основы практического использования
JIT и AOT компиляции, а также внутреннюю работу JAX. Это поможет
нам справляться с ограничениями JIT, поэтому важно знать, какие
ограничения существуют для JIT-компиляции.
216
5.3
Глава 5
Компиляция кода
Ограничения JIT-компиляции
JIT – это мощный механизм. Но и он имеет некоторые ограничения
и работает не всегда и не везде. Мы уже встречались с некоторыми
из таких ограничений, а здесь мы рассмотрим их в одном разделе.
Теперь, когда вы хорошо знакомы с работой внутреннего механизма
JIT, будет гораздо проще понять сущность этих ограничений, а также
инструментальных средств, предназначенных для их преодоления.
5.3.1
Чистые функции и функции, не являющиеся чистыми
Во-первых, как уже отмечалось в подразделе 5.1.2, JIT корректно ра
ботает только с чистыми функциями. Поэтому JIT может изменить
поведение конкретной функции, если вы полагаетесь на отсутствие
чистоты в ней, используете побочные эффекты или глобальное со
стояние. Например, скомпилированная функция, использующая
некоторое глобальное значение, не будет отображать дальнейшие
изменения этого значения, поскольку скомпилированная версия ис
пользует значение, существующее на момент компиляции (и трас
сировки).
5.3.2
Точные числовые данные
Во-вторых, JIT может изменять точные числовые данные при выво
де результатов функции. Это может происходить из-за оптимизации
в процессе компиляции. Например, XLA может переупорядочить
операции с плавающей точкой или исключить некоторые избыточ
ные вычисления, которые не должны производить какой-либо эф
фект с математической точки зрения (например, деление значения
на некоторое число, а затем умножение результата на то же самое
число), но, возможно, создающие некий эффект из-за накопления
арифметических погрешностей, связанных с проблемами арифме
тики с плавающей точкой.
5.3.3
Условные выражения, использующие значения входных
параметров
В-третьих, существуют ограничения, связанные с потоком управле
ния и условными выражениями, использующими значения входных
параметров. Например, цикл, зависящий от значения входного па
раметра, невозможно трассировать, но его можно заменить решени
ем со структурированными базисными элементами управления по
током выполнения из API jax.lax. Такие решения рассматривались
в подразделе 5.2.1.
217
Ограничения JIT-компиляции
5.3.4
Медленная компиляция
В-четвертых, jit-трансформированные функции иногда очень мед
ленно компилируются – даже в течение 10 и более секунд. Это мо
жет произойти, если исходный код генерирует чрезвычайно боль
шое внутреннее представление Jaxpr. Длинное представление Jaxpr
получается из больших циклов, которые развертываются при трас
сировке и интенсивно используют поток управления. Если Jaxpr со
держит сотни или тысячи строк, то вполне ожидаемой становится
медленная компиляция, поскольку время компиляции XLA возрас
тает приблизительно как квадрат количества операций, переданных
в компилятор.
Вы должны исключить такие циклы из программы. Существует
много способов сделать это, в том числе векторизация вычисле
ний (это обычная рекомендация для любых NumPy-подобных вы
числений) и замена циклов Python на структурированные базисные
элементы управления потоком выполнения из jax.lax. Одним из
вариантов также может стать изменение алгоритма (но не путайте
время компиляции со временем выполнения – второе обычно име
ет гораздо более важное значение). Кроме того, можно исключить
обертывание таких циклов с помощью jit. Но при этом остается
возможность применения трансформации jit для функций, распо
ложенных внутри цикла.
Рассмотрим учебный пример вычисления суммы с накоплением
(или нарастающий итог). Имеется массив элементов, и поставлена
задача – сформировать другой массив, где каждый элемент является
суммой всех элементов исходного массива, расположенных до соот
ветствующего элемента, включая его. В листинге 5.21 показана прос
тейшая реализация на языке Python.
Листинг 5.21
Простейшая реализация суммы с накоплением
def cumulative_sum(x):
acc = 0.0
y = []
for i in range(x.shape[0]):
acc += x[i]
y.append(acc)
return y
❶
j = jax.make_jaxpr(cumulative_sum)(jnp.ones(10000))
len(j.jaxpr.eqns)
>>> 30000
%time jax.jit(cumulative_sum)(jnp.ones(10000))
❷
Глава 5
218
Компиляция кода
>>> CPU times: user 2min 16s, sys: 1.47 s, total: 2min 18s
>>> Wall time: 2min 23s
❶ Функция для вычисления суммы с накоплением.
❷ Полученное представление Jaxpr содержит несколько тысяч строк.
❸ Для компиляции требуется достаточно длительное время.
❸
❸
В рассматриваемом здесь примере простой одноуровневый цикл
при развертывании преобразовался в огромное количество проме
жуточных выражений (равенств) Jaxpr. В этом представлении Jaxpr
30 000 выражений, поэтому для компиляции требуется более двух
минут.
Такой длинный цикл for можно заменить решением с использова
нием базисного элемента lax.scan (https://docs.jax.dev/en/latest/_au
tosummary/jax.lax.scan.html). Функция (базисный элемент) jax.lax.
scan() проходит по массиву (сканирует его) и вычисляет заданную
функцию для каждого элемента массива с накоплением состояния
(называемого carry). Мы используем carry для хранения накапли
ваемых сумм. Функция возвращает новое значение накапливаемо
го состояния и соответствующий элемент выходного массива. Это
равнозначно следующему коду на языке Python:
def scan(f, init, xs, length=None):
if xs is None:
xs = [None] * length
carry = init
ys = []
for x in xs:
carry, y = f(carry, x)
ys.append(y)
return carry, np.stack(ys)
Обновленный код показан в листинге 5.22.
Листинг 5.22 Реализация суммы с накоплением с использованием
lax.scan
def cumulative_sum_fast(x):
result, array = jax.lax.scan(
lambda carry, elem: (carry+elem, carry+elem), 1.0, x)
return array
j = jax.make_jaxpr(cumulative_sum_fast)(jnp.ones(10000))
len(j.jaxpr.eqns)
>>> 1
❶
%time cs = jax.jit(cumulative_sum_fast)(jnp.ones(10000))
>>> CPU times: user 145 ms, sys: 6.02 ms, total: 151 ms
❷
219
Ограничения JIT-компиляции
>>> Wall time: 213 ms
❷
❶ Полученное представление Jaxpr содержит только одну строку (хотя и достаточно
сложную).
❷ Компиляция выполняется очень быстро.
В этом варианте мы существенно сократили размер промежу
точного представления, и теперь время компиляции значительно
уменьшилось. Функция lax.scan() – это базисный элемент (primi
tive) JAX, поэтому итоговое представление и получилось таким кро
шечным.
5.3.5
Методы класса
Еще один хитроумный прием – уже используемые нами ранее ан
нотации @jit только для отдельных независимых функций. Но ан
нотация @jit может потребоваться и для методов класса. Простое
аннотирование метода класса с помощью @jit приводит к ошибке.
Листинг 5.23 Аннотация для метода класса
class ScaleClass:
def __init__(self, scale: jnp.array):
self.scale = scale
@jax.jit
def apply(self, x: jnp.array):
return self.scale * x
scale_double = ScaleClass(2)
scale_double.apply(10)
❶
❷
❸
>>> TypeError: Cannot interpret value of type <class '__main__.ScaleClass'>
as an abstract array; it does not have a dtype attribute
# TypeError: невозможно интерпретировать значение типа
# <class '__main__.ScaleClass'> как абстрактный массив;
# это значение не имеет атрибута dtype.
❶ Метод класса, который необходимо JIT-скомпилировать.
❷ Создание экземпляра класса.
❸ Вызов JIT-скомпилированного метода и получение ошибки.
В приведенном выше примере возникло исключение, потому
что для функции класса имеется первый параметр, представляю
щий собой экземпляр этого класса (здесь: ScaleClass). JAX не знает,
как обработать такой тип, так как он не реализован как интерфейс
абстрактного массива. Существует несколько способов устранения
возникшей проблемы. Во-первых, можно использовать вспомога
тельную функцию.
Глава 5
220
Компиляция кода
Листинг 5.24 Использование вспомогательной функции
для метода класса
from functools import partial
class ScaleClass:
def __init__(self, scale: jnp.array):
self.scale = scale
def apply(self, x: jnp.array):
return _apply_helper(self.scale, x)
@partial(jax.jit, static_argnums=0)
def _apply_helper(scale, x):
return scale*x
❶
❷
❸
scale_double = ScaleClass(2)
scale_double.apply(10)
>>> Array(20, dtype=int32, weak_type=True)
❹
❶ Теперь для метода не применяется JIT-компиляция, но используется вспомога-
тельная функция.
❷ Использование аннотации @jit с параметром scale, являющимся статическим.
❸ Вспомогательная функция размещается вне класса.
❹ Теперь все работает корректно.
В приведенном выше примере мы удалили аннотацию @jit из
метода класса, создали вспомогательную функцию за пределами
класса и использовали эту функцию в методе класса. Вспомогатель
ная функция аннотирована с помощью @jit, но мы воспользовались
functools.partial для определения статического параметра. Пер
вый параметр вспомогательной функции был помечен как статиче
ский, поскольку мы предполагаем, что он не будет часто изменять
ся, так как инициализирует экземпляр класса. То же самое значение
должно применяться ко многим различным аргументам.
Также можно сделать статическим параметр self в первом вари
анте исходного кода (листинг 5.23), но при этом потребуется внес
ти больше изменений в код, чтобы избежать неожиданностей. Этот
способ описан здесь: https://docs.jax.dev/en/latest/faq.html#strategy2-marking-self-as-static.
Другой, более гибкий подход: превратить класс ScaleClass в py
tree. Структуры pytree мы будем рассматривать в главе 10, но если
вас прямо сейчас интересует эта методика, то вы найдете ее описа
ние здесь: https://docs.jax.dev/en/latest/faq.html#strategy-3-makingcustomclass-a-pytree.
Резюме
5.3.6
221
Простые функции
Возможны такие ситуации, в которых функция уже имеет неболь
шой размер, и использование JIT-компиляции не обеспечит како
го-либо существенного ускорения. Кроме того, потребуется допол
нительное время для компиляции со всеми накладными расходами
на использование MLIR/XLA и даже время для копирования данных
на аппаратный акселератор, которое, вероятно, не является необ
ходимым. Всегда оценивайте (в числовых характеристиках) выгоду,
которую вы получаете от JIT-компиляции, и пытайтесь скомпилиро
вать наибольший возможный фрагмент вычислений. Такой подход
предоставляет компилятору бóльшую свободу для улучшения опти
мизации.
В этой главе мы рассмотрели способы ускорения выполнения
кода с помощью компиляции и соответствующей трансформации
jit(). В следующих главах мы узнаем, как сделать код еще более
производительным с помощью векторизации и распараллеливания
и соответствующих трансформаций vmap() и pmap(). Приложение D
содержит информацию о pjit(), экспериментальной методике рас
параллеливания, которая была окончательно унифицирована с ис
пользованием трансформации jit().
Упражнение 5.1
Попробуйте реализовать JIT-компилируемую функцию, принимаю
щую временные ряды (time series) и размер окна и вычисляющую
скользящее среднее значение (moving average).
Резюме
JAX использует JIT-компиляцию для генерации эффективного
кода для CPU, GPU и TPU с помощью трансформации jit() или
соответствующей аннотации @jit.
JIT-компиляция выполняется, когда JIT-трансформированная
функция выполняется в первый раз.
JAX предназначен для работы с функционально чистым кодом без
глобального состояния и побочных эффектов. Вы можете продол
жать писать и выполнять функции, не являющиеся чистыми, но
JAX не гарантирует их корректную работу.
Глава 5
222
Компиляция кода
Трансформации JAX сначала преобразовывают функции Python
в простой промежуточный язык Jaxpr (сокращение от JAX Expres
sion), используя трассировку.
По умолчанию для трассировки jit() использует трассировщик
ShapedArray, который имеет конкретную форму, но без конкрет
ных значений. Поэтому компилируемая функция работает со все
ми возможными входными данными, имеющими ту же форму.
Можно использовать средства управления потоком выполнения
и циклы языка Python внутри функций, если они не зависят от
значений входных параметров (но допускается зависимость от их
формы).
Конкретные аргументы можно пометить как статические, чтобы
использовать их конкретные значения для трассировки. Когда
функция вызывается с другим значением, выполняется повтор
ная компиляция.
Также можно использовать структурированные базисные элемен
ты (primitives) управления потоком выполнения из jax.lax, если
необходима проверка условия по значению входного параметра.
JAX использует MLIR и XLA для выполнения оптимизаций, незави
симых и зависимых от целевой платформы, и для генерации ори
ентированного на целевую платформу кода для CPU, GPU и TPU.
XLA и MLIR используют собственные специализированные вход
ные языки или промежуточные представления, называемые HLO
и StableHLO (или MHLO в старых версиях) соответственно.
В дополнение к JIT-компиляции JAX предоставляет API AOTкомпиляции для управления различными этапами процесса ком
пиляции.
Механизм JIT может изменять точные числовые данные в выводе
функций из-за оптимизаций в процессе работы.
Для JIT-трансформированных функций компиляция иногда мо
жет выполняться очень медленно, если исходный код генерирует
слишком большое внутреннее представление Jaxpr.
Для методов классов требуются специализированные методики
работы с механизмом JIT.
6
Векторизация кода
Темы главы:
методики векторизации исходного кода;
управление поведением vmap() с помощью параметров;
анализ типовых вариантов, в которых можно получить
преимущества от автоматической векторизации.
В главе 3 вы узнали, как ускорить вычисления, выполняя их на GPU
и TPU. Затем в главе 5 мы рассмотрели другую возможность уско
рения кода – компиляцию и XLA. А теперь мы будем изучать еще
два способа, позволяющих быстрее выполнять вычисления: автома
тическую векторизацию и распараллеливание. Эта глава полностью
посвящена автоматической векторизации, а в главах 7 и 8 рассмат
ривается распараллеливание вычислений.
Автоматическая векторизация (auto-vectorization) предоставля
ет пользователю несколько преимуществ. Во-первых, она упрощает
процесс программирования, позволяя писать более простые функ
ции для обработки одного элемента, а затем автоматически преоб
разовывает их в более сложные функции, работающие с пакетами
(или массивами) элементов. Во-вторых, автоматическая вектори
зация может ускорять вычисления, если имеющиеся аппаратные
224
Глава 6
Векторизация кода
ресурсы и логика программы позволяют одновременно выполнять
вычисления со многими элементами. Такой подход, как правило,
оказывается быстрее, чем обработка того же массива последователь
но, элемент за элементом. Обычно автоматическая векторизация
не превосходит по скорости векторизацию, выполненную вручную
(хотя и почти не уступает ей). Тем не менее автоматическая векто
ризация оказывается более быстрой в другом измерении: она обес
печивает высокую производительность разработчика и позволяет
сэкономить время по сравнению с векторизацией функции вручную.
В высокопроизводительных вычислениях и глубоком обучении
обычно имеются пакеты элементов, которые необходимо обрабаты
вать одновременно. На основе этой идеи создана методика минибатч градиентного спуска (mini-batch – небольшое подмножество
тренировочного набора данных). Все операции умножения матриц
в прямых и обратных проходах по нейронным сетям организованы
так, что умножение элементов может выполняться одновременно.
Иначе обработка становится слишком неэффективной. Поэтому ис
пользование векторизации имеет чрезвычайно важное значение.
Начнем с нескольких конкретных примеров с простыми функ
циями для демонстрации практического применения автоматиче
ской векторизации, чтобы вы могли полностью понять этот процесс.
Во втором разделе мы более подробно рассмотрим разнообразные
аспекты автоматической векторизации. Третий раздел содержит не
сколько примеров из реальной практики, в которых автоматическая
векторизация оказывается полезной.
6.1
Различные способы векторизации функции
Мы уже встречались с автоматической векторизацией и использо
ванием трансформации vmap() в главе 2. Я напомню принципы, на
которых основана автоматическая векторизация.
Часто в нашем распоряжении находится функция для обработки
одного элемента – для обработки одной точки данных, для приме
нения фильтра к изображению, для применения нейронной сети
к единственному элементу данных и т. п. На рис. 6.1 показан вариант
функции вычисления скалярного произведения, обрабатывающей
пару векторов. Числа в круглых скобках – это кортежи с формами
данных.
Рис. 6.1 Функция для обработки одной пары векторов
225
Различные способы векторизации функции
Во всех подобных вариантах, вероятнее всего, требуется приме
нение одной и той же функции к массиву элементов, или пакету, –
для обработки многочисленных точек данных в загрузчике данных,
одновременное применение фильтра к нескольким изображениям
или применение нейронной сети к пакету данных. На рис. 6.2 по
казана функция, обрабатывающая массивы векторов.
Рис. 6.2 Функция для обработки массивов векторов
Здесь мы обсудим, как перейти от функции, применяемой к одно
му элементу, к функции, применяемой к массиву (или пакету) эле
ментов. Начнем с простой функции, вычисляющей скалярное про
изведение двух векторов (или тензоров ранга 1 с типом Array, если
выражаться точно).
Листинг 6.1 Функция для вычисления скалярного произведения
двух векторов
def dot(v1, v2):
return jnp.vdot(v1, v2)
dot(jnp.array([1., 1., 1.]), jnp.array([1., 2., -1]))
>>> Array(2., dtype=float32)
❶ Использование функции vdot().
❷ Вычисление скалярного произведения двух векторов.
❶
❷
Это элементарная функция, она просто вычисляет скалярное про
изведение пары векторов.
Теперь предположим, что вместо двух векторов у нас имеется два
списка векторов (функции JAX не работают со списками языка Py
Глава 6
226
Векторизация кода
thon, поэтому в действительности мы подразумеваем массив типа
Array, содержащий «список» векторов; с технической точки зрения
это массив с дополнительным измерением, или тензор ранга 2 в на
шем случае) и необходимо вычислить список скалярных произведе
ний соответствующих элементов двух входных списков.
Листинг 6.2
Генерация двух списков векторов
rng_key = random.PRNGKey(42)
vs = random.normal(rng_key, shape=(20,3))
v1s = vs[:10,:]
v2s = vs[10:,:]
❶
❷
❸
❸
❶ Создание ключа генератора случайных чисел (более подробно об этом – в гла-
ве 9).
❷ Генерация двумерного массива случайных чисел.
❸ Разделение сгенерированного массива на две части: первые 10 чисел отправля-
ются в первый список, следующие 10 чисел – во второй список.
Существуют различные способы получения требуемого результа
та. Можно применить то, что я называю простейшими методиками,
которые могут работать (но могут и не сработать), но, как правило,
они являются неэффективными решениями. Можно вручную пере
писать функцию так, чтобы она работала с массивами (т. е. векто
ризовать ее). Или положиться на автоматическую векторизацию,
преобразовывающую функцию, работающую с единственным эле
ментом, в функцию, способную обработать массив элементов. Рас
смотрим подробнее все вышеперечисленные варианты.
6.1.1
Простейшие методики
Существуют две различные простейшие методики – предельно прос
тые и понятные, но далеко не всегда самые эффективные.
Первая простейшая методика
Первая простейшая методика – передача массивов в исходную функ
цию без каких-либо изменений в ее коде. Вообще говоря, существу
ют три возможных результата:
возможно, вы получите то, что хотите. Это происходит, пото
му что NumPy может использовать групповые (broadcasting)
операции
(https://numpy.org/doc/stable/user/basics.broadcast
ing.html) и универсальные функции (ufuncs) (https://numpy.org/
doc/stable/user/basics.ufuncs.html#ufuncs-basics) для векториза
ции операций с массивами. Если вы применяете функцию ак
тивации selu() из предыдущей главы (или почти любую дру
Различные способы векторизации функции
227
гую функцию активации), то получаете неплохой шанс на то,
что функция, предназначенная для работы с одним элементом,
автоматически должна работать и с массивом элементов. Это
удачный вариант, и он может возникнуть, но нельзя же все вре
мя зависеть от случая;
возможно, вы получите ошибку, если функция спроектирована
так, что даже групповые операции и универсальные функции
(ufuncs) не помогут. Это тоже не самый худший вариант, даже
если функция возвращает ошибку. По крайней мере, вы немед
ленно узнаете о том, что случилась неприятность, и можете ис
править ситуацию;
возможно, вы получите некоторый результат (подразумевает
ся, что не возникает никаких ошибок), который является непра
вильным и не соответствует тому, что вы хотели получить. Вот
это действительно плохой вариант, потому что вы можете счи
тать, что все прошло хорошо и программа работает, но на са
мом деле вы получили совсем не те результаты, которые долж
ны быть вычислены. Если полученный результат используется
для других вычислений, которые также используют групповые
операции и автоматическое преобразование типов, то ошибка
может оставаться незамеченной в течение длительного време
ни до тех пор, пока вы (как можно надеяться) окончательно не
убедитесь в том, что получаете неожиданные странные резуль
таты.
При использовании исходной функции из листинга 6.1 мы попа
даем в третий (самый плохой) вариант. Возможно, мы обнаружим
что-то неправильное в форме результата, так как это одно число
вместо массива скалярных произведений.
Листинг 6.3 Простейшая методика применения исходной функции
к двум спискам векторов
v1s.shape, v2s.shape
>>> ((10, 3), (10, 3))
dot(v1s, v2s)
>>> Array(1.0755965, dtype=float32)
❶
❷
❸
❶ Мы применяем исходную функцию к двум массивам с 10 элементами, являющи-
мися векторами с длиной 3.
❷ Функция вызывается без каких-либо ошибок.
❸ Но в результате выводится только одно число вместо 10.
В приведенном выше примере мы получили только одно число
вместо 10 попарных скалярных произведений. Эта ошибка может
228
Глава 6
Векторизация кода
оставаться незамеченной, если полученный результат используется
в последующих вычислениях без каких-либо проверок. Например,
результаты скалярных произведений могут использоваться как веса
для масштабирования некоторых других элементов, и при группо
вых операциях все последующие операции умножения также могут
быть выполнены без ошибок, и вы решите, что все в порядке.
Ошибки такого типа особенно трудно обнаружить, они могут про
являться только при более низком качестве применяемого алгорит
ма, чем предполагалось. Поэтому я предпочитаю не пользоваться
этой весьма ненадежной методикой или по крайней мере выпол
нять дополнительные проверки корректности результата (как ми
нимум по форме).
Вторая простейшая методика
Вторая простейшая методика – применение исходной функции
к каждому входному элементу переданного массива с последующим
объединением всех вычисленных результатов в новый массив.
Листинг 6.4 Простая генерация (вычисление) результатов по одному
(последовательно)
[dot(v1s[i],v2s[i]) for i in range(v1s.shape[0])]
>>> [Array(-0.9443626, dtype=float32),
>>> Array(0.8561607, dtype=float32),
>>> Array(-0.45202938, dtype=float32),
>>> Array(0.7629303, dtype=float32),
>>> Array(-2.06525, dtype=float32),
>>> Array(0.5056444, dtype=float32),
>>> Array(-0.5623387, dtype=float32),
>>> Array(1.5973439, dtype=float32),
>>> Array(1.7121218, dtype=float32),
>>> Array(-0.33462408, dtype=float32)]
❶
❷
❶ Исходная функция применяется поэлементно для конструктора списков Python.
❷ Результат имеет (почти) ожидаемую форму из 10 элементов.
В этом случае мы получили почти ожидаемый результат. Но при
этом необходимо учитывать два нюанса: первый – не слишком зна
чимый, второй – существенный.
Не слишком значимый нюанс заключается в том, что мы полу
чили список языка Python, состоящий из массивов типа Array, а не
один массив Array с многомерным массивом внутри. Вероятно, по
требуется использовать полученный результат в другом вычисле
нии. Но, как вы, возможно, помните из раздела 3.2.2, JAX намеренно
не принимает списки или кортежи в качестве входных данных для
своих функций, потому что это может привести к скрытому от поль
Различные способы векторизации функции
229
зователя снижению производительности, источник которого труд
но обнаружить. Поэтому потребуется преобразование полученного
списка в некоторую другую структуру.
Существенный нюанс – та же причина, по которой в функциях JAX
неприемлемы списки: снижение производительности. Предложен
ная методика проста, но обычно неэффективна. Ниже мы сравним
ее скорость со скоростью методик векторизации.
6.1.2
Векторизация вручную
Стандартным способом устранения неэффективности простейших
методик является переписывание и векторизация вручную функ
ции, работающей с одним элементом, для обеспечения возмож
ности приема пакетов элементов в качестве входных данных. Как
правило, это означает, что входные тензоры будут иметь дополни
тельное измерение для пакетов и потребуется соответствующим
образом переписать код вычислений. Это не вызывает затруднений
для простых вычислений, но все усложняется для нетривиальных
функций.
Воспользуемся весьма мощной функцией из интерфейса NumPy
с именем einsum(). Мы встретимся с ней в приложении D при об
суждении xmap(). А сейчас просто будем считать ее доступным сред
ством для векторизации исходной функции.
Листинг 6.5
Векторизация функции вручную
def dot_vectorized(v1s, v2s):
return jnp.einsum('ij,ij->i',v1s, v2s)
dot_vectorized(v1s, v2s)
>>> Array([-0.9443626 , 0.8561607 , -0.45202938, 0.7629303 ,
>>>
-2.06525
, 0.5056444 , -0.5623387 , 1.5973439 ,
>>>
1.7121218 , -0.33462408], dtype=float32)
❶
❷
❷
❷
❶ Мы переписали исходную функцию для поддержки массивов как входных дан-
ных.
❷ Получен ожидаемый результат.
Этот подход сработал превосходно, и мы получили результат,
имеющий необходимую форму. Решение с применением einsum() не
является единственным, и можно было бы использовать другие ме
тодики для написания векторизованной версии исходной функции.
Но общей основой для всех методик является тот факт, что нам при
ходится переписывать функцию для поддержки массивов. Для более
сложных функций такой подход может оказаться слишком трудоем
ким, и в большинстве случаев в этом процессе велика вероятность
ошибок.
Глава 6
230
6.1.3
Векторизация кода
Автоматическая векторизация
JAX предоставляет альтернативное решение – автоматическую
векторизацию (automatic vectorization). Трансформация vmap() вы
полняет преобразование исходной функции, работающей только
с одним элементом, в функцию, способную обрабатывать пакеты
элементов (см. рис. 6.3).
Входные
данные
Функция обработки
одного элемента
Вывод
результата
Трансформация vmap()
Входные
данные
Функция обработки
одного элемента
Вывод
результата
Рис. 6.3 Преобразование функции обработки одного элемента в функцию,
работающую с пакетами элементов, с помощью трансформации vmap()
Код с использованием vmap() очень прост и понятен.
Листинг 6.6
Автоматическая векторизация функции
dot_vmapped = jax.vmap(dot)
dot_vmapped(v1s, v2s)
>>> Array([-0.9443626 , 0.8561607 , -0.45202938, 0.7629303 ,
>>>
-2.06525
, 0.5056444 , -0.5623387 , 1.5973439 ,
>>>
1.7121218 , -0.33462408], dtype=float32)
❶
❷
❷
❷
❶ Трансформация vmap() создала другую функцию, которая поддерживает массивы
как входные данные.
❷ Получен ожидаемый результат.
Главное достоинство этой методики заключается в том, что она
не изменяет исходную функцию и позволяет получить требуемый
правильный результат.
Преимущество, которое мне особенно нравится, состоит в том,
что исходный код становится более ясным и простым для понима
ния. Это имеет большое значение в сфере ИТ, поскольку почти весь
исходный код требует сопровождения, и другие люди обязаны его
читать и понимать. Код обработки одного элемента обычно легко
понять, тогда как вручную векторизованный код часто бывает весь
Различные способы векторизации функции
231
ма серьезно оптимизирован, что затрудняет его сопровождение. Изза этого другие люди, сопровождающие код, тратят гораздо больше
времени на его понимание или при внесении изменений совершают
ошибки, причиной которых является непонимание.
Приведенный в листинге 6.5 пример с использованием einsum()
понять несложно, если только вы не имеете представления об этой
функции. И вы можете встретить действительно красивые приме
ры ручной векторизации, но я готов поспорить, что это нетипичный
подход.
Теперь у нас есть несколько реализаций, и мы сравним их поведе
ние и производительность.
6.1.4
Сравнение скорости выполнения
Как уже отмечалось ранее, простейшие методики могут обеспечить
интуитивно понятное решение. Но такое решение имеет некоторые
недостатки.
В данном случае для нас важны две метрики, связанные со време
нем: время разработки решения и время выполнения вычисления.
С оценкой времени на разработку решения почти все ясно: автома
тическая векторизация проще, чем формирование цикла, а также,
в нашем случае, проще даже, чем сохранение без изменений кода
функции, потому что при этом потребуется намного больше време
ни, чтобы понять, что сделано неправильно, и, наконец, написать
правильное решение. По сравнению с векторизацией вручную раз
работка с применением автоматической векторизации выполняется
гораздо быстрее.
Сравним вторые временные метрики – скорость выполнения всех
реализаций.
Листинг 6.7 Сравнение скорости выполнения различных вариантов
реализации
%timeit [dot(v1s[i],v2s[i]).block_until_ready() for i in range(v1s.shape[0])] ❶
>>> 4.58 ms ± 152 µs per loop
➥(mean ± std. dev. of 7 runs, 100 loops each)
❷
%timeit dot_vectorized(v1s, v2s).block_until_ready()
❷
>>> 91.9 µs ± 20 µs per loop
➥(mean ± std. dev. of 7 runs, 10000 loops each)
❸
%timeit dot_vmapped(v1s, v2s).block_until_ready()
❸
>>> 831 µs ± 7.7 µs per loop
➥(mean ± std. dev. of 7 runs, 1000 loops each)
❹
❹
Глава 6
232
Векторизация кода
dot_vectorized_jitted = jax.jit(dot_vectorized)
dot_vmapped_jitted = jax.jit(dot_vmapped)
❺
❺
# Разогрев.
dot_vectorized_jitted(v1s, v2s);
dot_vmapped_jitted(v1s, v2s);
%timeit dot_vectorized_jitted(v1s, v2s).block_until_ready()
>>> 6.93 µs ± 194 ns per loop
➥(mean ± std. dev. of 7 runs, 100000 loops each)
❻
%timeit dot_vmapped_jitted(v1s, v2s).block_until_ready()
>>> 7.6 µs ± 1.42 µs per loop
➥(mean ± std. dev. of 7 runs, 100000 loops each)
❶
❷
❸
❹
❺
❻
❻
Не следует забывать об асинхронной диспетчеризации.
Простейшая методика является самой медленной.
Вручную векторизованная функция показывает превосходную производительность.
Автоматически векторизованная функция медленнее, но все же остается очень быстрой.
JIT-трансформация обеих векторизованных функций.
После JIT-трансформации скорости решений с ручной и автоматической векторизацией стали гораздо более близкими друг к другу.
Здесь можно видеть, что реализация простейшей методики чрез
вычайно медленная и неэффективная. Вручную векторизованная
функция демонстрирует наилучшую производительность, а ав
томатически векторизованная функция немного медленнее, но
на порядок быстрее реализации простейшей методики. После JITкомпиляции обоих решений их скорости становятся сравнимыми;
версия с автоматической векторизацией чуть медленнее (но их до
верительные интервалы пересекаются, поэтому мы не можем оце
нить скорость более точно).
Рассмотрим внутреннее устройство и генерацию представлений
Jaxpr для этих функций.
Листинг 6.8
Получение внутренних представлений Jaxpr
jax.make_jaxpr(dot)(jnp.array([1., 1., 1.]
jnp.array([1., 1., -1]))
>>> { lambda ; a:f32[3] b:f32[3]. let
>>>
c:f32[] = dot_general[
>>>
dimension_numbers=(([0], [0]), ([], []))
>>>
preferred_element_type=float32
>>>
] a b
>>>
in (c,) }
jax.make_jaxpr(dot_vectorized)(v1s, v2s)
❶
❷
233
Управление поведением vmap()
>>> { lambda ; a:f32[10,3] b:f32[10,3]. let
>>>
c:f32[10] = dot_general[
>>>
dimension_numbers=(([1], [1]), ([0], [0]))
>>>
preferred_element_type=float32
>>>
] a b
>>>
in (c,) }
jax.make_jaxpr(dot_vmapped)(v1s, v2s)
>>> { lambda ; a:f32[10,3] b:f32[10,3]. let
>>>
c:f32[10] = dot_general[
>>>
dimension_numbers=(([1], [1]), ([0], [0]))
>>>
preferred_element_type=float32
>>>
] a b
>>>
in (c,) }
❸
❶ Представление Jaxpr исходной невекторизованной функции.
❷ Вручную векторизованная функция.
❸ Автоматически векторизованная функция.
Здесь можно видеть, что исходная невекторизованная функция
использует функцию dot_general() из пакета jax.lax (https://docs.
jax.dev/en/latest/_autosummary/jax.lax.dot_general.html).
Обе версии с использованием einsum() и vmap() содержат одина
ковый код, вызывающий функцию dot_general(). В этих решениях
нет циклов, и сгенерированный код эффективен.
6.2
Управление поведением vmap()
Трансформацию vmap() в базовом варианте использовать относи
тельно просто, но существует множество вариантов, в которых тре
буется более точное и детализированное управление. Например,
массивы могут быть скомпонованы по-разному, т. е. измерение па
кетов не является первым, или в качестве параметров используются
более сложные структуры, например словарь dict. Функция vmap()
предоставляет удобные способы работы с разнообразными структу
рами тензоров.
6.2.1
Управление осями массива для выполнения
преобразования
Вы можете управлять выбором осей массива для выполнения пре
образований. Для этого функция vmap() предоставляет параметр
in_axes. Значением этого параметра может быть целое число, None
или (возможно, вложенный) стандартный контейнер Python, такой
как кортеж (tuple), список (list) или словарь (dict).
Глава 6
234
Векторизация кода
Если параметр in_axes содержит целое число (по умолчанию уста
новлено значение 0), то заданная этим числом ось массива исполь
зуется для преобразования с участием всех аргументов функции.
В примере из листинга 6.6 этот параметр не был задан явно, поэтому
функция выполняла преобразование первой оси (с индексом 0) для
каждого аргумента, и измерением пакетов было измерение с индек
сом 0.
Предположим, что необходимо использовать различные индексы
для разных параметров. В этом случае можно воспользоваться кор
тежем, включающим целые числа и значения None, с длиной, равной
количеству позиционных аргументов исходной функции. Значение
None сообщает, что соответствующий параметр не требует преобра
зования. Общее правило: структура in_axes должна соответствовать
структуре передаваемых входных данных.
Вызов vmap() в коде листинга 6.6 равнозначен коду в листинге 6.9.
Листинг 6.9
Использование параметра in_axes
dot_vmapped = jax.vmap(dot, in_axes=(0,0))
❶
❶ Это равнозначно пропуску параметра in_axes в рассматриваемом примере.
На рис. 6.4 изображены два транспонированных массива, в кото
рых векторы расположены по горизонтальному измерению тензора.
Рис. 6.4 Функция для обработки транспонированных массивов векторов
Управление поведением vmap()
235
Если необходимо преобразовать транспонированные массивы,
как показано на рис. 6.4, то можно использовать значение парамет
ра in_axes=(1,1).
Рассмотрим еще более сложный пример с различными осями
и осями, которые не должны быть преобразованы. В функцию ска
лярного произведения добавлен новый параметр koeff. Функция
вычисляет то же скалярное произведение, что и ранее, но теперь
еще и с умножением на коэффициент, переданный как отдельный
параметр. Массивы, к которым нужно применить эту функцию,
структурированы по-разному: первый сохраняет форму, заданную
ранее, а второй представляет собой транспонированную версию ис
ходного массива, и измерение пакетов теперь имеет индекс 1. Такая
ситуация вполне может возникать естественным образом, если вы
обрабатываете данные из нескольких источников, скомпонованные
по-разному.
Листинг 6.10 Функция вычисления скалярного произведения
с масштабным коэффициентом
def scaled_dot(v1, v2, koeff):
return koeff*jnp.vdot(v1, v2)
v1s_ = v1s
v2s_ = v2s.T
k = 1.0
v1s_.shape, v2s_.shape
❶
❷
❸
❹
>>> ((10, 3), (3, 10))
❶ Добавлен еще один параметр (по сравнению с предыдущей версией).
❷ Первый массив тот же самый.
❸ Второй массив транспонирован.
❹ Значение коэффициента (постоянное для всех элементов массива).
В этом случае может потребоваться преобразование данных, что
бы привести все к одной схеме размещения. Другой вариант: можно
пропустить этот шаг (и, возможно, сократить количество вычисле
ний, и даже весьма существенно) и позволить пакетной функции
узнать, что входные массивы организованы по-разному. Эта новая
структура массива и функция scaled_dot() могут выглядеть так, как
показано на рис. 6.5.
Теперь необходимо применить эту функцию к массивам с по
мощью трансформации vmap(). Для работы со структурой аргу
мента вы должны передать параметр in_axes с кортежем, описы
вающим, какую ось нужно выбрать для преобразования в каждом
аргументе.
236
Глава 6
Векторизация кода
Рис. 6.5 Функция для вычисления скалярного произведения с масштабным
коэффициентом
Листинг 6.11 Использование параметра in_axes
для различающихся аргументов
scaled_dot_batched = jax.vmap(scaled_dot, in_axes=(0,1,None))
scaled_dot_batched(v1s_, v2s_, k)
❶
>>> Array([-0.9443626 , 0.8561607 , -0.45202938, 0.7629303 ,
>>>
-2.06525
, 0.5056444 , -0.5623387 , 1.5973439 ,
>>>
1.7121218 , -0.33462408], dtype=float32)
❶ Передача параметра in_axes, содержащего кортеж.
В рассматриваемом здесь примере мы сообщили трансформации
vmap() о необходимости итерации по оси с индексом 0 для первого
параметра, итерации по оси 1 для второго параметра, а также об от
сутствии итерации по третьему параметру, который содержит прос
той скаляр, а не массив.
Параметр in_axes также может быть стандартным контейнером
Python, возможно, вложенным.
Листинг 6.12 Использование параметра in_axes с контейнером
Python
def scaled_dot(data, koeff):
return koeff*jnp.vdot(data['a'], data['b'])
scaled_dot_batched=jax.vmap(scaled_dot,
in_axes=({'a':0,'b':1},None))
scaled_dot_batched({'a':v1s_, 'b': v2s_}, k)
❶
❷
Управление поведением vmap()
237
>>> Array([-0.9443626 , 0.8561607 , -0.45202938, 0.7629303 ,
>>>
-2.06525
, 0.5056444 , -0.5623387 , 1.5973439 ,
>>>
1.7121218 , -0.33462408], dtype=float32)
❶ Теперь функция принимает словарь и скалярное значение.
❷ Пометка осей для словаря и скалярного параметра.
В этом примере передается кортеж, содержащий словарь (dict)
и скалярное значение, что соответствует параметрам функции. Сло
варь определяет преобразование осей для каждого элемента и соот
ветствует первому параметру функции, а скалярное значение (здесь:
None) указывает на отсутствие преобразования для второго пара
метра функции.
6.2.2
Управление осями выходного массива
Также можно управлять схемой размещения получаемых результа
тов (вывода), если результат имеет более одного измерения и не
обходимо, чтобы измерение пакетов не являлось самым первым.
Подобное преобразование может потребоваться просто потому, что
следующая функция в конвейере принимает входные данные имен
но в таком формате.
Предположим, что имеется функция для масштабирования век
тора с некоторым коэффициентом и нужно получить на выходе ре
зультат, транспонированный по сравнению со схемой размещения
входного вектора, как показано на рис. 6.6.
Рис. 6.6 Функция для масштабирования и транспонирования векторов
Для этого существует параметр out_axes.
Листинг 6.13 Использование параметра out_axes
def scale(v, koeff):
return koeff*v
scale_batched = jax.vmap(scale,
in_axes=(0,None),
out_axes=(1))
❶
❷
❸
Глава 6
238
Векторизация кода
scale_batched(v1s, 2.0)
>>> Array([[-1.4672383 , -1.6510035 , 3.5308602 , -2.2189112 , 0.3024418 ,
>>>
0.7649379 , -4.028754 , -3.0968533 , 0.34476107, -2.9087348 ],
>>>
[-1.5357308 , -0.7061183 , 4.0082793 , -0.69232166, -3.2186441 ,
>>>
2.0812016 , 3.585087 , 0.15288436, 2.0001278 , 2.0246687 ],
>>>
[-1.6228952 , 1.5497094 , -3.2013843 , 0.50127184, -0.2000112 ,
>>>
1.6244346 , 0.17156784, -1.3113307 , -2.532448 , -1.60574
]],
>>>
dtype=float32)
❶ Простая функция для масштабирования вектора.
❷ Определение измерений пакетов для входных данных.
❸ Определение измерений пакетов для вывода результата.
В этом примере функция возвращает масштабированный вектор,
а пакетная версия возвращает массив масштабированных векторов.
Обратите внимание: необходимо, чтобы возвращаемый массив был
транспонированным, и измерение пакетов должно иметь индекс 1,
а не 0. Иногда это помогает избежать дополнительных трансформа
ций и достичь желаемой цели, используя только vmap().
6.2.3
Использование именованных аргументов
Иногда функции используют именованные аргументы вместо чисто
позиционных. Функции с именованными аргументами проще чи
тать, и труднее сделать ошибку при вызове таких функций, перепу
тав значения различных параметров.
Весьма важно знать следующее об аргументах, передаваемых как
ключевые слова: они всегда отображаются (выполняют преобразо
вание) на соответствующую им ось (с индексом 0). Иначе можно по
лучить неожиданное сообщение об ошибке.
Воспользуемся той же функцией scale() из предыдущего приме
ра, но сделаем ее второй параметр именованным аргументом, что
естественно для этой конкретной функции. После внесения изме
нения попытка выполнения кода предыдущего примера приводит
к ошибке.
Листинг 6.14
Использование именованных аргументов
def scale(v, koeff=1.0):
return koeff*v
❶
scale_batched = jax.vmap(scale,
in_axes=(0,None),
out_axes=(1))
scale_batched(v1s, koeff=2.0)
❷
239
Управление поведением vmap()
❸
>>> …
>>> ValueError: vmap in_axes specification must be
a tree prefix of the corresponding value, got
specification (0, None) for value tree PyTreeDef((*,)).
>>> …
# ValueError: спецификация vmap in_axes должна быть
# древовидным префиксом соответствующего значения,
# получена спецификация (0, None) для значения
# дерева PyTreeDef((*,)).
scale_batched = jax.vmap(scale,
in_axes=(0),
out_axes=(1))
❹
scale_batched(v1s, koeff=2.0)
❺
>>> …
>>> ValueError: vmap was requested to map its argument
along axis 0, which implies that its rank should be
at least 1, but is only 0 (its shape is ())
>>> …
# ValueError: vmap затребовал для преобразования свой
# аргумент по оси 0, предполагая, что его ранг должен
# быть как минимум 1, но получил только 0 (его форма ())
❶ Коэффициент koeff сделан именованным аргументом со значением по умол
чанию.
❷ Вызов функции с этим именованным параметром.
❸ Возникла какая-то непонятная ошибка.
❹ Корректировка in_axes, поскольку этот параметр работает только с позиционны-
ми аргументами.
❺ Возникает новая непонятная ошибка.
Мы получили какую-то непонятную ошибку, потому что параметр
in_axes предназначен для позиционных, а не для именованных ар
гументов (ключевых слов). Кажется, что ошибку легко исправить,
заменив значение параметра in_axes с (0, None) на (0) для соответ
ствия изменению в позиционных и именованных аргументах. Но
это не помогает, и мы получаем другую, тоже непонятную ошибку.
В данном случае vmap() пытается отобразить параметр, для которо
го нет измерения для преобразования. Это именно параметр koeff,
и, как было отмечено ранее, именованные параметры всегда ото
бражаются на соответствующую им первую ось.
Существует несколько способов устранения такой ошибки. Можно
вернуться к использованию позиционных параметров. Можно на
писать функцию-обертку, которая скроет этот параметр. Можно пе
реслать именованный параметр в массив с требуемой размерностью
для отображения. Мы продемонстрируем два последних подхода.
Глава 6
240
Векторизация кода
Листинг 6.15 Изменение кода для работы с именованными аргументами
from functools import partial
scale2 = partial(scale, koeff=2.0)
scale_batched = jax.vmap(scale2,
in_axes=(0),
out_axes=(1))
scale_batched(v1s)
❶
❷
>>> Array([[-1.4672383 , -1.6510035 , 3.5308602 , -2.2189112 , 0.3024418 ,
>>>
0.7649379 , -4.028754 , -3.0968533 , 0.34476107, -2.9087348 ],
>>>
[-1.5357308 , -0.7061183 , 4.0082793 , -0.69232166, -3.2186441 ,
>>>
2.0812016 , 3.585087 , 0.15288436, 2.0001278 , 2.0246687 ],
>>>
[-1.6228952 , 1.5497094 , -3.2013843 , 0.50127184, -0.2000112 ,
>>>
1.6244346 , 0.17156784, -1.3113307 , -2.532448 , -1.60574
]],
>>>
dtype=float32)
❸
scale_batched = jax.vmap(scale,
in_axes=(0),
out_axes=(1))
❹
scale_batched(v1s, koeff=jnp.broadcast_to(2.0, (v1s.shape[0],)))
❺
❹
>>> Array([[-1.4672383 , -1.6510035 , 3.5308602 , -2.2189112 , 0.3024418 ,
>>>
0.7649379 , -4.028754 , -3.0968533 , 0.34476107, -2.9087348 ],
>>>
[-1.5357308 , -0.7061183 , 4.0082793 , -0.69232166, -3.2186441 ,
>>>
2.0812016 , 3.585087 , 0.15288436, 2.0001278 , 2.0246687 ],
>>>
[-1.6228952 , 1.5497094 , -3.2013843 , 0.50127184, -0.2000112 ,
>>>
1.6244346 , 0.17156784, -1.3113307 , -2.532448 , -1.60574
]],
>>>
dtype=float32)
❻
❶ Создание функции partial с фиксированным значением параметра koeff.
❷ Отображается один позиционный параметр.
❸ Теперь все работает и выдает правильные результаты.
❹ Использование старой функции с позиционным и именованным параметрами.
❺ Перед вызовом функции мы изменяем именованный параметр так, чтобы он содержал ось
для отображения.
❻ Этот вариант тоже работает и выдает правильные результаты.
В приведенном выше примере мы создали новую функцию с по
мощью стандартного объекта Python partial, который предоставля
ет функцию с фиксированным значением параметра koeff (другой
вариант: можно самостоятельно написать функцию-обертку). При
использовании этой функции все работает правильно.
При другом подходе мы сохранили старую функцию и сделали
параметр ключевого слова массивом с осью для отображения. Этот
массив заполняется одинаковым значением коэффициента, кото
рое необходимо использовать (для этого применяется стандартная
241
Управление поведением vmap()
функция broadcast_to() из NumPy API). При таком подходе мы так
же получили правильные результаты.
6.2.4
Использование стиля декоратора
Как и многие другие функциональные трансформации, vmap() мож
но использовать с декораторами. В зависимости от конкретной си
туации такой подход может позволить получить более понятный код
без использования временных функций.
Перепишем функцию scale() из листинга 6.13 с применением де
коратора.
Листинг 6.16 Использование декоратора
from functools import partial
@partial(jax.vmap, in_axes=(0,None), out_axes=(1))
def scale(v, koeff):
return koeff*v
❶
❷
scale(v1s, 2.0)
>>> Array([[-1.4672383 , -1.6510035 , 3.5308602 , -2.2189112 , 0.3024418 ,
>>>
0.7649379 , -4.028754 , -3.0968533 , 0.34476107, -2.9087348 ],
>>>
[-1.5357308 , -0.7061183 , 4.0082793 , -0.69232166, -3.2186441 ,
>>>
2.0812016 , 3.585087 , 0.15288436, 2.0001278 , 2.0246687 ],
>>>
[-1.6228952 , 1.5497094 , -3.2013843 , 0.50127184, -0.2000112 ,
>>>
1.6244346 , 0.17156784, -1.3113307 , -2.532448 , -1.60574
]],
>>>
dtype=float32)
❸
❶ Импорт стандартного модуля Python functools.
❷ Перемещение вызова vmap в декоратор.
❸ Этот подход работает и выдает правильные результаты.
В приведенном выше примере мы переместили отдельный вызов
vmap() с его параметрами в декоратор и исключили дополнительный
функциональный объект.
Это удобно, если нет необходимости в отдельной функции, рабо
тающей с единственным элементом, и вас интересует исключитель
но пакетная функция. Таким образом, декоратор помогает написать
простой код для обработки одного элемента и преобразовать его
в код, предназначенный для работы с пакетами элементов.
6.2.5
Использование коллективных операций
Предположим, что вы пишете код, в котором требуется обмен данны
ми между различными элементами пакета (это тот же случай, в ко
тором вычисления распараллеливаются по нескольким устройствам
Глава 6
242
Векторизация кода
и необходима передача информации между устройствами). Для та
кого варианта JAX предоставляет коллективные операции (collective
operations (ops)) (https://docs.jax.dev/en/latest/jax.lax.html#paralleloperators). Обычно это операции с префиксом jax.lax.p*.
Коллективные операции в основном используются для распарал
леливания (это тема следующей главы), чтобы обеспечить обмен
информацией между устройствами, и изначально были предназна
чены для pmap(). Но они также работают с vmap() и способны обеспе
чить весьма полезное решение, если вы занимаетесь реализацией
чего-то вроде пакетной нормализации, когда требуется вычислять
статистические данные по пакетам и модифицировать элементы
пакетов.
Коллективные операции работают следующим образом:
при векторизации вычислений с помощью vmap() по некоторой
оси можно определить имя этой оси, используя аргумент axis_
name. Разрешается передавать любое имя, которое поможет вам
отличать именованную ось от всех прочих. Имя – это просто
метка, представленная строкой;
на именованную ось можно ссылаться внутри коллективных
операций, используя тот же аргумент axis_name. Операция бу
дет выполняться по этой именованной оси.
Рассмотрим простой вариант нормализации значений массива
так, чтобы они в сумме давали 1. В листинге 6.17 приведен соответ
ствующий код.
Листинг 6.17 Использование коллективных операций и параметра
axis_name
arr = jnp.array(range(50))
arr
>>> Array([ 0, 1, 2, 3,
>>>
11, 12, 13, 14,
>>>
21, 22, 23, 24,
>>>
31, 32, 33, 34,
>>>
41, 42, 43, 44,
>>>
dtype=int32)
❶
4,
15,
25,
35,
45,
5,
16,
26,
36,
46,
6,
17,
27,
37,
47,
7,
18,
28,
38,
48,
8, 9, 10,
19, 20,
29, 30,
39, 40,
49],
norm = jax.vmap(
lambda x: x/jax.lax.psum(x, axis_name='batch'),
axis_name='batch')
norm(arr)
>>> Array([0.
,
>>>
0.00408163,
>>>
0.00816326,
>>>
0.0122449 ,
0.00081633,
0.00489796,
0.00897959,
0.01306122,
0.00163265,
0.00571429,
0.00979592,
0.01387755,
0.00244898,
0.00653061,
0.01061224,
0.01469388,
❷
❸
❹
0.00326531,
0.00734694,
0.01142857,
0.0155102 ,
243
Варианты использования vmap() из реальной практики
>>>
>>>
>>>
>>>
>>>
>>>
>>>
0.01632653, 0.01714286,
0.02040816, 0.02122449,
0.0244898 , 0.02530612,
0.02857143, 0.02938776,
0.03265306, 0.03346939,
0.03673469, 0.03755102,
dtype=float32)
0.01795918,
0.02204082,
0.02612245,
0.03020408,
0.03428571,
0.03836735,
0.01877551,
0.02285714,
0.02693878,
0.03102041,
0.03510204,
0.03918367,
0.01959184,
0.02367347,
0.0277551 ,
0.03183673,
0.03591837,
0.04
],
jnp.sum(norm(arr))
>>> Array(1., dtype=float32)
❶
❷
❸
❹
❺
❺
Генерация массива для демонстрации.
Использование операции psum() с параметром axis_name=’batch’.
Использование параметра axis_name=’batch’ в vmap().
Применение созданной функции нормализации.
Проверка нормализованных значений.
В приведенном выше примере мы пометили ось, по которой век
торизуются вычисления, с помощью axis_name='batch' в списке ар
гументов vmap().
Внутри векторизованной функции используется простое выраже
ние x/jax.lax.psum(x, axis_name='batch'), где x – текущий элемент
обрабатываемого пакета, а jax.lax.psum(x, axis_name='batch') вы
полняет all-reduce суммирование по отображаемой оси axis_name.
Поэтому операция psum() возвращает сумму значений по изме
рению пакетов, а функция делит текущий элемент на полученную
сумму. Такой подход эффективно нормализует каждый элемент так,
чтобы сумма элементов была равна 1,0.
Существуют и другие коллективные операции: pmin() и pmax()
для вычисления минимума и максимума, pmean() для вычисления
среднего значения, all_gather() для объединения значений из всех
реплик и некоторые другие полезные операции (https://docs.jax.dev/
en/latest/jax.lax.html#parallel-operators), которые могут потребо
ваться для вычислений.
Теперь вы знакомы с основами автоматической векторизации,
и мы попробуем применить полученные знания в некоторых уже
известных вариантах из реальной практики.
6.3
Варианты использования vmap()
из реальной практики
Существует множество вариантов, в которых может потребоваться
применение автоматической векторизации. Рассмотрим несколько
типовых вариантов, где можно извлечь преимущества от использо
вания vmap().
Глава 6
244
Векторизация кода
Приведенные выше примеры в основном были ориентированы
на варианты обработки пакетов элементов функцией, изначально
предназначенной для обработки одного элемента. Несмотря на то
что многие варианты на определенном уровне можно свести к это
му общему шаблону, существуют семантически различные вариан
ты, с которыми вы встретитесь, и для них vmap() может оказаться
весьма полезным. Начнем с хорошо известного варианта, немного
расширим его, а затем рассмотрим некоторые другие варианты ис
пользования vmap().
6.3.1
Обработка пакетов данных
Вероятно, это самый простой и понятный вариант. Имеется функция
для обработки одного элемента какого-то набора (скажем, одно изо
бражение), и требуется применить эту функцию к пакету элементов.
Если вы не получаете каждый элемент по отдельности, а принимае
те (или по крайней мере можете принимать) элементы в пакетах, то
необходимость использования vmap() вполне очевидна.
Здесь существует один нюанс: элементы не должны обмениваться
информацией друг с другом, и это несомненно часто встречающаяся
ситуация. Это может оказаться не таким уж простым делом, если вы
имеете дело с некоторым типом нормализации, требующим статис
тических характеристик по пакетам или каких-либо текущих вычис
лений вне пределов каждого отдельного элемента.
Этот вариант применяется и к дополнению данных. Вы можете
иметь в своем распоряжении несколько функций для выполнения
разнообразных дополнений и сами решаете, какую из них (или не
которую их комбинацию) применить к каждому конкретному эле
менту. Здесь могут оказаться полезными базисные элементы управ
ления потоком выполнения из пакета jax.lax.
Рассмотрим вариант модели применения случайных дополнений
из главы 3.
Листинг 6.18
Дополнение одного элемента данных
add_noise_func = lambda x: x+10
horizontal_flip_func = lambda x: x+1
rotate_func = lambda x: x+2
adjust_colors_func = lambda x: x+3
augmentations = [
add_noise_func,
horizontal_flip_func,
rotate_func,
adjust_colors_func
]
❶
❶
❶
❶
❶
❶
❶
❶
Варианты использования vmap() из реальной практики
245
def random_augmentation(image, augmentations, rng_key):
'''A function that applies a random transformation to an image'''
# Функция, применяющая случайно выбранную трансформацию к изображению.
augmentation_index = random.randint(
key=rng_key, minval=0, maxval=len(augmentations), shape=())
augmented_image = lax.switch(augmentation_index, augmentations, image)
return augmented_image
❷
❸
image = jnp.array(range(100))
augmented_image = random_augmentation(image, augmentations, random.PRNGKey(211)) ❹
❶ Список из четырех функций-заглушек для обработки изображения (исключительно для де-
монстрационных целей).
❷ Функция для применения случайно выбранной трансформации из набора доступных.
❸ Некоторые фиктивные демонстрационные данные.
❹ Применение случайно выбранной трансформации.
В приведенном выше примере функция random_augmentation()
принимает изображение для дополнения, список функций допол
нения для выбора из него и ключ (или состояние) для генератора
случайных чисел (более подробно генерация случайных чисел рас
сматривается в главе 9). Функция использует lax.switch() для слу
чайного выбора одного из предложенных вариантов дополнения.
Чтобы применить эту функцию к пакету изображений, необхо
димо использовать различные ключи генератора случайных чисел
для каждого вызова (иначе каждый вызов будет использовать один
и тот же ключ, генерировать одно и то же «случайное» число и при
менять одну и ту же трансформацию). При таких условиях примене
ние vmap() становится относительно простым, и вам нужно обратить
внимание только на аргумент in_axes, чтобы пометить массивы, со
ответствующие измерениям пакетов.
Листинг 6.19
Дополнение многих элементов данных
images = jnp.repeat(
jnp.reshape(image, (1,len(image))),
10, axis=0)
images.shape
❶
>>> (10, 100)
rng_keys = random.split(random.PRNGKey(211), num=len(images))
random_augmentation_batch = jax.vmap(
random_augmentation, in_axes=(0,None,0))
augmented_images = random_augmentation_batch(
images, augmentations, rng_keys)
❶ Генерация пакета фиктивных демонстрационных данных.
❷
❸
❹
Глава 6
246
Векторизация кода
❷ Генерация массива ключей генератора случайных чисел.
❸ Автоматическая векторизация исходной функции.
❹ Применение векторизованной функции к пакету данных.
В приведенном выше примере мы сообщили функции vmap()
о том, что изображения размещены в пакете по первому измере
нию тензора изображений, список дополнений не является пакетом
и остается одним и тем же для каждого вызова, а ключи для генера
тора случайных чисел также представлены в виде пакета. Возможно,
это не самая эффективная реализация, поскольку имеет смысл вы
нести генератор случайных чисел за пределы функции и передавать
непосредственно массив случайных чисел вместо ключей. Но этот
вариант реализации функции и измерение ее производительности
я предлагаю читателям выполнить как учебное упражнение. Таким
образом, вы можете продолжить создание конвейера дополнения
и модификации данных, добавляя полезные функции и усложняя
логику дополнения.
6.3.2
Пакетная обработка моделей нейронных сетей
Это еще один простой и понятный вариант использования, похожий
на предыдущий во многих аспектах. Но сначала необходимо обра
тить особое внимание на несколько фактов.
Подобный подход предоставляет значительное преимущество
для разработчиков: они получают возможность писать более прос
той код, не беспокоясь об измерении пакетов. Такой код обычно
гораздо проще понять (и изменять), чем написанный вручную код
векторизации. При этом вы пользуетесь почти всеми преимущест
вами полностью векторизованного кода, включая его высокую про
изводительность.
Мы уже применяли vmap() для такой тренировки нейронной сети
в разделе 2.4. Здесь повторно приводится соответствующий код из
главы 2 (листинг 6.20).
Листинг 6.20
Пакетирование прогнозов нейронной сети
import jax.numpy as jnp
from jax.nn import swish
from jax import vmap
def predict(params, image):
"""Function for per-example predictions."""
# Функция прогнозирования по каждому образцу.
activations = image
for w, b in params[:-1]:
outputs = jnp.dot(w, activations) + b
activations = swish(outputs)
❶
247
Варианты использования vmap() из реальной практики
final_w, final_b = params[-1]
logits = jnp.dot(final_w, activations) + final_b
return logits
batched_predict = vmap(predict, in_axes=(None, 0))
❶ Функция прогнозирования по каждому образцу.
❷ Создание функции, работающей с пакетами изображений.
❷
У нас имеется полностью понятная функция для прогнозирова
ния по отдельному элементу, затем мы создаем ее пакетную версию,
применяя трансформацию vmap(). Для такого простого примера
можно без затруднений написать функцию, работающую с пакета
ми, но такой подход становится неэффективным для более сложных
вариантов.
Однако здесь есть один нюанс. Многие современные модели ис
пользуют статистические характеристики по пакетам и/или состоя
ние внутри модели. Например, именно такой вариант представляет
собой хорошо известная нормализация пакетов. В подобных случаях
вы не находитесь в режиме «пишем функцию для обработки одного
элемента, затем применяем к ней vmap». А кроме того, для таких мо
делей не так-то просто использовать vmap().
В дальнейшем мы вернемся к более сложным нейронным сетям,
когда будем подробно рассматривать фреймворк высокого уровня
для нейросетей Flax. Предположим, что вы создаете сложные ней
ронные сети с нуля, используя только чистый JAX. В таком случае,
возможно, потребуется использование коллективных операций из
подраздела 6.2.3, чтобы реализовать взаимодействие между элемен
тами пакетов или векторизацию кода вручную.
6.3.3
Поэлементные градиенты
Существует вариант, связанный с тренировкой нейронной сети, –
получение градиентов по отдельным выборкам, как уже отмечалось
в подразделе 4.2.3. В тех случаях, когда требуются градиенты по от
дельным выборкам, но в то же время весьма нежелательно снижение
производительности пакетной обработки, JAX предоставляет прос
той способ решения этой задачи:
1 создать функцию для формирования прогноза по одной выбор
ке, например predict(model_params, x);
функцию потерь для вычисления оценки одного прог
ноза, например loss_fn(model_parameters, x, y);
3 создать функцию вычисления градиентов для прогноза по од
ной выборке с помощью трансформации grad(), т. е. применить
grad(loss_fn);
4 создать пакетную версию функции вычисления градиента с по
мощью трансформации vmap(): vmap(grad(loss_fn)). Другим
2 создать
Глава 6
248
Векторизация кода
способом сделать это невозможно – сначала получив вектори
зованную версию, а затем вычисляя ее градиент, – потому что
потери всегда являются скалярным значением, поэтому grad()
вернет ошибку, если применить эту трансформацию к функции,
работающей с векторами. Ранее мы уже выполняли похожую
работу, но в то время не было необходимости в вычислении
потерь по каждому элементу с последующим объединением
результатов внутри функции вычисления потерь, поэтому она
возвращала одно значение для пакета, и вычисление соответ
ствующего градиента работало нормально;
5 дополнительно (но не обязательно): скомпилировать получен
ную функцию с помощью трансформации jit(), чтобы сделать
ее более эффективной для выполнения на внутреннем аппарат
ном оборудовании, получив в итоге jit(vmap(grad(loss_fn)))
(model_params, batch_x, batch_y).
Возможно, все это было бы трудно сделать в других фреймворках,
если метод градиентного спуска обновляет используемые градиен
ты, объединенные по пакетам, и не предоставляет простого способа
извлечения градиентов по каждой отдельной выборке из внутрен
него механизма процедуры. В листинге 6.21 показан небольшой
пример вычисления градиентов по отдельным выборкам для задачи
линейной регрессии, которую мы использовали в главе 4.
Листинг 6.21
Вычисление градиентов по отдельным выборкам
from jax import grad, vmap, jit
x = jnp.linspace(0, 10*jnp.pi, num=1000)
e = 10.0*random.normal(random.PRNGKey(42), shape=x.shape)
y = 65.0 + 1.8*x + 40*jnp.cos(x) + e
model_parameters = jnp.array([1., 1.])
def predict(theta, x):
w, b = theta
return w * x + b
❶
❶
❶
❷
def loss_fn(model_parameters, x, y):
prediction = predict(model_parameters, x)
return (prediction-y)**2
❸
grads_fn = jit(vmap(grad(loss_fn), in_axes=(None, 0, 0)))
❹
batch_x, batch_y = x[:32], y[:32]
grads_fn(model_parameters, batch_x, batch_y)
>>> Array([[
0.
, -213.84189 ],
249
Варианты использования vmap() из реальной практики
>>>
>>>
>>> …
>>>
❶
❷
❸
❹
[ -5.541931, -176.22874 ],
[ -11.923036, -189.57124 ],
[-177.83148 , -182.41585 ]], dtype=float32)
Некоторые зашумленные данные для задачи регрессии.
Простая модель линейной регрессии.
Функция для вычисления ошибки прогноза по одному элементу (экземпляру).
Создание JIT-скомпилированной пакетной версии функции вычисления гради
ента.
Мы использовали простую функцию потерь, применимую к одно
му элементу (предыдущая версия этой функции потерь из главы 4
использовала агрегацию потерь). Затем мы применили последова
тельность трансформаций: сначала получили функцию вычисления
градиента, потом векторизовали ее, чтобы она стала применимой
к пакетам данных, и, наконец, выполнили JIT-компиляцию итоговой
функции. Обратите внимание, с какой легкостью все эти трансфор
мации комбинируются в JAX.
6.3.4
Векторизация циклов
Вернемся к примеру использования фильтров изображений из гла
вы 3. У нас имеется код в листинге 6.22 для функций обработки изо
бражений.
Листинг 6.22
Применение матричных фильтров к изображению
import jax.numpy as jnp
from jax.scipy.signal import convolve2d
from skimage.io import imread
from skimage.util import img_as_float32
from matplotlib import pyplot as plt
kernel_blur = jnp.ones((5,5))
kernel_blur /= jnp.sum(kernel_blur)
❶
❶
❶
❶
❶
❷
❷
def color_convolution(image, kernel):
❸
channels = []
for i in range(3):
color_channel = image[:,:,i]
filtered_channel = convolve2d(color_channel, kernel, mode="same")
filtered_channel = jnp.clip(filtered_channel, 0.0, 1.0)
channels.append(filtered_channel)
final_image = jnp.stack(channels, axis=2)
return final_image
img = img_as_float32(imread('The_Cat.jpg'))
img_blur = color_convolution(img, kernel_blur)
❹
❺
Глава 6
250
Векторизация кода
plt.figure(figsize = (12,10))
plt.imshow(jnp.hstack((img_blur, img)))
❻
❻
Импорт всех необходимых пакетов.
Подготовка ядра для фильтра blur (размытие).
Функция для применения ядра фильтра к цветному изображению.
Загрузка тестового изображения из репозитория книги (вы можете заменить его
любым изображением по вашему выбору).
❺ Применение фильтра к изображению.
❻ Вывод результата.
❶
❷
❸
❹
Самой главной в этом примере была функция color_convolution(), применяющая ядро фильтра к изображению. Ядро фильтра –
это просто прямоугольная матрица весов, в нашем случае – матрица
с одинаковыми элементами. Это означает, что для создания нового
пиксела в получаемом (выходном) изображении каждый соседний
пиксел должен быть взят с тем же весом. Если вы забыли основной
принцип этого способа обработки изображений, то обратитесь к со
ответствующему разделу главы 3.
Результатом является отфильтрованное изображение, в нашем
случае – размытое изображение, показанное слева на рис. 6.7. Ис
пользовался тот же файл, что и в главе 3, доступный в репозитории
GitHub книги. Вы можете заменить этот файл изображения на любой
файл по вашему выбору, если не имеете возможности получить до
ступ к репозиторию.
Рис. 6.7 Отфильтрованное и исходное изображения
Давайте посмотрим, где можно применить автоматическую век
торизацию для такого способа обработки изображений. Если мыс
251
Варианты использования vmap() из реальной практики
лить на высоком уровне, то существует очевидный вариант исполь
зования пакетной обработки изображений – возможно, потребуется
применение одного фильтра к многочисленным изображениям или
применение различных фильтров к одному изображению. Разуме
ется, также можно объединить оба варианта. Это самый простой
способ применения vmap(), и мы уже обсудили его в подразделе 6.3.1.
Если заглянуть поглубже в код, то также можно заметить еще
одно место возможного применения vmap(): внутренний цикл по
цветовым каналам внутри функции color_convolution(). Этот цикл
отлично подходит для автоматической векторизации, так как уже
обрабатывает тензор со специально выделенным измерением, по
хожим на пакет, – в данном случае это измерение цветового кана
ла, – и обработка каждого канала выполняется независимо от любо
го другого канала.
Перепишем функцию color_convolution(), используя vmap() вмес
то цикла.
Листинг 6.23
Векторизация цикла внутри функции
import jax.numpy as jnp
from jax.scipy.signal import convolve2d
from skimage.io import imread
from skimage.util import img_as_float32
from matplotlib import pyplot as plt
kernel_blur = jnp.ones((5,5))
kernel_blur /= jnp.sum(kernel_blur)
def matrix_filter(channel, kernel):
filtered_channel = convolve2d(channel, kernel, mode="same")
filtered_channel = jnp.clip(filtered_channel, 0.0, 1.0)
return filtered_channel
color_convolution_vmap = jax.vmap(
matrix_filter, in_axes=(2, None), out_axes=2)
img = img_as_float32(imread('The_Cat.jpg'))
❶
❶
❷
❶
img_blur = color_convolution_vmap(img, kernel_blur)
plt.figure(figsize = (12,10))
plt.imshow(jnp.hstack((img_blur, img)))
❶ Замена цикла на vmap().
❷ Мы выделили цикл в отдельную функцию.
В приведенном выше примере мы заменили цикл внутри функ
ции color_convolution() на более простую функцию и выполнили
автоматическую векторизацию ее вызова. Код действительно стал
более простым, а если измерить производительность, то мы увидим,
что код к тому же стал быстрее выполняться.
Глава 6
252
Листинг 6.24
Векторизация кода
Тестирование внесенных изменений
%timeit color_convolution(
img, kernel_blur).block_until_ready()
❶
>>> 405 ms ± 2.68 ms per loop (mean ± std. dev. of 7 runs, 1 loop each)
%timeit color_convolution_vmap(
img, kernel_blur).block_until_ready()
❷
>>> 184 ms ± 12.8 ms per loop (mean ± std. dev. of 7 runs, 10 loops each)
❶ Исходная функция с циклом.
❷ Функция после автоматической векторизации с использованием vmap().
Здесь можно видеть, что функция стала выполняться быстрее более
чем в два раза, а кроме того, теперь ее легче читать. Функцию можно
сделать еще более быстрой, если применить трансформацию jit().
Листинг 6.25
Добавление трансформации jit()
color_convolution_jit = jax.jit(color_convolution)
color_convolution_vmap_jit = jax.jit(color_convolution_vmap)
color_convolution_jit(img, kernel_blur);
color_convolution_vmap_jit(img, kernel_blur);
%timeit color_convolution_jit(
img, kernel_blur).block_until_ready()
>>> 337 ms ± 1.33s per loop (mean ± std. dev. of 7 runs, 1 loop each)
%timeit color_convolution_vmap_jit(
img, kernel_blur).block_until_ready()
❶
❶
❷
❷
❸
❹
>>> 143 ms ± 588 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
❶
❷
❸
❹
JIT-компиляция функции.
Разогрев для выполнения реальной компиляции.
Исходная функция + JIT.
Функция после автоматической векторизации + JIT.
Здесь можно видеть, что JIT-компиляция исходной функции соз
дает версию, более медленную, чем функция с автоматической век
торизацией, даже без JIT-компиляции. При сочетании JIT-компиля
ции и автоматической векторизации полученная в итоге функция
работает еще быстрее.
Кроме того, можно использовать второй вызов vmap(), чтобы
функция работала с пакетами изображений и/или с фильтрами,
а итоговая комбинация будет выглядеть так:
jit(vmap(vmap(matrix_filter(...))))
Резюме
253
Этот пример и пример вычисления градиентов по каждой выбор
ке особенно подчеркивают преимущество комбинируемых транс
формаций, которые можно выполнять в JAX.
Резюме
Существует несколько способов векторизации кода: многократ
ные вызовы функции и объединение результатов вручную (в дей
ствительности эти способы не являются настоящей векториза
цией), векторизация, выполняемая вручную, и автоматическая
векторизация.
Автоматическая векторизация выполняет преобразование функ
ции, обрабатывающей один элемент, в функцию, одновременно
обрабатывающую множество элементов.
Автоматическая векторизация с использованием vmap() немного
уступает по скорости векторизации вручную, но делает код более
простым и понятным.
Можно управлять выбором осей для отображения (преобразова
ния), используя параметр in_axes.
Можно управлять схемой компоновки результатов с помощью па
раметра out_axes.
Использование декоратора помогает писать меньше кода, если
вас интересуют только функции, которые обрабатывают пакеты
данных.
Параметр axis_name позволяет использовать коллективные опе
рации, если в коде требуется обмен информацией между различ
ными элементами пакетов.
Пакетная обработка данных – это самый простой и очевидный
вариант использования для применения автоматической векто
ризации.
Трансформация vmap() также часто используется для пакетной
обработки моделей нейронных сетей. Но такой подход может ока
заться не самым простым, если используются какие-либо функции
для обработки статистических характеристик пакетов данных.
Можно с легкостью вычислять градиенты для отдельных выборок
(элементов), комбинируя трансформации grad() и vmap().
Векторизация циклов – еще один весьма полезный вариант ис
пользования трансформации vmap().
7
Распараллеливание
вычислений
Темы главы:
использование режима параллельного выполнения для
распараллеливания вычислений с применением pmap();
управление поведением pmap() с помощью параметров;
реализация процесса тренировки нейронной сети
с распараллеливанием по данным;
выполнение кода в конфигурации со многими хостами.
В этой главе продолжается изучение трансформаций JAX. Мы нач
нем с подробного рассмотрения распараллеливания, или одно
временного выполнения вычислений на нескольких устройствах
в параллельном режиме. Это особенно важно, если вы имеете дело
с тренировкой крупномасштабной нейронной сети, с имитацией по
годных условий или поведения океана, а также с любой другой за
дачей, в которой хотя бы некоторая часть вычислений не зависит от
прочих частей, следовательно, может выполняться в параллельном
режиме. В этом случае вы можете выполнять быстрее все вычисле
ние в целом, т. е. затрачивать меньше времени.
В JAX существует несколько механизмов распараллеливания вы
числений. Самым простым и понятным является pmap(), так как он
позволяет получить явное управление выполнением вычислений.
Распараллеливание вычислений
255
Механизм pmap() подробно рассматривается в этой главе. В следу
ющей главе описывается сегментирование тензоров – новый упро
щенный способ неявного распараллеливания, позволяющий ком
пилятору автоматически распределять функции по устройствам.
В приложении D дополнительно рассматриваются два эксперимен
тальных и отчасти неактуальных механизма: xmap() и pjit() для тех,
кто интересуется историей развития методик распараллеливания
в JAX или вынужден работать со старым кодом, использующим эти
методики.
Трансформация pmap(), или параллельное отображение, имеет
интерфейс, аналогичный vmap(), поэтому разумно начать изучение
именно с интерфейса. В pmap() используется методика распаралле
ливания, которую называют SPMD (single-program, multiple-data –
одна программа, множественные данные). По методике SPMD вы
выполняете одну программу на нескольких (многих) устройствах.
Основная идея заключается в разделении данных на фрагменты,
и каждое устройство одновременно обрабатывает собственный
фрагмент данных, используя один и тот же код. Таким образом,
можно обрабатывать одновременно больше данных, просто добав
ляя дополнительные устройства (но при этом масштабирование
обычно становится сублинейным из-за издержек на обмен инфор
мацией).
В этой главе мы начнем с определения условий обобщенной за
дачи, которую можно распараллелить по нескольким устройствам
(Cloud TPU с восемью ядрами TPU), а также рассмотрим, как рабо
тает распараллеливание при использовании pmap(). Далее следует
подробное описание параметров pmap() и применения коллектив
ных операций. В конце главы мы применим на практике все полу
ченные знания и приобретенные навыки для решения реальной за
дачи тренировки нейронной сети. Мы разработаем программу для
тренировки с распараллеливанием по данным для нейронной сети
классификации изображений из главы 2. Хотя пример MNIST может
показаться слишком простым и тривиальным, это неплохой пример
модели, и вы можете с легкостью наблюдать за изменениями в коде,
тогда как при этом другие несущественные части (набор данных
и размеры тензоров) остаются теми же самыми. Я преднамеренно
выбрал использование уже знакомого нам примера классификации
изображений MNIST в этой главе и в следующей, чтобы особо вы
делить изменения в коде и позволить вам с легкостью сравнивать
различные методики.
После этого мы узнаем, как выполнить код в крупных средах со
многими хостами, например TPU Pods. Иногда пользовательские
функции (или нейронные сети) становятся настолько большими, что
одно устройство GPU/TPU не справляется с их обработкой, поэтому
требуется выполнение вычислений в кластере. Такая ситуация часто
Глава 7
256
Распараллеливание вычислений
возникает при тренировке и логическом выводе больших языковых
моделей LLM (large language model). Современные LLM, такие как
GPT-4, PaLM 2, Gemini, самые свежие версии LLaMa 2, Falcon и мно
гие другие требуют использования систем, содержащих многочис
ленные GPU. JAX позволяет не только распределять (или разделять,
или сегментировать) обработку данных по различным компьютерам
(называемую распараллеливанием по данным (data parallelism)), но
также разделять крупномасштабные вычисления на части, выпол
няемые различными компьютерами (так называемое распаралле
ливание модели (model parallelism)). Можно использовать pmap()
в конфигурации со многими хостами для распределения вручную
конкретного вычисления в кластере (хотя существуют и другие спо
собы сделать это; см. главу 8 и приложение D). Если в вашем распо
ряжении нет кластера GPU или TPU и не требуется тренировка круп
номасштабной модели, то вы можете пропустить эту часть.
7.1
Распараллеливание вычислений
с помощью pmap()
Начнем с простой задачи, в которой имеется несколько устройств
аппаратного ускорения (например, TPU или GPU) и функция, работу
которой можно распределить по этим устройствам в параллельном
режиме.
Процесс описывается следующим образом:
подготовка системы с несколькими устройствами;
данных, предназначенных для обработки; они
представлены как некоторые (возможно, большие) тензоры;
3 принятие решения о том, как входной тензор(ы) можно разде
лить на независимые фрагменты для отдельной обработки;
4 изменение формы входного тензора так, чтобы он имел допол
нительную ось с различными фрагментами, размещенными по
этой оси;
5 трансформация функции, используемой для обработки данных,
с помощью pmap() (и дополнительно с применением vmap(),
если функция не векторизована, но каждый фрагмент пред
ставляет пакет элементов данных);
6 применение трансформированной функции к тензору (тензо
рам) с измененной формой;
7 изменение формы полученного в результате тензора для ис
ключения введенного дополнительного измерения.
1
2 подготовка
Мы будем использовать Cloud TPU с восемью вычислительными
ядрами и функцию вычисления скалярного произведения двух век
торов, с которой мы работали в предыдущей главе.
257
Распараллеливание вычислений с помощью pmap()
7.1.1
Установка начальных условий задачи
Начнем с настройки системы с несколькими устройствами. Исполь
зуем среду Cloud TPU, предоставляющую восемь ядер TPU в обла
ке. Это тот же тип настройки, который мы выполняли в подразде
ле 3.2.5, когда перемещали обработку изображений на TPU.
Для использования Cloud TPU выполните процедуру, описанную
в главе 3 или в приложении C.
После создания рабочей среды Cloud TPU и установления соеди
нения Colab (или вашего локального блокнота Jupyter) с этой средой
выполните код, приведенный в листинге 7.1, чтобы проверить до
ступность устройств TPU.
Листинг 7.1
Настройка Cloud TPU в Google Colab
from jax.lib import xla_bridge
print(xla_bridge.get_backend().platform)
❶
>>> tpu
❶
import jax
jax.local_devices()
>>> [TpuDevice(id=0,
➥core_on_chip=0),
>>> TpuDevice(id=1,
➥core_on_chip=1),
>>> TpuDevice(id=2,
➥core_on_chip=0),
>>> TpuDevice(id=3,
➥core_on_chip=1),
>>> TpuDevice(id=4,
➥core_on_chip=0),
>>> TpuDevice(id=5,
➥core_on_chip=1),
>>> TpuDevice(id=6,
➥core_on_chip=0),
>>> TpuDevice(id=7,
➥core_on_chip=1)]
❶
❷
process_index=0, coords=(0,0,0),
process_index=0, coords=(0,0,0),
process_index=0, coords=(1,0,0),
process_index=0, coords=(1,0,0),
process_index=0, coords=(0,1,0),
process_index=0, coords=(0,1,0),
process_index=0, coords=(1,1,0),
process_index=0, coords=(1,1,0),
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❸
❶ Проверка с целью определить, какие внутренние компоненты JAX используются.
❷ Полный импорт JAX.
❸ Доступны восемь устройств TPU.
Здесь можно видеть, что в текущем процессе Python с process_index=0 нам предоставлен доступ к восьми ядрам TPU, размещенным
на четырех микросхемах (кортеж coords определяет индекс микро
схемы в двоичном формате, а значение core_on_chip означает номер
ядра на микросхеме; более подробную информацию об использова
Глава 7
258
Распараллеливание вычислений
нии TPU можно узнать здесь: https://moocaholic.medium.com/hard
ware-for-deep-learning-part-4-asic-96a542fe6a81#e04d).
Если такая система недоступна, то можно эмулировать конфигу
рацию с произвольным количеством устройств, используя специ
альный флаг XLA --xla_force_host_platform_device_count. Необхо
димо установить этот флаг перед импортом JAX.
Листинг 7.2
Эмуляция системы с несколькими устройствами на CPU
import os
os.environ['XLA_FLAGS'] = '--xla_force_host_platform_device_count=8' ❶
import jax
jax.devices("cpu")
❷
WARNING:jax._src.lib.xla_bridge:No GPU/TPU found, falling back to CPU. (Set
TF_CPP_MIN_LOG_LEVEL=0 and rerun for more info.)
# ПРЕДУПРЕЖДЕНИЕ: jax._src.lib.xla_bridge: не найдены устройства GPU/TPU,
# возврат на CPU. Установите TF_CPP_MIN_LOG_LEVEL=0 и выполните повторный
# запуск для получения более подробной информации.
❸
[CpuDevice(id=0),
CpuDevice(id=1),
❸
CpuDevice(id=2),
❸
CpuDevice(id=3),
❸
CpuDevice(id=4),
CpuDevice(id=5),
CpuDevice(id=6),
CpuDevice(id=7)]
❶ Установка переменной среды с указанием требуемого количества устройств пе-
ред импортом JAX.
❷ Импорт JAX.
❸ Доступны восемь устройств CPU.
Использование нескольких устройств CPU помогает прототипи
ровать, отлаживать и тестировать код для многих устройств перед
выполнением его в дорогостоящей системе TPU и GPU. Даже при ис
пользовании Google Colab такой подход может способствовать уско
ренному созданию прототипа, поскольку среда времени выполне
ния CPU быстрее перезапускается.
Наличие многоядерного CPU (в наше время это обычное явле
ние) – простой способ распараллеливания работы по нескольким
ядрам. Этот вариант также работает, если количество устройств CPU,
указанное в флаге XLA, превышает число CPU в используемой систе
ме. Такой поход работает, даже если доступно единственное физиче
ское устройство. В этом случае вы не получите увеличения скорости,
но извлечете пользу при тестировании семантики параллельной реа
лизации.
259
Распараллеливание вычислений с помощью pmap()
Если вам повезло и вы являетесь владельцем компьютера с не
сколькими GPU, то можете использовать их для распараллеливания.
ПРИМЕЧАНИЕ Трансформация pmap() требует, чтобы все
применяемые устройства были одинаковыми. Это становится
проблемой для владельцев компьютеров с несколькими раз
личными GPU, так как невозможно использовать pmap() для
распараллеливания вычислений по различным моделям GPU.
Воспользуемся функцией, с которой мы работали в предыдущей
главе для применения vmap(), – функцией вычисления скалярного
произведения двух векторов.
Листинг 7.3 Функция вычисления скалярного произведения
двух векторов
import jax.numpy as jnp
def dot(v1, v2):
return jnp.vdot(v1, v2)
❶
dot(jnp.array([1., 1., 1.]), jnp.array([1., 2., -1]))
❷
>>> Array(2., dtype=float32)
❶ Использование функции vdot().
❷ Вычисление скалярного произведения двух векторов.
Если вместо двух векторов, как в предыдущей главе, имеются два
списка векторов и необходимо вычислять скалярные произведения
соответствующих элементов этих списков, то на этот раз списки мо
гут иметь значительно больший размер, чтобы эмулировать вариант,
в котором один акселератор не способен эффективно выполнить все
операции умножения в параллельном режиме. Здесь может помочь
разделение работы между несколькими акселераторами.
Листинг 7.4
Генерация двух длинных списков векторов
from jax import random
rng_key = random.PRNGKey(42)
vs = random.normal(rng_key, shape=(20_000_000,3))
v1s = vs[:10_000_000,:]
v2s = vs[10_000_000:,:]
v1s.shape, v2s.shape
❶
❷
❸
❸
>>> ((10000000, 3), (10000000, 3))
❶ Создание ключа генератора случайных чисел (более подробно об этом – в гла-
ве 9).
Глава 7
260
Распараллеливание вычислений
❷ Генерация двумерного массива случайных чисел.
❸ Разделение полученного массива на две части: первые 10 элементов размещают-
ся в первом списке, следующие 10 элементов – во втором.
Теперь у нас есть все требуемые ингредиенты: система с несколь
кими устройствами, данные и функция, которую необходимо при
менить. Наша цель – распределить вычисление всех запланирован
ных скалярных произведений по доступным устройствам.
7.1.2
Использование pmap (почти) так же, как vmap
В главе 3 мы не обсуждали такую тему: почему при наличии восьми
устройств TPU JAX по умолчанию выполняет вычисления только на
одном устройстве. Если необходимо задействовать все доступные
устройства, то самым простым способом выполнения вычислений
является отображение функции и указание каждому устройству вы
полнять один индекс из этого отображения. Простой способ рас
пределения вычислений по различным устройствам – применение
механизма параллельного отображения pmap().
Как и vmap(), pmap() отображает функцию вдоль осей массива. Раз
личие заключается в том, что vmap() векторизует функцию, добавляя
измерение пакетов в каждую простейшую операцию функции (сна
чала в представлении Jaxpr, которое после компиляции транслирует
ся в соответствующие операции HLO). Вместо этого pmap() компили
рует функцию с помощью XLA (поэтому отдельная трансформация
jit() не нужна), создавая реплики функции для устройств, и выпол
няет каждую реплику на отдельном устройстве в параллельном ре
жиме.
Модель «одна программа, множественные данные»
Целью pmap() является выражение парадигмы «одна программа, множественные данные» SPMD (single-program, multiple-data). SPMD представляет собой ответвление классификации архитектур компьютеров,
предложенной Майклом Дж. Флинном (Michael J. Flynn) в 1966 г.
Флинн определил четыре исходные единицы классификации на основе количества параллельных потоков инструкций и данных, доступных
в конкретной архитектуре:
SISD (single instruction single data – одна инструкция, один поток
данных): это последовательный компьютер без распараллеливания
по данным и по инструкциям;
SIMD (single instruction multiple data – одна инструкция, множественные данные): вариант, когда одна инструкция одновременно применяется к нескольким потокам данных. Современные процессоры
имеют специальные векторные инструкции для реализации этого
Распараллеливание вычислений с помощью pmap()
261
типа распараллеливания, например MMX, SSE и AVX в семействе x86.
Позже, в 1972 г., Флинн определил три дополнительные подкатегории
SIMD для массива, конвейерных и векторных процессоров;
MISD (multiple instructions single data – множественные инструкции,
один поток данных): множественные инструкции работают с одним
потоком данных. Эта редко встречающаяся архитектура в основном
используется для устойчивых к критическим ошибкам приложений,
таких как компьютер управления полетом;
MIMD (multiple instructions multiple data – множественные инструкции, множественные данные): несколько (много) процессоров выполняют различные инструкции для различных данных.
Более подробное описание можно найти в этом видеоролике: https://
www.youtube.com/watch?v=KVOc6369-Lo.
Категория MIMD иногда делится на следующие подкатегории, не являющиеся частью классификации Флинна:
SPMD (single program multiple data – одна программа, множественные данные): при этом типе распараллеливания одна и та же программа (например, некоторая функция нейронной сети) выполняется
одновременно на нескольких устройствах (таких как GPU или TPU),
но входные данные для каждого экземпляра работающей программы
могут быть различными (например, это могут быть сегменты массива
или разные пакеты данных), т. е. здесь мы имеем множественные потоки данных. SPMD работает на более высоком уровне абстракции по
сравнению с SIMD, поэтому они не являются взаимно исключающими.
Программа SPMD может использовать SIMD-инструкции на каждом
устройстве, на котором она работает, если устройство предоставляет
такие возможности;
MPMD (multiple programs multiple data – множественные программы,
множественные данные): по крайней мере две различные программы
работают с множественными потоками данных.
Рассмотрим примеры поведения pmap() по сравнению с пове
дением vmap(). В листинге 7.5 мы пытаемся создать две версии ис
ходной функции. Одна автоматически векторизована и скомпи
лирована, другая представляет собой распараллелеленную версию
исходной функции, полученную простой заменой vmap() на pmap()
без явного вызова jit().
Листинг 7.5 Применение vmap и pmap к учебным массивам
dot_batched = jax.jit(jax.vmap(dot))
x_vmap = dot_batched(v1s, v2s)
x_vmap.shape
❶
❷
❸
Глава 7
262
Распараллеливание вычислений
>>> (10000000,)
dot_parallel = jax.pmap(dot)
x_pmap = dot_parallel(v1s,v2s)
❸
❹
❺
>>> ...
>>> ValueError: compiling computation that requires 10000000 logical devices,
but only 8 XLA devices are available (num_replicas=10000000)
❻
# ValueError: компиляция вычисления, которое требует 10000000 логических
# устройств, но доступно только 8 XLA-устройств (num_replicas=10000000)
❶
❷
❸
❹
❺
❻
Создание автоматически векторизованной и скомпилированной функции.
Вызов автоматически векторизованной и скомпилированной функции.
Проверка формы результата.
Создание распараллеленной функции.
Вызов распараллеленной функции.
Ошибка параллельного отображения.
В этом листинге наблюдается несколько интересных моментов:
автоматически векторизованная и скомпилированная версия
функции dot_batched() успешно применяется к предложенным
массивам;
в распараллеленной версии функции возникает ошибка, так
как она предполагает, что каждый элемент отображаемой оси
связывается с отдельным устройством, а у нас нет достаточного
количества устройств для такого отображения, поэтому простая
замена vmap() на pmap() не работает.
Важное различие между vmap() и pmap() заключается в том, что
размер отображаемой оси обязательно должен быть меньше или
равен количеству доступных локальных XLA-устройств, которое
возвращает функция jax.local_device_count(). У нас имеется толь
ко 8 устройств, тогда как размер отображаемой оси равен 10 мил
лионам.
Здесь выполняются шаги 3 и 4 из описания процесса, приведенно
го в начале главы: мы решаем, как входные тензоры должны разде
ляться на независимые фрагменты для отдельной обработки. Мы из
меняем форму тензоров соответствующим образом для получения
дополнительной оси с различными фрагментами, размещенными
по ней.
Разделим массив, содержащий векторы, на восемь фрагментов,
но не будем делить само измерение векторов (содержащее три эле
мента). Необходимо так перекомпоновать массивы, чтобы распа
раллеливаемое измерение соответствовало количеству устройств.
Предполагается обеспечение распараллеливания для обработки по
этой оси и передача каждой из восьми сформированных групп на
отдельное устройство.
263
Распараллеливание вычислений с помощью pmap()
Листинг 7.6
Реструктуризация массивов
v1s.shape
>>> (10000000, 3)
v1sp = v1s.reshape((8, v1s.shape[0]//8, v1s.shape[1]))
v2sp = v2s.reshape((8, v2s.shape[0]//8, v2s.shape[1]))
v1sp.shape
>>> (8, 1250000, 3)
❶
❷
❷
❸
❶ Невозможно применить pmap по первой (основной) оси, так как не имеется до-
статочное количество устройств.
❷ Изменение формы массивов: разделение на восемь фрагментов.
❸ Теперь массивы содержат восемь фрагментов, в каждом из которых содержится
1,25 миллиона элементов.
В приведенном выше примере массивы были реструктурированы
так, чтобы первая (основная) ось (которую необходимо использо
вать как отображаемую ось) содержала восемь элементов, что соот
ветствует количеству имеющихся устройств. Обратите внимание:
здесь применялось целочисленное деление (операция //), потому
что обычное деление возвращает число с плавающей точкой, кото
рое невозможно использовать в качестве формы.
Во многих случаях требуемое измерение может оказаться неде
лимым на количество доступных устройств. В подобной ситуации
можно воспользоваться дополняющим выравниванием (padding),
т. е. дополнением массива некоторым фиктивным значением (на
пример, нулем), чтобы сделать измерение делимым. После обработ
ки нужно удалить фиктивные значения (и результаты вычислений
по ним). Разумеется, такой подход сработает только в том случае,
если фиктивные значения не влияют на результат вычисления. Ина
че потребуется еще и корректировка самого вычисления.
Мы почти готовы к переходу к шагам 5 и 6 из описания процесса,
приведенного в начале главы, т. е. к трансформации функции с при
менением pmap() и применению трансформированной функции
к тензору (тензорам) с измененной формой.
Применим функцию, трансформированную с помощью pmap(),
к реструктурированному массиву.
Листинг 7.7
Применение pmap к реструктурированным массивам
x_pmap = dot_parallel(v1sp,v2sp)
x_pmap.shape
>>> (8,)
❶
❷
❶ Теперь распараллеливание отображения проходит успешно.
❷ Получена не та форма, которая ожидалась.
Глава 7
264
Распараллеливание вычислений
Теперь выполнение функции завершается успешно. Но при взгля
де на форму результата мы понимаем, что это не то, чего мы ожи
дали. Предполагалось получение огромного количества скалярных
произведений в каждой группе, а именно 1,25 миллиона чисел. То
есть мы должны были увидеть форму (8, 1250000), а не (8, ).
Это та же проблема векторизации, которая была описана в пре
дыдущей главе. Функция dot() предназначена для работы с одним
элементом, но в ней не возникает ошибка при передаче нескольких
элементов, и она просто вычисляет скалярное произведение тензо
ров самого высшего ранга (а это совсем не то, что нам нужно).
После разделения исходного массива на восемь фрагментов каж
дый фрагмент становится массивом меньшего размера, поэтому
нам нужна векторизованная функция для корректной обработки
сформированных входных массивов. Решение уже известно: можно
воспользоваться vmap(). Итак, сначала выполняется векторизация
функции с помощью vmap(), затем применяется pmap() для распа
раллеливания.
Листинг 7.8
Добавление vmap в код
dot_parallel = jax.pmap(jax.vmap(dot))
x_pmap = dot_parallel(v1sp,v2sp)
❶
x_pmap.shape
>>> (8, 1250000)
type(x_pmap)
>>> jaxlib.xla_extension.ArrayImpl
❷
❸
❶ Распараллеливание автоматически векторизованной функции.
❷ Теперь полученная форма корректна.
❸ Типом результата является Array.
В приведенном выше примере мы сначала создали автоматиче
ски векторизованную версию функции, затем распараллелили ее по
нескольким устройствам. Теперь все работает правильно, и мы полу
чили результат в ожидаемой корректной форме.
Если вы используете более старую версию JAX, то, возможно, обра
тили внимание на тип ShardedDeviceArray. Ранее существовала спе
циальная версия типа DeviceArray, которая логически выглядела как
единый массив, хотя он физически распределялся по устройствам, ис
пользуемым для вычисления. В новых версиях JAX это тип jaxlib.xla_
extension.ArrayImpl, являющийся хорошо знакомым нам типом Array.
Последующие вызовы pmap() могут выполняться без перемеще
ния данных в зависимости от конкретного вычисления. Если для
этих данных вызывается функция, отличная от pmap(), то незаметно
Распараллеливание вычислений с помощью pmap()
265
для пользователя происходит обмен информацией для сбора полу
чаемых значений на одном устройстве.
Осталось решить одну небольшую задачу – шаг 7 общего процес
са: необходимо реструктурировать полученный в итоге массив для
устранения артефактов распараллеливания, а затем удалить допол
нительную ось, если не предполагаются какие-либо действия с полу
ченными результатами в параллельном режиме.
Листинг 7.9
Исключение дополнительных измерений
x_pmap = x_pmap.reshape((x_pmap.shape[0]*x_pmap.shape[1]))
x_pmap.shape
❶
>>> (10000000,)
jax.numpy.all(x_pmap == x_vmap)
>>> Array(True, dtype=bool)
❷
❷
❶ Удаление отображаемой оси.
❷ Проверка, чтобы убедиться в том, что получен тот же результат, что и при исполь-
зовании vmap().
В рассматриваемом здесь примере мы сделали полученный мас
сив плоским и удалили отображаемую ось, поскольку она больше не
нужна. Сравнение текущего результата с результатом, полученным
при использовании vmap(), показывает, что они одинаковы.
Подведем общий итог проделанной работы:
1 изначально
в нашем распоряжении находился массив (или не
сколько массивов), который требовалось обработать конкрет
ной функцией;
2 мы разделили исходный массив на группы, количество которых
соответствовало числу доступных аппаратных устройств;
3 поскольку каждое устройство должно обрабатывать пакет дан
ных, мы подготовили пакетную версию функции с помощью
vmap();
4 функция была скомпилирована и реплицирована на доступные
устройства с помощью pmap();
5 трансформация pmap() запустила вычисления на каждом
устройстве с соответствующим фрагментом данных, что позво
лило получить итоговые массивы на каждом устройстве;
6 мы исключили разделение на фрагменты, удалили измерение,
по которому происходило отображение и обработка. После это
го выходной массив не содержит артефактов процесса распа
раллеливания.
Как можно видеть, трансформация pmap() в целом аналогична
трансформации vmap(), но они выполняют различные действия, по
этому (как правило) невозможно просто заменить одну на другую.
Глава 7
266
Распараллеливание вычислений
Тем не менее трансформации pmap() и vmap() иногда могут ока
зываться взаимозаменяемыми и равнозначными с точки зрения
внешнего наблюдателя, например если имеется небольшое коли
чество обрабатываемых элементов, не превышающее число доступ
ных устройств, скажем в нашем примере нужно обработать не более
восьми элементов.
Листинг 7.10 Вариант, в котором vmap() и pmap() выглядят (и работают)
одинаково
vs = random.normal(rng_key, shape=(16,3))
v1s = vs[:8,:]
v2s = vs[8:,:]
jax.vmap(dot)(v1s,v2s)
❶
❶
❶
❷
>>> Array([ 0.51048726, -0.7174605 , -0.20105815, -0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ], dtype=float32)
jax.pmap(dot)(v1s,v2s)
❸
>>> Array([ 0.51048726, -0.7174605 , -0.20105815, -0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ], dtype=float32)
❶ Генерация небольших массивов с восемью элементами (равно количеству доступных
устройств).
❷ Автоматическая векторизация функции.
❸ Распараллеливание функции.
Здесь vmap() и pmap() использовались одинаково, и полученные
результаты одинаковы. В старых версиях JAX можно заметить един
ственное различие в типе полученного массива: DeviceArray или
ShardedDeviceArray.
Еще одним не столь явно выраженным различием является ско
рость, поскольку для таких маленьких функций и массивов не имеет
смысла распараллеливать вычисления. Если все можно вычислить
на одном устройстве, то распараллеливание по нескольким устрой
ствам только добавит накладные расходы на обмен информацией
и в целом замедлит вычисление. Так как осмысленная оценка про
изводительности должна сравнивать JIT-скомпилированные версии
кода (иначе код может оказаться неоптимальным), добавим явные
вызовы jit() в код, приведенный в листинге 7.11.
Листинг 7.11 Измерение различий (в производительности)
между vmap() и pmap()
dot_v = jax.jit(jax.vmap(dot))
x = dot_v(v1s,v2s)
dot_pjo = jax.jit(jax.pmap(dot))
❶
❶
❷
267
Распараллеливание вычислений с помощью pmap()
❷
x = dot_pjo(v1s,v2s)
>>>… UserWarning: The jitted function dot includes
a pmap. Using jit-of-pmap can lead to inefficient
data movement, as the outer jit does not preserve
sharded data representations and instead collects
input and output arrays onto a single device. Consider
removing the outer jit unless you know what you're doing.
# UserWarning: JIT-скомпилированная функция dot включает
# pmap. Использование jit-of-pmap может привести к
# неэффективному перемещению данных, так как внешний
# jit не сохраняет представления сегментированных данных,
# а вместо этого собирает входные и выходные массивы
# на одном устройстве. Рекомендуется удалить внешний
# jit, если вы точно не знаете, что именно хотите сделать.
See https://github.com/google/jax/issues/2926.
warnings.warn(
dot_pji = jax.pmap(jax.jit(dot))
x = dot_pji(v1s,v2s)
dot_p = jax.pmap(dot)
x = dot_p(v1s,v2s)
%timeit dot_v(v1s,v2s).block_until_ready()
❸
❸
❹
❹
>>> 122 µs ± 5.53 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit dot_pjo(v1s,v2s).block_until_ready()
>>> 2.08 ms ± 30.1 µs per loop (mean ± std. dev. of 7 runs, 10 loops each)
%timeit dot_pji(v1s,v2s).block_until_ready()
>>> 1.51 ms ± 43 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit dot_p(v1s,v2s).block_until_ready()
>>> 1.54 ms ± 51 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶
❷
❸
❹
Компиляция и разогрев версии vmap().
Компиляция и разогрев версии pmap() с внешним jit().
Компиляция и разогрев версии pmap() с внутренним jit().
Компиляция и разогрев версии pmap() без явного вызова jit().
В приведенном выше примере были созданы следующие скомпи
лированные версии:
функция, трансформированная с помощью vmap();
внешний вызов jit() поверх функции, трансформированной
с помощью pmap();
трансформированная с помощью pmap() функция с внутренней
трансформацией jit();
трансформированная с помощью pmap() функция без явного
вызова jit().
Глава 7
268
Распараллеливание вычислений
По итогам сравнения было сделано несколько интересных наблю
дений:
трансформированная с помощью vmap() версия является самой
быстрой;
трансформированная с помощью pmap() версия с явной внут
ренней трансформацией jit() и без нее выдает одинаковые
результаты;
трансформированная с помощью pmap() версия с явной внеш
ней трансформацией jit() является самой медленной, и JAX
выдает информативное предупреждающее сообщение о том,
что такая комбинация может ухудшить производительность,
потому что внешний вызов jit() не сохраняет представления
сегментированных данных, а вместо этого собирает входные
и выходные массивы на одном устройстве.
После получения общего представления о том, что делает pmap(),
можно более подробно рассмотреть параметры этой трансфор
мации.
7.2
Управление поведением pmap()
Поскольку интерфейс pmap() очень похож на интерфейс vmap(), мож
но предположить, что существуют аналогичные способы управле
ния отображаемыми осями входных и выходных тензоров. Тензоры
могут быть организованы по-разному, и если измерение, которое
необходимо отобразить и обработать, не является первым (основ
ным), то потребуется сообщить pmap() об особенностях организации
входных данных. Также имеется возможность использования имен
осей и выполнения коллективных операций, если необходим обмен
информацией между устройствами. Для этого используются соот
ветствующие параметры pmap().
7.2.1
Управление отображением осей входных
и выходных данных
Как и при работе с vmap(), можно управлять отображением осей
входных и выходных тензоров с помощью специальных параметров.
Варианты точно такие же, как для управления поведением vmap()
(см. раздел 6.2). Например, массивы могут быть скомпонованы поразному, и измерение, требующее распараллеливания, не является
первым. Или в качестве входных параметров функции используются
более сложные структуры, например словарь dict.
269
Управление поведением pmap()
Использование параметра in_axes
По умолчанию pmap() предполагает, что все входные данные для
функции являются отображаемыми и обрабатываемыми, и, как
и для vmap(), можно использовать параметр in_axes, чтобы опреде
лить, какую ось позиционных аргументов нужно отображать. Аргу
менты, помеченные значением None, становятся общедоступными
(полностью скопированными на каждое устройство), тогда как цело
численные значения определяют, какие оси позиционных аргумен
тов необходимо отобразить и обработать.
В старых версиях JAX поддерживалось отображение и обработка
только первых (основных) осей (с индексом 0), но сейчас таких огра
ничений не существует.
Мы можем воспользоваться всеми примерами из предыдущей
главы, где мы работали с vmap(). Единственным изменением будет
использование pmap(). При этом существует определенное ограни
чение, так как примеры будут работать только для вариантов, в ко
торых размер отображаемой оси достаточно мал, для того чтобы
соответствовать весьма ограниченному количеству доступных ап
паратных устройств (например, восьми). Для более реалистичных
вариантов может потребоваться комбинирование vmap() и pmap(),
чтобы подготовить пакеты для каждого устройства, как это было
сделано в предыдущем разделе.
Начнем с простого примера применения параметра in_axes со
значением, установленным по умолчанию.
Листинг 7.12
Явное использование параметра in_axes
vs = random.normal(rng_key, shape=(16,3))
v1s = vs[:8,:]
v2s = vs[8:,:]
def dot(v1, v2):
return jnp.vdot(v1, v2)
dot_pmapped = jax.pmap(dot, in_axes=(0,0))
dot_pmapped(v1s, v2s)
>>> Array([ 0.51048726, -0.7174605 , -0.20105815,
➥-0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ],
➥ dtype=float32)
❶
❷
❸
❹
❶
❶
❶
❷
❸
❹
❹
❹
❹
Генерация небольшого массива для демонстрационных целей.
Та же функция dot, которая использовалась ранее.
Создание распараллеленной функции с явно передаваемым параметром in_axes.
Полученный массив содержит скалярные произведения, вычисленные в параллельном режиме.
Глава 7
270
Распараллеливание вычислений
В приведенном выше примере создан небольшой массив с пер
вым (основным) измерением размером 8, равным количеству име
ющихся аппаратных устройств. Если у вас другая конфигурация, то,
возможно, потребуется изменение размера массивов для соответ
ствия количеству устройств.
Поскольку поведением по умолчанию pmap() является использо
вание первых осей для отображения всех параметров функции, при
мер из листинга 7.10 может использовать явно передаваемое значе
ние параметра in_axes=(0,0).
В листинге 7.13 один из двух входных массивов транспонирует
ся, поэтому отображение должно быть выполнено по его второй оси
(с индексом 1).
Листинг 7.13
Использование параметра in_axes для неосновных осей
v1s.T.shape, v2s.shape
>>> ((3, 8), (8, 3))
jax.pmap(dot, in_axes=(1,0))(v1s.T, v2s)
❶
❶
>>> Array([ 0.51048726, -0.7174605 , -0.20105815, -0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ], dtype=float32)
v1s.T.shape, v2s.T.shape
>>> ((3, 8), (3, 8))
jax.pmap(dot, in_axes=(1,1))(v1s.T, v2s.T)
❷
❷
>>> Array([ 0.51048726, -0.7174605 , -0.20105815, -0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ], dtype=float32)
❶ Отображаемые измерения имеют индексы 1 и 0 для первого и второго массивов соответ-
ственно.
❷ Теперь отображаемые измерения для обоих массивов имеют индекс 1.
Если имеется параметр, который нужно скопировать (или раз
множить) на каждое аппаратное устройство как есть, без разделе
ния, то можно воспользоваться значением None в соответствующей
позиции параметра in_axes. В листинге 7.14 мы используем уже
знакомую функцию для вычисления масштабированного скалярно
го произведения из предыдущей главы, где последний ее параметр
koeff не должен разделяться.
Листинг 7.14 Использование параметра in_axes для множественной передачи
параметра функции без изменений
def scaled_dot(v1, v2, koeff):
return koeff*jnp.vdot(v1, v2)
❶
271
Управление поведением pmap()
v1s_ = v1s
v2s_ = v2s.T
k = 1.0
❷
❸
❹
v1s_.shape, v2s_.shape
>>> ((8, 3), (3, 8))
scaled_dot_pmapped = jax.pmap(scaled_dot, in_axes=(0,1,None))
❺
scaled_dot_pmapped(v1s_, v2s_, k)
>>> Array([ 0.51048726, -0.7174605 , -0.20105815, -0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ], dtype=float32)
❶
❷
❸
❹
❺
❻
❻
❻
Последний параметр koeff необходимо скопировать на все устройства.
Этот параметр будет отображен по оси номер 0.
Этот параметр будет отображен по оси номер 1.
Этот параметр будет передан без изменений (как есть) на все устройства.
Передача параметра in_axes.
Функция работает, как предполагалось.
В приведенном выше примере последний параметр функции –
коэффициент масштабирования реплицируется на все устройства.
Как и при работе с vmap(), также можно использовать более слож
ные структуры в качестве параметров функции и определять соот
ветствующее значение in_axes. В примере, показанном в листин
ге 7.15, используется dict – словарь языка Python.
Листинг 7.15 Использование параметра in_axes для определения контейнера
Python
def scaled_dot(data, koeff):
return koeff*jnp.vdot(data['a'], data['b'])
❶
scaled_dot_pmapped = jax.pmap(scaled_dot, in_axes=({'a':0,'b':1},None))
❷
scaled_dot_pmapped({'a':v1s_, 'b': v2s_}, k)
>>> Array([ 0.51048726, -0.7174605 , -0.20105815,
➥-0.26437205, -1.3696793 ,
>>>
2.744793 , 1.7936493 , -1.1743435 ],
➥dtype=float32)
❶ Теперь функция принимает словарь и скалярное значение.
❷ Пометка осей для словаря и скалярного параметра.
❸ Функция работает, как предполагалось.
❸
❸
❸
❸
Использование параметра out_axes
Также можно управлять схемой компоновки выходного тензора
с помощью параметра out_axes. В листинге 7.16 требуется масшта
Глава 7
272
Распараллеливание вычислений
бировать входные векторы функции с заданным коэффициентом,
а затем по некоторой причине вернуть результат в транспониро
ванном виде с отображаемой осью при выводе с индексом 1, а не 0.
Листинг 7.16
Использование параметра out_axes
def scale(v, koeff):
return koeff*v
scale_pmapped = jax.pmap(scale,
in_axes=(0,None),
out_axes=(1))
res = scale_pmapped(v1s, 2.0)
v1s.shape, res.shape
>>> ((8, 3), (3, 8))
❶
❷
❸
❹
❶
❷
❸
❹
❹
Простая функция масштабирования вектора.
Определение пакетных измерений для входных данных.
Определение пакетных измерений для выводимого результата.
Теперь выводится транспонированный и масштабированный входной вектор.
В приведенном выше примере отображаемым входным измере
нием является измерение с индексом 0, охватывающее все векторы,
которые необходимо обработать, тогда как измерение с индексом 1
содержит компоненты отдельного вектора. Для вывода необходим
транспонированный тензор с компонентами отдельного вектора,
размещенный по измерению с индексом 0, а индексы масштабиро
ванного вектора должны располагаться по измерению с индексом 1.
Существует особый случай, когда для параметра out_axes уста
новлено значение None. Такая установка автоматически возвраща
ет значение на первое устройство и должна использоваться, толь
ко если вы уверены в том, что все значения одинаковы на каждом
устройстве.
Пример большого массива
Ранее демонстрировались простые учебные примеры, но в реальной
практике вам, вероятнее всего, придется работать с массивами го
раздо большего размера, чем количество используемых устройств.
В предыдущем разделе потребовалось скомбинировать pmap()
и vmap(), чтобы распределить работу по нескольким акселераторам
с помощью pmap() и использовать пакеты на каждом акселераторе
посредством vmap().
Создадим большие массивы для варианта вычисления скалярно
го произведения и обработаем их в параллельном режиме в той же
конфигурации Cloud TPU (или в любой другой конфигурации, до
Управление поведением pmap()
273
ступной вам для использования). В этот раз мы используем транспо
нированные версии больших массивов (в которых строки и столбцы
меняются местами), взятых из начала главы.
Теперь рассматриваемый пример становится большим и слож
ным. Необходимо внимательно следить за тем, какие измерения
имеются в наличии и как и когда использовать их. Рассмотрим все
подробности на схеме, показанной на рис. 7.1.
Рис. 7.1 Схема обработки данных для примера большого транспонированного массива
Сначала мы создаем два больших массива, по соответствующим
элементам которых необходимо вычислить скалярные произведе
ния. Массивы являются транспонированными версиями массивов
из начала главы и имеют форму (3, 10 000 000). По первой оси (с ин
дексом 0) размещаются элементы каждого отдельного вектора (по
три элемента в каждом векторе). Вторая ось (с индексом 1) содержит
сами векторы (по 10 миллионов в каждом массиве), а индекс по этой
оси соответствует конкретному вектору в массиве.
Для распараллеливания нужно использовать вторую ось, так как
ее относительно легко разделить на группы и вычислять скалярные
произведения в каждой группе отдельно. Если использовать только
vmap(), это должна быть ось, отображенная для обработки. Но для
применения pmap() необходимо сначала создать группы, поскольку
мы изменили форму исходных массивов так, что вторая ось разде
274
Глава 7
Распараллеливание вычислений
лена на две. После такой трансформации массив имеет форму (3, 8,
1 250 000), где первая ось остается неизменной (продолжает хранить
компоненты вектора), а вместо старой второй оси мы получили две
новые: ось для групп (с индексом 1) и ось для векторов внутри каж
дой группы (с индексом 2).
Затем мы применяем pmap() с передачей параметра in_axes=(1,1),
который сообщает о том, что оба входных тензора должны отобра
жаться по второй оси (с индексом 1) из существующих трех. Сле
довательно, вычисление будет распределено по восьми отдельным
устройствам, и каждое из них принимает свою отдельную группу
данных.
Внутри каждой группы продолжает находиться массив меньшего
размера (или пакет), состоящий из векторов, поэтому применяется
vmap() для отображения функции, обрабатывающей один элемент,
на все имеющиеся векторы. В функцию vmap() также необходимо пе
редать параметр in_axes, и хотя он выглядит точно так же, как и для
pmap(), т. е. jax.vmap(dot, in_axes=(1,1)), его смысл совершенно дру
гой. Трансформированная с помощью vmap() функция видит только
группу векторов, переданную на текущее конкретное устройство,
поэтому принимает массив с формой (3, 1 250 000). Отображение на
ось с индексом 1 здесь является отображением по оси, содержащей
1 250 000 элементов, а не по оси, содержащей восемь групп.
Результатом этого вычисления является массив с формой (8,
1 250 000), где каждая из восьми групп вычислений на отдельном
устройстве возвращает 1 250 000 скалярных произведений.
Завершающим действием становится удаление искусственно соз
данного измерения групп и объединение содержимого всех восьми
групп в единый массив, т. е. в итоге мы получаем массив результатов
с 10 миллионами скалярных произведений, вычисленных в парал
лельном режиме на восьми различных устройствах.
Теперь, чтобы лучше понять процесс, рассмотрим код в листин
ге 7.17.
Листинг 7.17 Пример обработки большого массива в параллельном
режиме
vs = random.normal(rng_key, shape=(20_000_000,3))
v1s = vs[:10_000_000,:].T
v2s = vs[10_000_000:,:].T
v1s.shape, v2s.shape
>>> ((3, 10000000), (3, 10000000))
v1sp = v1s.reshape((v1s.shape[0], 8, v1s.shape[1]//8))
❶
❶
❷
❸
275
Управление поведением pmap()
v2sp = v2s.reshape((v2s.shape[0], 8, v2s.shape[1]//8))
v1sp.shape, v2sp.shape
❸
>>> ((3, 8, 1250000), (3, 8, 1250000))
❹
dot_parallel = jax.pmap(
jax.vmap(dot, in_axes=(1,1)),
in_axes=(1,1)
)
❺
❻
x_pmap = dot_parallel(v1sp,v2sp)
x_pmap.shape
>>> (8, 1250000)
x_pmap = x_pmap.reshape((x_pmap.shape[0]*x_pmap.shape[1]))
x_pmap.shape
>>> (10000000,)
jax.numpy.all(x_pmap == x_vmap)
>>> Array(True, dtype=bool)
❼
❽
❾
❾
❶ Теперь используются транспонированные версии исходных массивов.
❷ Первое измерение содержит компоненты вектора. Второе измерение содержит
векторы.
❸ Разделение второго измерения, содержащего векторы, на два новых измерения:
группы и векторы.
❹ Мы получили восемь групп векторов.
❺ Оповещение vmap о необходимости использования второго измерения для ото❻
❼
❽
❾
бражения и обработки (vmap не видит измерение групп, поэтому его вторым измерением является измерение векторов).
Оповещение pmap о необходимости использования второго измерения (групп),
которое невидимо для vmap.
Получение восьми групп вычисленных скалярных произведений.
Удаление измерения групп.
Проверка с целью убедиться в том, что полученный результат является корректным (полностью совпадающим с результатом, полученным с использованием
vmap в начале главы).
Очевидно, что в листинге 7.17 приведен код реализации процесса,
показанного на рис. 7.1. Вы можете заменить константу 8 на вызов
функции jax.device_count(), чтобы код соответствовал вашей кон
фигурации.
Снова хочу обратить особое внимание на внимательное отслежи
вание измерений тензоров, поскольку при работе только с индексами
осей легко запутаться. Следует помнить, что параметр in_axes=(1,1)
имеет различный смысл для pmap() и vmap() в приведенном выше
276
Глава 7
Распараллеливание вычислений
коде. В таком коде очень легко сделать ошибку. В следующей главе
мы рассмотрим более надежные способы распараллеливания кода
с меньшей вероятностью возникновения ошибок.
При работе с индексами тензоров легко запутаться и исполь
зовать неправильные индексы или некорректно изменить форму
массива. Например, если мы меняем форму исходных массивов так,
чтобы измерение групп стало новым последним измерением тензо
ра, то должна измениться семантика векторов, и будут вычисляться
скалярные произведения по другим векторам. Поэтому рекоменду
ется быть особенно внимательными при выполнении всех операций
с индексами, проверять и перепроверять их, чтобы полностью убе
диться в правильности совершаемых действий.
7.2.2
Использование именованных осей и коллективных
операций
Во всех предыдущих примерах в этой главе вычисления не зависели
друг от друга, и этого было достаточно для выполнения простых па
раллельных операций. Но для более сложных вариантов обработки
данных, например при выполнении глобальной нормализации или
при работе со значениями глобального минимума/максимума, воз
можно, потребуется обмен данными и передача информации между
устройствами.
Обмен информацией между процессами
Из предыдущей главы, темой которой являлась автоматическая век
торизация, нам известно, что JAX предоставляет коллективные опе
рации (collective operations (ops)). Список доступных коллективных
операций приведен в табл. 7.1 (а также здесь: https://docs.jax.dev/en/
latest/jax.lax.html#parallel-operators).
Я пропустил операцию pdot() из списка на сайте документации,
потому что она не документирована и ее будущее неясно (https://
github.com/google/jax/discussions/13851).
Все перечисленные в табл. 7.1 операции выполняются по оси,
определяемой параметром axis_name, передаваемым при вызове
конкретной коллективной операции. Тот же параметр axis_name при
вызове pmap() или vmap() связывает заданное имя с отображаемой
осью, поэтому все коллективные операции могут ссылаться на это
имя.
Управление поведением pmap()
Таблица 7.1
277
Коллективные операции
Операция
Описание
all_gather(x, axis_name, *[, ...])
Генерирует значения x во всех репликах
all_to_all(x, axis_name, split_axis, ...[, ...])
Материализует отображенные оси
и выполняет отображение других осей
Вычисляет сумму с распространением
результата (all-reduce sum) по x по
трансформированной с помощью pmap оси
axis_name
Аналогична psum(x, axis_name), но каждое
устройство сохраняет только часть результата.
Это похоже на первую часть psum без
объединения результатов
Вычисляет максимум с распространением
результата (all-reduce max) по x по
трансформированной с помощью pmap оси
axis_name
Вычисляет минимум с распространением
результата (all-reduce min) по x по
трансформированной с помощью pmap оси
axis_name
Вычисляет среднее значение
с распространением результата (all-reduce
mean) по x по трансформированной с по
мощью pmap оси axis_name
Выполняет коллективную перестановку
в соответствии с условием перестановки perm
Служит удобной оберткой для jax.lax.
ppermute с альтернативной кодировкой
перестановки
Меняет местами ось axis_name,
трансформированную с помощью pmap,
и неотображаемую ось axis
Возвращает индекс по отображаемой оси
axis_name
psum(x, axis_name, *[, axis_index_groups])
psum_scatter(x, axis_name, *[, ...])
pmax(x, axis_name, *[, axis_index_groups])
pmin(x, axis_name, *[, axis_index_groups])
pmean(x, axis_name, *[, axis_index_groups])
ppermute(x, axis_name, perm)
pshuffle(x, axis_name, perm)
pswapaxes(x, axis_name, axis, *[, ...])
axis_index(axis_name)
Коллективные операции
Коллективные операции (collective operations (ops)) используются в параллельном программировании для обмена информацией между всеми
процессами в группе процессов (в отличие от двухточечного (point-topoint) обмена информацией между любыми специальными процессами). Это определяется стандартом Message Passing Interface (MPI).
Стандарт MPI содержит три класса операций: синхронизация, перемещение данных и коллективные вычисления.
В приведенном ниже списке описаны некоторые широко применяемые
функции (или паттерны). Это неполный список возможных операций, но
и этого вполне достаточно для понимания коллективных операций JAX:
Глава 7
278
Распараллеливание вычислений
broadcast: эта функция используется для распределения данных из
одного обрабатывающего элемента во все обрабатывающие элементы. В JAX использование значения None для параметра in_axes
означает, что аргумент не имеет дополнительной оси для отображения и должен распространяться по всем устройствам. Также можно
использовать специальный параметр static_broadcasted_argnums,
чтобы определить, какие позиционные аргументы необходимо интерпретировать как статические (константы времени компиляции) и распространяемые по всем устройствам;
scatter: этот паттерн используется для распределения данных из одного обрабатывающего элемента во все обрабатывающие элементы,
но в отличие от функции broadcast, отправляющей одно и то же сообщение каждому обрабатывающему элементу, scatter разделяет сообщение и передает только одну его часть каждому обрабатывающему
элементу. Именно такое действие выполняется, когда мы используем
отображение оси с помощью pmap() или vmap();
reduce: действие этого паттерна обратно broadcast. Он используется
для сбора данных из различных вычислительных устройств и объединения их в общий результат с помощью некоторой заданной функции
(например, суммирования);
all-reduce: это особый вариант паттерна reduce, в котором результат
операции reduce необходимо распространить по всем обрабатывающим элементам. Функции psum(), pmax(), pmin(), pmean() выполняют
операции all-reduce;
gather: этот паттерн используется для сохранения данных из всех обрабатывающих элементов в одном обрабатывающем элементе;
all-gather: этот паттерн используется для сохранения данных из всех
обрабатывающих элементов во всех обрабатывающих элементах;
all-to-all: этот паттерн, также называемый total exchange, используется, если требуется, чтобы каждый обрабатывающий элемент передавал свое сообщение всем прочим обрабатывающим элементам.
Если вы хотите более подробно узнать об MPI, то можете начать с введения: https://pdc-support.github.io/introduction-to-mpi. Также можно
ознакомиться с материалами курса «Designing and Building Applications for Extreme Scale Systems» здесь: https://wgropp.cs.illinois.edu/
courses/cs598-s15/.
В листинге 7.18 показан стандартный пример нормализации мас
сива в параллельном режиме с применением коллективной опера
ции psum(). Необходимо вычислить сумму всех элементов и разде
лить на нее каждый элемент.
279
Управление поведением pmap()
Листинг 7.18 Использование коллективной операции и параметра
axis_name
arr = jnp.array(range(8))
norm = jax.pmap(
lambda x: x/jax.lax.psum(x, axis_name='p'),
axis_name='p')
norm(arr)
❶
❷
❸
❹
>>> Array([0.
, 0.03571429, 0.07142857, 0.10714287, 0.14285715,
>>>
0.17857143, 0.21428573, 0.25
], dtype=float32)
jnp.sum(norm(arr))
>>> Array(1., dtype=float32)
❶
❷
❸
❹
❺
❺
Генерация массива для демонстрационных целей.
Использование коллективной операции psum() с параметром axis_name=’p’.
Использование параметра axis_name=’p’ в pmap().
Применение созданной функции нормализации.
Проверка нормализованных значений.
В приведенном выше примере параметр axis_name внутри вызова
pmap() присваивает имя трансформируемой с помощью этой функ
ции оси. Отображаемая ось определяется значением аргумента in_
axes, который здесь отсутствует, и по умолчанию используется ось
с индексом 0.
Мы присваиваем имя оси, отображаемой во время вызова pmap(),
и внутри распараллеленного вычисления используем функцию
psum(), вычисляющую сумму с распространением результата по
трансформированной с помощью pmap() оси axis_name='p'. Резуль
татом является вычисление суммы всех элементов по именованной
оси и деление каждого элемента на полученную сумму.
Это тривиальный учебный пример работы с очень маленьким
массивом, в котором каждый элемент может быть обработан на
собственном отдельном устройстве, но его с легкостью можно рас
ширить для работы с более крупными массивами, для которых по
требуется разделение на меньшие части и организация обработки
полученных подмассивов на отдельных устройствах.
Следует всегда помнить о том, что необходимо использовать две
различные операции: для агрегации значений внутри группы на од
ном устройстве и коллективную операцию для обмена информаци
ей между устройствами и объединения значений.
В листинге 7.19 выполняется нормализация более крупного мас
сива.
Глава 7
280
Листинг 7.19
Распараллеливание вычислений
Пример нормализации большого массива
arr = jnp.array(range(200))
arr = arr.reshape(8, 25)
arr.shape
❶
❷
>>> (8, 25)
norm = jax.pmap(
lambda x: x/jax.lax.psum(jnp.sum(x), axis_name='p'),
axis_name='p')
narr = norm(arr)
narr.shape
❸
❹
>>> (8, 25)
jnp.sum(narr)
>>> Array(1., dtype=float32)
❺
❶ Генерация массива, количество элементов в котором больше имеющихся аппа-
ратных устройств.
❷ Изменение формы массива: разделение на группы, количество которых равно
числу доступных XLA-устройств.
❸ Агрегация значений еще и внутри каждой группы.
❹ Применение функции нормализации.
❺ Проверка полученных нормализованных значений.
В листинге 7.19 единственное существенное изменение сдела
но внутри вызова psum(). Применяя функцию jnp.sum(), мы бе
рем сумму каждого массива, находящегося на каждом конкретном
устройстве, а затем, используя psum(), мы передаем полученные
суммы на все устройства и вычисляем общую сумму. Далее, при
меняя массовую передачу данных (broadcasting) в стиле NumPy,
на каждом устройстве делим каждый элемент массива на вычис
ленную общую сумму, чтобы получить нормализованную версию
массива. Мы не стали изменять форму итогового массива на перво
начальную (плоскую), поскольку здесь это неважно, но вы можете
легко проверить и убедиться в том, что теперь все элементы мас
сива в сумме дают 1.
Многие коллективные операции с распространением результата
предоставляют дополнительный (необязательный) параметр axis_
index_groups, позволяющий выполнять коллективные операции
для групп, охватывающих ось по частям, а не всю ось целиком. Этот
параметр представляет собой список списков, содержащих индек
сы оси. Каждый список является группой, по которой выполняется
коллективная операция.
Управление поведением pmap()
281
ПРИМЕЧАНИЕ Группы обязательно должны охватывать все
индексы оси только однократно. Кроме того, все группы не
пременно должны иметь одинаковый размер.
В приведенном ниже коде (листинг 7.20) предыдущий пример
нормализации изменен: теперь нормализация массива выполня
ется по нескольким группам. Определяются четыре группы: группа
для устройств с индексами 0 и 1, вторая группа для устройств с ин
дексами 2 и 3 и т. д. Коллективная операция будет выполняться для
каждой группы отдельно.
Листинг 7.20
Нормализация по группам
arr = jnp.array(range(200))
arr = arr.reshape(8, 25)
arr.shape
>>> (8, 25)
norm = jax.pmap(
lambda x: x/jax.lax.psum(
jnp.sum(x),
axis_name='p',
axis_index_groups=[[0,1], [2,3], [4,5], [6,7]]
),
axis_name='p')
❶
narr = norm(arr)
narr.shape
>>> (8, 25)
jnp.sum(narr)
>>> Array(4., dtype=float32)
jnp.sum(narr[:2]), jnp.sum(narr[2:4]), jnp.sum(narr[4:6]), jnp.sum(narr[6:])
>>> (Array(1., dtype=float32),
>>> Array(1., dtype=float32),
>>> Array(1.0000001, dtype=float32),
>>> Array(1., dtype=float32))
❶
❷
❸
❹
❷
❸
❹
❹
❹
❹
Определение индексов для четырех групп.
Теперь итоговая сумма по массиву равна 4.
Проверка каждой группы по отдельности.
Сумма элементов в каждой группе равна 1 (с погрешностью округления в одном случае).
Мы предоставили той же функции psum() список, определяющий,
какие группы должна вычислять эта функция. Выполнены четыре
Глава 7
282
Распараллеливание вычислений
операции psum(), и в результате все элементы нормализованы по со
ответствующим группам. Элементы в первых двух вырезках из ис
ходного массива в сумме дают 1, в следующих двух вырезках сумма
элементов также равна 1 и т. д. В сумме элементов в четвертой и пя
той вырезках присутствует погрешность округления.
Вложенные отображения
Благодаря функциональной сущности JAX можно с легкостью ком
бинировать различные трансформации, например выполнять
вложенные вызовы pmap() или смешивать pmap() и vmap(). Комби
нирование vmap() и pmap() особенно полезно, так как это обычная
ситуация при распараллеливании пакетной обработки по несколь
ким различным устройствам (компьютерам). Выполнение вложен
ных вызовов pmap() – редкий случай, особенно при небольшом ко
личестве устройств, предназначенных для распараллеливания, хотя
и такой вариант имеет смысл, если в коде имеется несколько вло
женных циклов.
Мы уже рассматривали комбинирование pmap() и vmap(), но те
перь можно расширить этот пример для использования различных
именованных осей, чтобы управлять осью при выполнении коллек
тивных операций. Здесь мы создаем комбинацию pmap() и vmap()
с двумя отдельными коллективными операциями: одна операция
pmax() выполняется внутри каждого массива, размещенного на от
дельном устройстве, вторая операция pmax() выполняется между не
сколькими различными устройствами.
В листинге 7.21 создается функция, вычисляющая отношение
между наибольшим элементом по всем пакетам (размещенным на
различных устройствах) и наибольшим элементом внутри отдель
ного пакета (размещенного на конкретном устройстве).
Листинг 7.21
Смешанное применение коллективных операций
arr = jnp.array(range(200))
arr = arr.reshape(8, 25)
arr.shape
>>> (8, 25)
f = jax.pmap(
jax.vmap(
lambda x: jax.lax.pmax(x,
axis_name='v')/jax.lax.pmax(x, axis_name='p'),
axis_name='v'
),
axis_name='p')
f(arr)
❶
❷
❸
Управление поведением pmap()
283
>>> Array([[0.13714285, 0.13636364, 0.1355932 , 0.13483146, 0.1340782 ,
>>>
0.13333334, 0.13259669, 0.13186814, 0.13114755, 0.13043478,
>>>
0.12972972, 0.12903225, 0.12834224, 0.12765957, 0.12698412,
>>>
0.12631579, 0.12565446, 0.125
, 0.12435233, 0.12371133,
>>>
0.12307693, 0.12244899, 0.12182741, 0.12121212, 0.12060301],
>>>
...
>>>
[1.1371429 , 1.1306819 , 1.1242937 , 1.1179775 , 1.1117318 ,
>>>
1.1055555 , 1.0994476 , 1.0934067 , 1.0874318 , 1.0815217 ,
>>>
1.0756756 , 1.0698925 , 1.0641712 , 1.0585107 , 1.05291 ,
>>>
1.0473684 , 1.0418848 , 1.0364584 , 1.0310881 , 1.0257732 ,
>>>
1.0205128 , 1.0153062 , 1.0101522 , 1.0050505 , 1.
]],
>>>
dtype=float32)
❶ Функция, использующая две коллективные операции по различным осям.
❷ Ось для vmap().
❸ Ось для pmap().
Массивы были сформированы так, что наибольшим значением
в каждом пакете является самый последний элемент, а самое боль
шое значение по всем пакетам – это элемент из последнего пакета.
Поэтому итоговые вырезки массива (или итоговый пакет) начинают
ся с большего значения, а затем значения постепенно уменьшаются
по мере прохождения по вырезке массива, поскольку наибольший
элемент внутри пакета является константой, тогда как наибольший
элемент по всем пакетам постепенно увеличивается. Последнее зна
чение самого последнего итогового пакета равно 1, потому что здесь
оба максимальных значения одинаковы.
Параметр axis_name позволяет с легкостью управлять выбором
оси для использования в каждой коллективной операции.
Можно одновременно идентифицировать несколько отображае
мых осей при использовании одной коллективной операции. В этом
случае передается кортеж, содержащий все имена осей.
Мы можем переписать пример общей нормализации из листин
га 7.19 так, чтобы исключить вызов jnp.sum() на каждом устройстве,
заменив его общим суммированием, включая также оси пакетов
(введенные трансформацией vmap()). Сумма вычисляется с приме
нением коллективной операции по обеим осям одновременно по
средством передачи кортежа axis_name=('p','v') в функцию psum().
Листинг 7.22 Общая нормализация с одновременным
использованием двух осей
arr = jnp.array(range(200))
arr = arr.reshape(8, 25)
arr.shape
>>> (8, 25)
Глава 7
284
Распараллеливание вычислений
norm = jax.pmap(
jax.vmap(
lambda x: x/jax.lax.psum(x, axis_name=('p','v')),
axis_name='v'
),
axis_name='p')
❶
narr = norm(arr)
narr.shape
❷
>>> (8, 25)
jnp.sum(narr)
>>> Array(1., dtype=float32)
❸
❶ Одновременное использование обеих осей.
❷ Применение функции нормализации.
❸ Проверка нормализованных значений.
Также можно использовать вложенные трансформации pmap(),
при этом существует единственное ограничение: количество XLAустройств не должно быть меньше произведения размеров отобра
жаемых осей. В приведенном ниже простом примере (листинг 7.23)
представлена небольшая матрица, в которой выполняется отобра
жение строк и столбцов. Мы используем ту же самую общую норма
лизацию по двум осям одновременно, что и в предыдущем примере.
Листинг 7.23
Пример вложенных трансформаций pmap()
arr = jnp.array(range(8)).reshape(2,4)
arr
❶
>>> Array([[0, 1, 2, 3],
[4, 5, 6, 7]], dtype=int32)
n = jax.pmap(
jax.pmap(
lambda x: x/jax.lax.psum(x, axis_name=('rows','cols')),
axis_name='cols'
),
axis_name='rows')
❷
jnp.sum(n(arr))
>>> Array(1., dtype=float32)
❶ Генерация небольшой матрицы.
❷ Выполнение вложенной трансформации pmap() по строкам и по столбцам.
❸ Проверка результата нормализации.
❸
Эту функцию можно сделать более ясно выраженной, если ис
пользовать стиль декоратора. В приведенном ниже примере мы за
Пример программы тренировки нейронной сети с распараллеливанием
285
меняем два вложенных вызова pmap() на два декоратора. Такой код
обычно проще читать и понимать.
Листинг 7.24 Вложение трансформаций pmap() с использованием
стиля декораторов
from functools import partial
@partial(jax.pmap, axis_name='rows')
@partial(jax.pmap, axis_name='cols')
def n(x):
return x/jax.lax.psum(x, axis_name=('rows','cols'))
❶
❷
❷
jnp.sum(n(arr))
>>> Array(1., dtype=float32)
❶ Импорт декоратора partial.
❷ Использование двух декораторов.
❸ Проверка результата нормализации.
❸
В варианте с вложенными вызовами имена осей разрешаются
в соответствии с правилами лексической области видимости. Пер
вый декоратор отвечает за внешний вызов, второй – за внутренний.
Теперь мы обладаем всеми требуемыми навыками и знаниями
для решения реальной задачи тренировки нейронной сети. В следу
ющем разделе мы разработаем программу тренировки с распарал
леливанием по данным нейронной сети классификации изображе
ний из главы 2.
7.3
Пример программы тренировки нейронной
сети с распараллеливаниемпо данным
Мы разрабатываем то, что можно было бы назвать примером SPMD
MNIST классификации изображений. Существуют различные спо
собы распараллеливания процесса тренировки нейронной сети, но
в данном случае мы воспользуемся методикой тренировки с распа
раллеливанием по данным.
Распараллеливание по данным и распараллеливание модели
Распараллеливание по данным (data parallelism) – это тип распараллеливания, при котором одна и та же операция выполняется в параллельном режиме для элементов некоторого набора данных. Данные разделяются на N фрагментов, распределенных по N узлам, и каждый узел
обрабатывает собственную часть данных.
Глава 7
286
Распараллеливание вычислений
В сценариях глубокого обучения большой набор данных обычно разделяется на фрагменты, распределенные по различным узлам. На каждом
узле имеется точная копия модели и происходит обработка собственного фрагмента данных. Результаты обработки (градиенты) передаются
между узлами и агрегируются (объединяются), так что каждый узел обновляет веса своей копии модели, чтобы сохранять единообразие обрабатывающей функции.
Распараллеливание модели (model parallelism) – это тип распараллеливания, при котором большая модель вычислений разделяется и распределяется по различным узлам, например вычисление отдельных слоев
нейронной сети на нескольких компьютерах. Такой подход позволяет
тренировать крупные модели, с которыми невозможно работать на одном компьютере.
Большие нейронные сети также можно тренировать с использованием
комбинации распараллеливания по данным и распараллеливания модели. Например, большая языковая модель компании Google под названием PaLM (https://arxiv.org/abs/2204.02311), которая имеет 540 млрд
параметров, тренировалась именно таким способом. Следует отметить,
что PaLM использовала JAX.
В следующей главе вы узнаете о других механизмах распаралле
ливания, которые можно использовать для тренировки с распарал
леливанием модели.
7.3.1
Подготовка данных и структуры нейронной сети
Напомним немного о настройках примера классификации изобра
жений из главы 2. Имеются рукописные изображения, представляю
щие цифры, взятые из базы данных MNIST. Некоторые изображения
показаны на рис. 7.2.
Рис. 7.2
Рукописные изображения цифр из базы данных MNIST
Пример программы тренировки нейронной сети с распараллеливанием
287
В главе 2 мы создали простую нейронную сеть типа MLP (multilayer
perceptron – многослойный перцептрон) с несколькими полностью
связанными слоями. Здесь мы снова используем эту сеть, но теперь
ее тренировка организована так, чтобы получить возможность ее
распределения по многим вычислительным устройствам, в нашем
случае – по восьми ядрам TPU.
Загрузка и подготовка набора данных почти такая же, как и ранее.
Единственное различие состоит в том, что теперь мы принимаем бо
лее крупные пакеты из набора данных, для которых в дальнейшем
можно изменить форму, чтобы получить набор пакетов для несколь
ких вычислительных устройств.
Листинг 7.25
Загрузка набора данных
import tensorflow as tf
import tensorflow_datasets as tfds
data_dir = '/tmp/tfds'
data, info = tfds.load(name="mnist",
data_dir=data_dir,
as_supervised=True,
with_info=True)
data_train = data['train']
data_test = data['test']
HEIGHT = 28
WIDTH = 28
CHANNELS = 1
NUM_PIXELS = HEIGHT * WIDTH * CHANNELS
NUM_LABELS = info.features['label'].num_classes
NUM_DEVICES = jax.device_count()
BATCH_SIZE = 32
❶
def preprocess(img, label):
"""Resize and preprocess images."""
return (tf.cast(img, tf.float32)/255.0), label
train_data = tfds.as_numpy(
data_train.map(preprocess).batch(
NUM_DEVICES*BATCH_SIZE).prefetch(1)
)
test_data = tfds.as_numpy(
data_test.map(preprocess).batch(
NUM_DEVICES*BATCH_SIZE).prefetch(1)
)
❷
❷
len(train_data)
>>> 235
❸
Глава 7
288
Распараллеливание вычислений
❶ Новая константа для определения количества вычислительных устройств.
❷ Запрос более крупных пакетов с размером 32 * (количество устройств).
❸ Этот набор данных содержит 235 больших пакетов.
В приведенном выше примере мы использовали функцию jax.device_count() для получения количества доступных вычислительных
устройств и увеличили размер пакета соответствующим образом,
чтобы позволить каждому устройству работать с размером паке
та 32.
Структура нейронной сети точно такая же, как в главе 2. Для удоб
ства в листинге 7.26 приведен ее код.
Листинг 7.26
Структура многослойного перцептрона (MLP)
import jax
import jax.numpy as jnp
from jax import grad, jit, vmap, value_and_grad
from jax import random
from jax.nn import swish, logsumexp, one_hot
LAYER_SIZES = [28*28, 512, 10]
PARAM_SCALE = 0.01
❶
def init_network_params(sizes, key=random.PRNGKey(0), scale=1e-2):
"""Initialize all layers for a fully-connected neural network
with given sizes"""
# Инициализация всех слоев для полностью связанной нейронной сети
# с заданными размерами.
❷
def random_layer_params(m, n, key, scale=1e-2):
"""A helper function to randomly initialize
weights and biases of a dense layer"""
# Вспомогательная функция для случайно выбранной инициализации
# весов и отклонений плотного слоя.
w_key, b_key = random.split(key)
return (scale * random.normal(w_key, (n, m)),
scale * random.normal(b_key, (n,)))
keys = random.split(key, len(sizes))
return [random_layer_params(m, n, k, scale)
for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
init_params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
def predict(params, image):
"""Function for per-example predictions."""
# Функция для прогнозов по одному образцу.
activations = image
for w, b in params[:-1]:
outputs = jnp.dot(w, activations) + b
❸
❹
Пример программы тренировки нейронной сети с распараллеливанием
289
activations = swish(outputs)
final_w, final_b = params[-1]
logits = jnp.dot(final_w, activations) + final_b
return logits
batched_predict = vmap(predict, in_axes=(None, 0))
❶
❷
❸
❹
❺
Определение количества нейронов в каждом полностью связанном слое.
Функция для случайно выбранной инициализации параметров.
Подготовка начальных параметров.
Функция прямого прохода.
Генерация пакетной функции прямого прохода.
❺
Здесь следует напомнить основной принцип: параметры нейрон
ной сети отделяются от использующей их функции прямого прохода,
чтобы функция не имела собственного состояния и являлась функ
ционально чистой. Мы генерируем случайно выбранные начальные
параметры и сохраняем их в переменной init_params. Также под
готавливается функция predict(), выполняющая прямой проход,
и создается пакетная версия этой функции с использованием vmap().
Инструкция in_axes=(None,0) здесь сообщает о том, что вместо ото
бражения первого параметра функции (параметры нейронной сети)
мы отображаем измерение с номером 0 для второго параметра (для
данных, которые обрабатывает функция).
Теперь у нас есть данные и структура нейронной сети, и мы гото
вы к реализации распараллеливания процедуры тренировки, опи
санной ниже:
1 данные
должны распределяться по всем доступным устрой
ствам. Параметры модели по-прежнему будут реплицироваться
на каждое устройство, чтобы на нем существовала собственная
полная локальная копия параметров модели для выполнения
шага обновления градиента;
2 изменение функции update() для агрегации градиентов по
всем устройствам с использованием коллективной операции
jax.lax.psum(). В остальном код остается неизменным, т. е. па
раметры модели обновляются локально и возвращаются обнов
ленные параметры модели и значение потерь;
3 трансформация функции update() с помощью pmap() с указани
ем на то, что данные отображаются (распределяются по устрой
ствам), но параметры модели реплицируются;
4 обновление цикла тренировки так, чтобы после приема пакета
данных из загрузчика можно было изменить форму пакета, до
бавив дополнительную ось в pmap();
5 агрегация потерь, полученных на каждом устройстве.
Теперь все изменения достаточно точно определены. Рассмотрим
их подробнее.
Глава 7
290
7.3.2
Распараллеливание вычислений
Реализация процедуры тренировки
с распараллеливанием по данным
Теперь необходимо подготовить функцию потерь и функцию для
обновления параметров нейронной сети. Изменения заключаются
в отображении функции обновления на несколько вычислительных
устройств. Для этого потребуется внести некоторые существенные
изменения по сравнению с тренировкой модели в главе 2.
Начнем с самой общей идеи. Нам нужен большой массив трени
ровочных образцов и разделение этого массива по одной выбранной
оси с распределением по нескольким устройствам так, чтобы каж
дое устройство имело собственную часть набора данных (или более
конкретно: часть текущего большого пакета). Для выполнения ите
рации тренировочного процесса каждому устройству также потре
буется копия параметров модели, поэтому необходима репликация
начальных параметров нейросети на всех устройствах без разделе
ния и без каких-либо изменений. Следовательно, каждое устройство
будет иметь собственную копию параметров нейросети и часть на
бора данных и сможет выполнить шаг вычисления градиента ло
кально. При этом на каждом устройстве будут находиться собствен
ные градиенты, поэтому необходимо объединить (просуммировать)
все вычисленные градиенты по всем устройствам и передать объ
единенные градиенты на каждое устройство, чтобы предоставить
устройствам возможность локального выполнения шага обновления
параметров. После завершения этого шага каждое устройство снова
будет иметь одинаковые значения параметров, и процесс можно по
вторить.
Углубляясь в конкретные подробности, мы предполагаем, что
функция update() будет работать со следующими сущностями:
с копией параметров нейронной сети (с первым параметром
функции params). Параметры будут одинаковыми для всех
устройств, но поскольку они должны быть размещены локаль
но на каждом устройстве, придется их скопировать. Можно ре
плицировать параметры модели вручную или положиться на
саму трансформацию pmap(), пометив первый параметр функ
ции update() значением None в параметре in_axes. В автомати
зированном варианте pmap()-трансформированная функция
реплицирует параметры модели на каждом устройстве перед
началом вычисления. При репликации вручную необходи
мо подготовить специальную версию структуры, содержащую
параметры нейронной сети с дополнительной осью и уже ре
плицированными копиями по этой оси. Затем трансформация
pmap() должна отобразить это дополнительное измерение зна
чений параметров. В любом случае каждое устройство получит
291
Пример программы тренировки нейронной сети с распараллеливанием
собственную копию весов каждого слоя. Мы выбрали автомати
зированный способ;
с тренировочными данными, состоящими из пакета изображе
ний (параметр x) и соответствующих меток (параметр y). Здесь
мы просто разделяем большой пакет на группы пакетов мень
шего размера (в нашем случае размер равен 32) с отдельной
осью для отображения, чтобы каждое устройство получило соб
ственную часть большого пакета;
с некоторыми дополнительными параметрами – в нашем слу
чае это номер эпохи epoch_number, – позволяющими постепен
но снижать скорость обучения. Этот параметр будет передан
на каждое устройство в режиме broadcast во время вызова
функции.
Функция update() возвращает параметры модели и значения по
терь – много обновленных значений, по одному из каждого устрой
ства. На следующей итерации обновления снова потребуется ре
пликация параметров модели, но перед этим необходимо выбрать
параметры модели, взятые с одного из устройств. Это наиболее
подходящее место для применения особого варианта с параметром
out_axes функции pmap() посредством установки значения None
для первого возвращенного параметра и приема этого значения
только с первого устройства. Мы уверены в том, что обновленные
параметры модели одинаковы на каждом устройстве, так как они
были одинаковыми в начале шага обновления (поскольку являлись
реплицированными), а далее мы выполнили шаг вычисления гра
диента, поскольку градиенты агрегированы по всем устройствам.
Второе возвращаемое значение содержит потери, различные на всех
устройствах, потому что на них обрабатываются разные части дан
ных, поэтому второе значение в параметре out_axes равно 0.
В листинге 7.27 реализованы описанные выше изменения.
Листинг 7.27
Функции loss и update
from functools import partial
INIT_LR = 1.0
DECAY_RATE = 0.95
DECAY_STEPS = 5
NUM_EPOCHS = 20
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits = batched_predict(params, images)
log_preds = logits - logsumexp(logits)
return -jnp.mean(targets*log_preds)
❶
292
Глава 7
Распараллеливание вычислений
@partial(jax.pmap,
axis_name='devices',
in_axes=(None, 0, 0, None),
out_axes=(None,0))
❷
def update(params, x, y, epoch_number):
loss_value, grads = value_and_grad(loss)(params, x, y)
grads = [(jax.lax.psum(dw, 'devices'),
❸
jax.lax.psum(db, 'devices'))
❸
for dw, db in grads]
❸
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)], loss_value
❶ Та же функция потерь, что и ранее.
❷ Использование pmap() для распараллеливания функции и отображения только
данных; параметры реплицируются, но не отображаются.
❸ Накопление градиентов на каждом устройстве.
Функция потерь loss() абсолютно та же самая, никаких измене
ний, но в функцию update() внесены важные изменения.
Мы использовали декоратор, чтобы определить, что параметры,
связанные с данными, имеют измерение с номером 0 для отображе
ния, а первый и четвертый параметры передаются на каждое устрой
ство в режиме broadcast. Это описывает параметр in_axes=(None, 0,
0, None). В итоге каждый вызов функции update() получает часть на
бора данных для тренировки и копию параметров нейронной сети,
поэтому создается возможность вычисления значения потерь и гра
диентов локально на каждом устройстве.
После локального вычисления градиентов на каждом устройстве
необходимо объединить полученные градиенты и отправить резуль
тат объединения обратно на каждое устройство, чтобы обеспечить
наличие одинаковых результатов. Это делается с помощью коллек
тивной операции psum() для весов и отклонений каждого слоя.
Теперь каждое устройство имеет в своем распоряжении сумму
градиентов по всем устройствам, и с этого момента мы готовы вы
полнить шаг обновления градиентов. Это делается локально на каж
дом устройстве, и поскольку значения первоначальных параметров
и объединенные градиенты одинаковы на всех устройствах, мы по
лучаем один и тот же результат на каждом устройстве. Берем обнов
ленные параметры с первого устройства и распространяем их в ре
жиме broadcast на следующей итерации для повторения процесса.
Предположим, что мы воспользовались некоторой стохастичностью
в процессе обновления; каждое устройство должно было получить
собственную версию параметров. В этом случае также потребова
лось бы объединение этих параметров каким-то способом, но это не
наш вариант.
293
Пример программы тренировки нейронной сети с распараллеливанием
Можно лучше понять, что происходит внутри описанного процес
са, если наблюдать за формой тензоров.
Листинг 7.28
Формы тензоров данных
train_data_iter = iter(train_data)
x, y = next(train_data_iter)
x.shape, y.shape
>>> ((256, 28, 28, 1), (256,))
x = jnp.reshape(x, (NUM_DEVICES, BATCH_SIZE, NUM_PIXELS)
y = jnp.reshape(
one_hot(y, NUM_LABELS),
(NUM_DEVICES, BATCH_SIZE, NUM_LABELS))
x.shape, y.shape
>>> ((8, 32, 784), (8, 32, 10))
updated_params, loss_value = update(init_params, x, y, 0)
loss_value
>>> Array([0.5771865 , 0.5766423 , 0.5766001 ,
➥0.57689124, 0.57701343,
>>>
0.57676095, 0.57668227, 0.5764269 ],
➥ dtype=float32)
❶
❶
❷
❸
❸
❸
❹
❺
❺
❶ Получение пакета данных из тренировочного набора.
❷ Здесь используется пакет, содержащий 256 элементов.
❸ Изменение формы данных для используемого MLP и добавление отдельного из-
мерения для распараллеливания по восьми устройствам.
❹ Выполнение одного шага обновления градиентов.
❺ Каждое устройство содержит собственное значение потерь, потому что выполня-
ло тренировку с собственной частью набора данных.
В рассматриваемом здесь примере выполнено множество опера
ций изменения формы тензора.
В исходной версии в главе 2 принимался пакет с 32 элементами,
представляющими тензор изображения (28×28×1), и тензор меток
(с одним скалярным значением для каждого изображения). Затем
мы «выпрямляли» каждое изображение, превращая его в последо
вательность из 784 пикселов, и производили прямое кодирование
с одним активным состоянием (one-hot encoding) каждой метки
в вектор с 10 элементами. Для обоих тензоров предварительно уста
навливалось измерение пакетов, равное 32.
Теперь в дополнение к вышеописанным трансформациям мы раз
деляем измерение пакетов (пакет стал больше, он содержит 256 эле
ментов) на два: измерение устройств (в данном случае его размер
294
Глава 7
Распараллеливание вычислений
равен 8 по количеству доступных устройств) и измерение пакетов
(со старым размером 32). Здесь форма (8, 32, 784) означает, что мы
имеем пакет из 32 «выпрямленных» изображений длиной 784 пиксе
ла для каждого из восьми устройств. То же самое относится и к фор
ме (8, 32, 10), где имеется пакет из 32 векторов из 10 элементов
в прямой кодировке с одним активным состоянием для каждого из
восьми устройств.
После выполнения одного шага обновления мы получаем обнов
ленные параметры (взятые из первого устройства) и список потерь
из каждого устройства. Значения потерь немного отличаются, по
скольку каждое устройство вычислило потери и соответствующие
градиенты по отдельным фрагментам данных. Чтобы получить объ
единенные потери, можно просто просуммировать все вычислен
ные их значения.
При таком подходе выполняются две потенциально большие опе
рации обмена информацией между устройствами, связанные с ве
сами модели. Во-первых, мы передаем параметры модели на все
устройства в начале каждой итерации. Во-вторых, происходит об
мен градиентами между всеми устройствами, и измерение градиен
тов совпадает с измерением весов, так как каждый исследуемый вес
производит собственный градиент.
Можно было бы применить более эффективный подход с посто
янным хранением отдельной копии параметров модели на каждом
устройстве и локальным их обновлением с использованием объеди
ненных градиентов. Этот подход мы реализуем в следующей главе
с применением методики сегментирования тензоров, а также про
делаем то же самое с использованием pmap() в разделе 10.2.
Для завершения необходимо реализовать полный цикл трени
ровки. Здесь почти все понятно и просто. Воспользуемся теми же
функциями вычисления точности, что и в главе 2. Во многом тре
нировочный цикл остается тем же самым, но все же требуется из
менение формы данных, описанное выше. Теперь имеется изме
рение устройств для распараллеливания и измерение пакетов для
использования их на каждом устройстве. Единственное изменение
заключается в том, что мы не используем константу BATCH_SIZE,
а вместо нее вводится выражение num_elements/NUM_DEVICES, потому
что последний большой пакет из генератора данных может содер
жать меньше элементов, чем требуется для полного пакета, поэтому
мы сохраняем измерение устройств неизменным, но размер пакета
для устройства может уменьшиться в таком особом случае. Вы обя
зательно должны отрегулировать эту часть при реализации процес
са тренировки с распараллеливанием по данным для своих наборов
данных. Размер набора может не позволять разделить его на пакеты
равной величины для всех доступных устройств, и последний пакет
Пример программы тренировки нейронной сети с распараллеливанием
295
потребует более тщательной обработки (иначе вы получите неожи
данную ошибку).
Листинг 7.29
Полный цикл тренировки
@jit
❶
def batch_accuracy(params, images, targets):
images = jnp.reshape(images, (len(images), NUM_PIXELS))
predicted_class = jnp.argmax(batched_predict(params, images), axis=1)
return jnp.mean(predicted_class == targets)
def accuracy(params, data):
accs = []
for images, targets in data:
accs.append(batch_accuracy(params, images, targets))
return jnp.mean(jnp.array(accs))
❶
import time
params = init_params
for epoch in range(NUM_EPOCHS):
start_time = time.time()
losses = []
for x, y in train_data:
num_elements = len(y)
x = jnp.reshape(x,
(NUM_DEVICES, num_elements//NUM_DEVICES, NUM_PIXELS))
y = jnp.reshape(one_hot(y, NUM_LABELS),
(NUM_DEVICES, num_elements//NUM_DEVICES, NUM_LABELS))
params, loss_value = update(params, x, y, epoch)
losses.append(jnp.sum(loss_value))
epoch_time = time.time() - start_time
train_acc = accuracy(params, train_data)
test_acc = accuracy(params, test_data)
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
print("Training set loss {}".format(jnp.mean(jnp.array(losses))))
print("Training set accuracy {}".format(train_acc))
print("Test set accuracy {}".format(test_acc))
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
…
>>>
Epoch 0 in 16.44 sec
Training set loss 3.3947665691375732
Training set accuracy 0.9246563911437988
Test set accuracy 0.9237304925918579
Epoch 1 in 8.64 sec
Training set loss 3.055901050567627
Training set accuracy 0.9457446336746216
Test set accuracy 0.946582019329071
Epoch 19 in 9.53 sec
❷
❸
❸
❹
❺
Глава 7
296
Распараллеливание вычислений
>>> Training set loss 2.8397674560546875
>>> Training set accuracy 0.9905917048454285
>>> Test set accuracy 0.98095703125
❶
❷
❸
❹
❺
Функции для вычисления точности.
Начало цикла тренировки.
Изменение формы данных для MLP и распараллеливания.
Обновление параметров с применением распараллеливания по данным.
Объединение значений потерь для эпохи.
Обратите внимание на то, как вычисляется точность. Это делается
только на одном устройстве, но, возможно, потребуется распарал
леливание и этого вычисления (предлагается в качестве отдельного
учебного упражнения).
Вы, вероятно, также заметили, что мы объединяем потери, сумми
руя их значения. Кроме того, мы не используем jit() для функции
обновления, так как это и не требуется. Напомню: pmap() автомати
чески во внутреннем режиме выполняет JIT-компиляцию функции.
Этот пример тренировки с распараллеливанием по данным так
же является примером еще одной комбинации функций трансфор
мации, поскольку здесь объединены трансформации grad(), vmap()
и pmap(). А также скрыто используется трансформация jit().
7.4
Использование конфигураций
с несколькими хостами
Мы почти закончили, но осталась еще одна тема, заслуживающая
внимания: использование JAX в средах с несколькими хостами CPU
или JAX-процессами.
Ранее мы ограничивали конфигурации системами с одним хос
том, включающими максимум одну плату TPU с восемью ядрами. Но
если мы имеем дело с TPU, то можем получить в свое распоряжение
более крупные системы, именуемые TPU Pods, или их подмножество
под названием TPU Pod slices (https://cloud.google.com/tpu/docs/sys
tem-architecture-tpu-vm#tpu_slices).
JAX использует модель программирования мультиконтроллера
(multicontroller programming model). В этой модели каждый процесс
JAX Python выполняется независимо, и одна и та же программа JAX
Python работает в каждом процессе. Такой подход отличается от дру
гих распределенных систем, где один контроллер управляет многи
ми рабочими узлами.
Это означает, что вы обязательно должны вручную запустить JAXпрограмму на каждом хосте, и в настоящее время не существует
приемлемого способа управления несколькими процессами JAX Py
thon из одного блокнота Colab.
Использование конфигураций с несколькими хостами
297
Прежде чем мы начнем изучение конфигураций с несколькими
хостами, необходимо ознакомиться с некоторыми особенностями:
во-первых, вы обязательно должны создать экземпляр кластера
и инициализировать его. Обычно это делается с помощью вызова
jax.distributed.initialize(). Такое действие не является обяза
тельным для TPU (хотя рекомендуется), но требуется для GPU. Не
обходимо выполнить такую операцию до начала каких-либо вы
числений с применением JAX. В дальнейшем вы сможете получить
индексы конкретных процессов, вызвав функцию jax.process_index().
Во-вторых, как вы помните, существует различие между глобаль
ными и локальными устройствами, описанное в подразделе 3.2.3:
к локальным устройствам (local devices) можно обращаться ло
кально из процесса JAX Python. Существуют GPU, напрямую со
единенные с хостом в GPU-системе. Для Cloud TPU это все ядра
TPU на общей плате, соединенной с хостом (максимум восемь
ядер на один хост). Ранее мы работали только с локальными
устройствами;
глобальные устройства (global devices) – это устройства, доступ
ные всем процессам JAX на различных хостах. Например, для
подмножества v2-32 TPU Pod slice 32 ядра TPU v2 соединены
с четырьмя хостами, на каждом из которых имеется общая пла
та TPU v2, содержащая восемь ядер.
Каждый процесс может видеть только свои локальные устройства,
хотя способен обмениваться информацией со всеми глобальными
устройствами, используя коллективные операции.
Можно воспользоваться pmap() для выполнения вычислений,
распределенных по нескольким процессам. Каждый экземпляр
pmap() видит только свои локальные устройства, поэтому необхо
димо подготовить данные для каждого процесса отдельно. Каждый
процесс JAX выполняет вычисления локально со своим фрагмен
том данных.
Но при вызове коллективных операций внутри функции они ра
ботают со всеми глобальными устройствами, поэтому кажется, что
pmap() обрабатывает массив, сегментированный по различным хос
там. Каждый хост видит и обрабатывает только собственный сег
мент, но хосты обмениваются информацией с помощью коллектив
ных операций.
Напишем программу, выполняющую распараллеленное вычис
ление скалярного произведения в кластере, а затем вычисляющую
общую сумму вычисленных скалярных произведений. Програм
ма работает на всех хостах TPU. Скопируйте код из листинга 7.30
в файл с именем worker.py, затем распределите его по кластеру
и запустите.
298
Глава 7
Распараллеливание вычислений
Листинг 7.30. Программа для кластера TPU Pod slice
import jax
import jax.numpy as jnp
from jax import random
jax.distributed.initialize()
print('== Running worker: ', jax.process_index())
def dot(v1, v2):
return jnp.vdot(v1, v2)
rng_key = random.PRNGKey(42 + 10*jax.process_index())
vs = random.normal(rng_key, shape=(2_000_000,3))
v1s = vs[:1_000_000,:]
v2s = vs[1_000_000:,:]
# Общее количество ядер TPU в Pod.
device_count = jax.device_count()
# Количество ядер TPU, подключенных к этому хосту.
local_device_count = jax.local_device_count()
if jax.process_index() == 0:
print('-- global device count:', jax.device_count())
#print('global devices:', jax.devices())
print('-- local device count:', jax.local_device_count())
#print('local devices:', jax.local_devices())
print('-- JAX version:', jax.__version__)
❶
❷
❸
❸
❸
❸
❹
❺
❻
v1sp = v1s.reshape(
(local_device_count,
v1s.shape[0]//local_device_count,
v1s.shape[1]))
v2sp = v2s.reshape(
(local_device_count,
v2s.shape[0]//local_device_count,
v2s.shape[1]))
if jax.process_index() == 0:
print('-- v1sp shape: ', v1sp.shape)
dots = jax.pmap(jax.vmap(dot))(v1sp,v2sp)
if jax.process_index() == 0:
print('-- dots shape: ', dots.shape)
# Форма (8, 125000, 3)
# Форма (8, 125000)
global_sum = jax.pmap(
lambda x: jax.lax.psum(jnp.sum(x), axis_name='p'),
axis_name='p'
)(dots)
if jax.process_index() == 0:
❼
❽
❾
299
Использование конфигураций с несколькими хостами
print('-- global_sum shape: ', global_sum.shape)
# Форма (8,)
print(f'== Worker {jax.process_index()} global sum: {global_sum}')
dots = dots.reshape((dots.shape[0]*dots.shape[1]))
if jax.process_index() == 0:
print('-- result shape: ', dots.shape)
# Форма (1000000,)
local_sum = jnp.sum(dots)
print(f'== Worker {jax.process_index()} local sum: {local_sum}')
❿
⓫
print(f'== Worker {jax.process_index()} done')
❶ Инициализация кластера.
❷ Каждый рабочий узел выводит сообщение со своим идентификатором.
❸ Генерация случайного тензора, различного на каждом хосте, с использованием
❹
❺
❻
❼
❽
❾
❿
⓫
разных значений seed для генератора случайных чисел.
Получение значения счетчика глобальных устройств (здесь: 32).
Получение значения счетчика локальных устройств (здесь: 8).
Процесс 0 выводит информацию.
Каждый хост обеспечивает распараллеливание вычисления скалярного произведения по своим восьми ядрам TPU.
Отдельная трансформация pmap() выполняет psum() по всем глобальным устройствам.
Каждый сегмент на каждом ядре TPU вычисляет сумму всех элементов.
Проверка: каждый процесс должен вычислить одну и ту же глобальную сумму.
Каждый процесс также вычисляет собственную локальную сумму.
В приведенном выше примере мы сначала инициализируем клас
тер с помощью функции jax.distributed.initialize(). Далее каж
дый компьютер подготавливает собственный массив случайных
векторов. Мы использовали различные значения seed для генерато
ра случайных чисел (более подробно об этом – в следующей главе),
чтобы получить отличающиеся значения для каждого рабочего узла.
Значения seed основаны на вызове функции jax.process_index(),
возвращающей разные числа для различных JAX-процессов в клас
тере. Данные сегментируются, и полученные сегменты загружаются
по отдельности для каждого рабочего узла.
Для диагностики мы выводим номера глобальных и локальных
устройств из процесса JAX с индексом 0.
Затем каждый хост с платой TPU (и восемью микросхемами TPU)
выполняет распараллеленное и векторизованное вычисление ска
лярного произведения, независимо обрабатывая собственный фраг
мент данных. Здесь отсутствует какой-либо обмен информацией
между различными хостами.
Ситуация меняется при следующем вызове pmap(), где вычисля
ется глобальная сумма. Каждый хост выполняет jax.lax.psum(jnp.
sum(x)) на каждом из своих восьми ядер TPU. На всех хостах вход
ные данные для pmap() имеют форму (8, 125 000), которая разделя
ется по первому измерению, и каждое ядро TPU получает вектор из
300
Глава 7
Распараллеливание вычислений
125 000 значений (скалярных произведений из предыдущего эта
па). Вычисляется сумма этих значений с использованием jnp.sum(),
и получается единственное число, а затем выполняется коллектив
ная операция. Все ядра TPU на всех хостах – всего 32 ядра – обме
ниваются своими локальными суммами и вычисляют общую сумму
с накоплением результата с помощью jax.lax.psum(). В итоге каж
дый рабочий узел получает массив из восьми элементов, содержа
щий одно и то же число, которое и является глобальной суммой по
всему кластеру.
Далее вычисляются локальные суммы на каждом рабочем узле
и выводятся для визуальной проверки, позволяющей увидеть, что
на каждом рабочем узле имеется собственная сумма, и эти значения
объединяются в глобальную сумму, равную значению, полученному
в итоговом массиве.
Теперь создадим небольшой кластер TPU Pod slice для выполне
ния программы.
ПРИМЕЧАНИЕ Использование TPU Pod slice стоит дорого.
Всегда проверяйте текущую цену (https://cloud.google.com/
tpu/pricing#pod-pricing) и не забывайте удалять кластер TPU
Pod slice, когда он больше не нужен.
Доступные кластеры TPU Pod и зоны их предоставления можно
посмотреть здесь: https://cloud.google.com/tpu/docs/regions-zones.
При написании кода этого примера я обнаружил, что кластер
v2-32 TPU Pod slice доступен в зонах us-central1-a и europe-west4-a.
Следовательно, мы можем создать его:
$gcloud compute tpus tpu-vm create tpu-pod
➥--zone europe-west4-a --accelerator-type v2-32
➥--version tpu-vm-base
Create request issued for: [tpu-pod]
Waiting for operation [projects/true-poet-371617/
➥locations/europe-west4-a/operations/
➥operation-1676557945942-5f4d210d03
➥274-dd57b5fd-0bfd83f9] to complete...done.
Created tpu [tpu-pod].
Примите мои поздравления, если вы впервые используете TPUсуперкомпьютер.
Кластер TPU Pod создан, теперь необходимо подготовить его для
запуска JAX-программы. В первую очередь нужно установить JAX:
$gcloud compute tpus tpu-vm ssh tpu-pod
➥--zone europe-west4-a --worker=all
➥--command="pip install 'jax[tpu]>=0.2.16' -f
https://storage.googleapis.com/jax-releases/libtpu_releases.html"
Использование конфигураций с несколькими хостами
301
SSH: Attempting to connect to worker 0...
SSH: Attempting to connect to worker 1...
SSH: Attempting to connect to worker 2...
SSH: Attempting to connect to worker 3...
Looking in links: https://storage.googleapis.com/jax-releases/libtpu_
releases.html
Collecting jax[tpu]>=0.2.16
Downloading jax-0.4.3.tar.gz (1.2 MB)
…
Successfully installed jax-0.4.3 jaxlib-0.4.3 libtpu-nightly0.1.dev20230207
numpy-1.24.2 opt-einsum-3.3.0 scipy-1.10.0
Successfully installed jax-0.4.3 jaxlib-0.4.3 libtpu-nightly0.1.dev20230207
numpy-1.24.2 opt-einsum-3.3.0 scipy-1.10.0
Возможно, вы увидите сообщения об ошибках SSH-соединения,
если заранее не установили соединения с указанными компьютера
ми. Чтобы узнать, как устранить эти ошибки, обратитесь к руковод
ству https://cloud.google.com/tpu/docs/jax-pods.
Теперь мы полностью готовы к распределению JAX Python про
граммы по всем хостам кластера TPU Pod slice:
$gcloud compute tpus tpu-vm scp worker.py
➥tpu-pod: --worker=all --zone=europe-west4-a
WARNING: Cannot retrieve keys in ssh-agent. Command may stall.
SCP: Attempting to connect to worker 0...
SCP: Attempting to connect to worker 1...
SCP: Attempting to connect to worker 2...
SCP: Attempting to connect to worker 3...
worker.py
| 1 kB |
1.8 kB/s | ETA: 00:00:00 |
worker.py
| 1 kB |
1.8 kB/s | ETA: 00:00:00 |
worker.py
| 1 kB |
1.8 kB/s | ETA: 00:00:00 |
worker.py
| 1 kB |
1.8 kB/s | ETA: 00:00:00 |
100%
100%
100%
100%
На завершающем этапе выполняется наша программа:
gcloud compute tpus tpu-vm ssh tpu-pod
➥--zone europe-west4-a --worker=all --command "python3 worker.py"
SSH: Attempting to connect to worker 0…
SSH: Attempting to connect to worker 1…
SSH: Attempting to connect to worker 2…
SSH: Attempting to connect to worker 3…
== Running worker: 2
== Worker 2 global sum: [-45.239502 -45.239502 -45.239502 -45.239502
➥-45.239502 -45.239502 -45.239502 -45.239502]
== Worker 2 local sum: -367.02166748046875
== Worker 2 done
== Running worker: 0
-- global device count: 32
Глава 7
302
-----==
➥
-==
==
==
==
➥
==
==
==
==
➥
==
==
Распараллеливание вычислений
local device count: 8
JAX version: 0.4.3
v1sp shape: (8, 125000, 3)
dots shape: (8, 125000)
global_sum shape: (8,)
Worker 0 global sum: [-45.239502 -45.239502
-45.239502 -45.239502 -45.239502 -45.239502
result shape: (1000000,)
Worker 0 local sum: 1397.5830078125
Worker 0 done
Running worker: 1
Worker 1 global sum: [-45.239502 -45.239502
-45.239502 -45.239502 -45.239502 -45.239502
Worker 1 local sum: -992.9242553710938
Worker 1 done
Running worker: 3
Worker 3 global sum: [-45.239502 -45.239502
-45.239502 -45.239502 -45.239502 -45.239502
Worker 3 local sum: -82.87371826171875
Worker 3 done
-45.239502
-45.239502]
-45.239502
-45.239502]
-45.239502
-45.239502]
Здесь вы видите общий объединенный вывод результатов, полу
ченных на всех рабочих узлах. На каждом рабочем узле имеется соб
ственный массив из восьми элементов с глобальной суммой, и все
рабочие узлы выводят свою локальную сумму, так что вы можете
проверить результаты и убедиться, что все вычислено верно.
Все операции обмена информацией между 32 ядрами TPU в клас
тере завершились успешно.
После завершения работы не забудьте удалить кластер TPU Pod,
если он больше не нужен:
$gcloud compute tpus tpu-vm delete tpu-pod --zone europe-west4-a
You are about to delete tpu [tpu-pod]
Do you want to continue (Y/n)?
Delete request issued for: [tpu-pod]
Waiting for operation [projects/true-poet-371617/
➥locations/europe-west4-a/operations/operation➥1676559317550-5f4d262914486-f827e901-ce9ce3f7]
➥ to complete...done.
Deleted tpu [tpu-pod].
Важно помнить о том, что на платформах с несколькими хостами
входные данные для функций, трансформированных с помощью
pmap(), обязательно должны иметь размер главной оси, равный
количеству локальных, а не глобальных устройств. Каждый хост
способен без каких-либо затруднений обращаться только к своим
локальным устройствам. Обмен информацией между глобальными
устройствами выполняется с использованием коллективных опе
раций.
Резюме
303
О мультипроцессных средах более подробно можно узнать здесь:
https://docs.jax.dev/en/latest/multi_process.html.
Мы завершили изучение методики явного распараллеливания
с применением pmap(). В следующей главе мы рассмотрим другой
подход к организации распараллеливания с использованием сег
ментирования тензоров.
Резюме
Распараллеливание может повысить скорость кода, выполняя вы
числения на нескольких устройствах в параллельном режиме.
Параллельное отображение, или трансформация pmap(), пред
ставляет собой простой и понятный способ распределения вычис
лений по различным устройствам. Эта трансформация позволяет
писать SPMD-программы и выполнить их на нескольких устрой
ствах.
Трансформация pmap() компилирует функцию с помощью XLA
(поэтому нет необходимости в отдельной трансформации jit()),
реплицирует скомпилированную функцию на устройствах и вы
полняет каждую реплику на отдельном устройстве в параллель
ном режиме.
Размер отображаемой оси обязательно должен быть меньше или
равен количеству доступных локальных XLA-устройств, которое
возвращается функцией jax.local_device_count().
Необходимо изменить форму массивов, чтобы измерение, по ко
торому выполняется распараллеливание, соответствовало коли
честву устройств.
Можно управлять выбором оси для отображения, используя пара
метр in_axes.
Можно управлять схемой размещения выходных данных с по
мощью параметра out_axes.
Параметр axis_name позволяет использовать коллективные опе
рации, если в коде требуется обмен информацией между различ
ными устройствами.
Можно воспользоваться необязательным параметром axis_index_groups, позволяющим выполнять коллективные операции для
групп, покрывающих ось, а не для всей оси одновременно.
Благодаря функциональной сущности JAX можно с легкостью
комбинировать различные трансформации, например создавать
вложенные вызовы pmap() или объединять pmap() и другие транс
формации, такие как vmap() или grad().
Использование стиля декораторов для pmap() иногда помогает не
много уменьшить объем исходного кода.
Глава 7
304
Распараллеливание вычислений
Трансформацию pmap() можно использовать для реализации
процесса тренировки нейронной сети с распараллеливанием по
данным.
JAX использует модель программирования мультиконтроллера,
где каждый процесс JAX Python выполняется независимо, и одна
и та же программа JAX Python работает в каждом процессе.
В средах со многими хостами каждая трансформация pmap() мо
жет видеть только свои локальные устройства и каждый процесс
JAX выполняет вычисления локально с собственным фрагментом
данных.
Если коллективные операции вызываются внутри функции, то
они работают на всех глобальных устройствах, и это выглядит,
как если бы трансформация pmap() работала с массивом, сегмен
тированным по нескольким хостам. Каждый хост видит и обраба
тывает только собственный сегмент, но все хосты обмениваются
информацией с помощью коллективных операций.
8
Использование
сегментирования
тензоров
Темы главы:
использование сегментирования тензоров для
организации распараллеливания с помощью XLA;
реализация распараллеливания по данным и по тензорам
для тренировки нейронных сетей.
В этой главе представлен альтернативный и более новый способ рас
параллеливания вычислений в JAX с использованием сегментиро
вания тензоров (tensor sharding). Вариант использования тот же, что
и в предыдущей главе: выполнение отдельных частей вычисления
в параллельном режиме и ускоренное выполнение всего вычисления
в целом. Такой подход особенно полезен для различных способов
распараллеливания процесса тренировки нейронных сетей вне за
висимости от того, применяется распараллеливание по данным или
распараллеливание модели. Такую методику также можно применять
для (логического) вывода в больших моделях, для которых недоста
точно одного GPU. Но эта новейшая методика способна обеспечить
преимущества не только в области глубокого обучения, но и в дру
гих сферах деятельности. Если вы работаете с большими тензорами
в биоинформатике, космологии, моделировании погоды или в ка
306
Глава 8
Использование сегментирования тензоров
кой-либо другой отрасли науки, то сегментирование тензоров может
предоставить простой способ распараллеливания вычислений.
Распараллеливание с использованием pmap(), описанное в преды
дущей главе, предоставляет возможность явно сообщить компиля
тору, что необходимо сделать, используя код, предназначенный для
каждого устройства, и явно вызываемые коллективные операции
для обмена информацией. Другой подход позволяет компилятору
автоматически распределять функции по устройствам без указания
слишком большого количества подробностей низкого уровня. Сег
ментирование тензоров (или применение распределенных масси
вов) относится ко второму подходу. Такой вариант распараллелива
ния вычислений стал доступным с версии 0.4.1 JAX вместе с новым
типом jax.Array.
В приложении D описаны две экспериментальные методики рас
параллеливания, которые в настоящее время не являются актуаль
ными, а именно xmap() и pjit(). Возможно, вы пожелаете ознако
миться с этими темами либо потому, что интересуетесь историей
разработки механизма распараллеливания в JAX, либо потому, что
вам необходимо понимать и поддерживать код, использующий эти
методики. Кроме того, изучение xmap() и pjit() помогает лучше по
нять сегментирование тензоров, хотя знание этих методик не требу
ется и изучения текущей главы вполне достаточно.
Мы уже знакомы с типом jax.Array, который рассматривался
в главе 3. Здесь мы будем изучать этот тип еще подробнее. В сово
купности с jit() jax.Array обеспечивает автоматическое распарал
леливание на основе компилятора. Я называю это неявным распа
раллеливанием, поскольку при этом не используются какие-либо
специальные конструкции языка или библиотечные функции для
явного распараллеливания кода, а вместо этого вы размещаете тен
зоры на устройствах таким способом, чтобы вычисления выполня
лись в параллельном режиме на различных устройствах.
Это простая идея. В JAX вычисления логически вытекают из раз
мещения данных. Поэтому сегментирование тензоров можно рас
сматривать как расширение методики, описанной в подразделе
3.2.3, где тензоры размещались на различных устройствах (CPU,
GPU, TPU), а вычисления выполнялись именно там, где размещены
тензоры. При сегментировании тензоров происходит почти то же
самое, но теперь можно разделять тензоры и распределять их сег
менты по различным устройствам. Во время вычисления JAX при
нимает на себя ответственность за то, как наиболее эффективно вы
полнить вычисление без ненужных перемещений данных.
Тип jax.Array – это универсальный тип массива, объединяющий
(и заменяющий) типы DeviceArray, ShardedDeviceArray и GlobalDeviceArray из предыдущих версий JAX. Новый тип jax.Array помо
гает сделать механизм распараллеливания одним из главных функ
Основы сегментирования тензоров
307
циональных средств JAX. Он упрощает и унифицирует внутренние
механизмы JAX и позволяет объединить jit() и pjit() (именно по
этому вас может заинтересовать содержимое приложения D).
Функции, созданные с помощью jit(), могут работать с распре
деленными массивами (массивами, сегментированными по раз
личным устройствам) без копирования данных на одно устройство.
Компилятор определяет сегментирование для промежуточных пере
менных на основе сегментирования входных данных. Также можно
воздействовать на сегментирование промежуточных переменных,
устанавливая ограничения. Компилятор вставляет коллективные
операции там, где они необходимы. Возможно, вы слышали о такой
же особенности pjit(), если ранее работали с этой трансформацией,
а теперь аналогичное поведение обобщено и в jit().
ПРИМЕЧАНИЕ Функциональные возможности, требуемые
для jax.Array, не поддерживаются механизмом времени вы
полнения Colab TPU, даже если вы используете более старую
версию JAX, поддерживающую Colab TPU. Потребуется на
стройка Colab TPU, описанная в приложении C.
Сначала рассмотрим полный пример вычисления хорошо знако
мого нам скалярного произведения, затем обсудим все заложенные
в основу концепции и в завершение реализуем распараллеливание
процесса тренировки нейронной сети с применением новой мето
дики.
8.1
Основы сегментирования тензоров
Перепишем наш старый добрый пример вычисления скалярного
произведения для использования сегментирования тензоров. В лис
тинге 8.1 массивы сегментируются перед вычислениями, и соответ
ствующий измененный код выделен полужирным шрифтом в лис
тинге. Затем все вычисления выполняются в параллельном режиме.
В код, связанный с самими вычислениями, не внесено никаких из
менений.
Листинг 8.1 Распараллеливание вычислений скалярного
произведения с использованием распределенных
массивов
from jax.experimental import mesh_utils
from jax.sharding import PositionalSharding
from jax import random
❶
❶
def dot(v1, v2):
❷
Глава 8
308
Использование сегментирования тензоров
return jnp.vdot(v1, v2)
rng_key = random.PRNGKey(42)
vs = random.normal(rng_key, shape=(8_000,10_000))
v1s = vs[:4_000,:]
v2s = vs[4_000:,:]
v1s.shape, v2s.shape
❸
❸
❸
>>> (4000, 10000), (4000, 10000))
jax.debug.visualize_array_sharding(v1s)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
┌────────────────────────────────────────────────────────────┐
│
│
│
│
│
│
│
│
│
TPU 0
│
│
│
│
│
│
│
│
│
└────────────────────────────────────────────────────────────┘
❹
sharding = PositionalSharding(mesh_utils.create_device_mesh((8,1))) ❺
sharding
>>> PositionalSharding([[{TPU 0}]
>>>
[{TPU 1}]
>>>
[{TPU 2}]
>>>
[{TPU 3}]
>>>
[{TPU 6}]
>>>
[{TPU 7}]
>>>
[{TPU 4}]
>>>
[{TPU 5}]])
v1sp = jax.device_put(v1s, sharding)
v2sp = jax.device_put(v2s, sharding)
type(v1sp)
>>> jaxlib.xla_extension.ArrayImpl
jax.debug.visualize_array_sharding(v1sp)
>>>
>>>
>>>
>>>
>>>
>>>
┌────────────────────────────────────────────────────────────┐
│
TPU 0
│
├────────────────────────────────────────────────────────────┤
│
TPU 1
│
├────────────────────────────────────────────────────────────┤
│
TPU 2
│
❻
❻
❻
❻
❻
❻
❻
❻
❼
❼
❽
❾
Основы сегментирования тензоров
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
309
├────────────────────────────────────────────────────────────┤
│
TPU 3
│
├────────────────────────────────────────────────────────────┤
│
TPU 6
│
├────────────────────────────────────────────────────────────┤
│
TPU 7
│
├────────────────────────────────────────────────────────────┤
│
TPU 4
│
├────────────────────────────────────────────────────────────┤
│
TPU 5
│
└────────────────────────────────────────────────────────────┘
d = jax.vmap(dot)(v1sp, v2sp)
d.shape
❿
>>> (4000,)
jax.debug.visualize_array_sharding(d)
⓫
⓬
>>> ┌───────┬───────┬───────┬───────┬───────┬───────┬───────┬───────┐
>>> │ TPU 0 │ TPU 1 │ TPU 2 │ TPU 3 │ TPU 6 │ TPU 7 │ TPU 4 │ TPU 5 │
>>> └───────┴───────┴───────┴───────┴───────┴───────┴───────┴───────┘
❶ Импорт модулей для поддержки сегментирования.
❷ Знакомая функция вычисления скалярного произведения двух векторов.
❸ В этот раз мы генерируем более широкие векторы, чем раньше.
❹ Визуализация сегментирования: весь тензор размещается на одном устройстве.
❺ Создание двумерной сетки для сегментирования первой оси тензора второго
ранга по восьми TPU.
❻ Функция create_device_mesh() возвращает самый производительный порядок
устройств для заданной формы.
❼ Сегментирование тензоров по их первым осям.
❽ Проверка типа.
❾ Проверка выполнения сегментирования тензора по нескольким устройствам.
❿ Выполнение вычислений.
⓫ Получена ожидаемая форма.
⓬ Результат также распределен по нескольким устройствам.
Приведенный выше пример содержит много различных особен
ностей. Рассмотрим самые важные его части.
Во-первых, мы создали более широкую, чем обычно, версию мас
сивов с векторами. Теперь каждый вектор содержит 10 000 элемен
тов. Это сделано для следующего примера, в котором мы продемон
стрируем, как легко сегментировать вычисления для распределения
их по многим измерениям. В саму функцию не вносились никакие
изменения. Мы воспользовались функцией jax.debug.visualize_array_sharding() для визуального представления схемы размещения
тензора. Изначально весь тензор размещался на одном TPU.
310
8.1.1
Глава 8
Использование сегментирования тензоров
Сетка устройств
Мы создаем сетку устройств (device mesh) – n-мерный массив
устройств, сформированный функцией mesh_utils.create_device_
mesh(). Эта функция возвращает наиболее производительный по
рядок устройств заданной формы. Это важно, так как аппаратные
устройства обычно скомпонованы по некоторой топологической
схеме (например, 2D- или 3D-тор) и соединены не полностью. Толь
ко между соседними элементами существует некоторое высокоско
ростное соединение.
В приведенном выше примере мы создали двумерную (2D) сетку
размером (8, 1) для сегментирования первой оси тензора ранга 2 по
восьми TPU. Здесь можно видеть, что устройства расположены не
в порядке от единицы до восьми, а по более сложной схеме из сооб
ражений производительности.
Если вы эмулируете систему с несколькими устройствами на CPU
с применением флага XLA --xla_force_host_platform_device_count
в соответствии с инструкцией, приведенной в листинге 7.2, то не
обходимо добавить параметр устройств в функцию create_device_
mesh(). Вызов должен выглядеть следующим образом:
>>> mesh_utils.create_device_mesh((8, 1), devices=jax.devices("cpu"))
Если требуется сегментировать только второе измерение тензора,
то используется сетка устройств с формой (1, 8) вместо (8, 1).
ПРИМЕЧАНИЕ Форма тензора, необходимая для сегменти
рования, должна совпадать с формой сегментирования. Это
означает, что обе формы обязательно должны иметь равную
длину (одинаковое количество осей) – именно поэтому мы
использовали сетку устройств размером (8, 1), а не просто
(8, ). Кроме того, количество элементов по каждой оси тензо
ра должно быть кратным (делиться без остатка) размеру со
ответствующей оси сегментирования.
8.1.2
Позиционное сегментирование
Далее создается объект PositionalSharding, представляющий схему
распределенной памяти. Этот объект фиксирует порядок устройств
и начальную форму.
Мы используем уже знакомую функцию jax.device_put(), кото
рая весьма интенсивно применялась в главе 3, для передачи данных
на устройство. Функция принимает сегментированный объект вмес
то конкретного устройства и размещает данные соответствующим
образом. Можно проверить тот факт, что размещение тензора изме
няется и его первая ось (с индексом 0) сегментирована по первой
311
Основы сегментирования тензоров
оси сетки устройств. Вторая ось сегментируемого объекта имеет раз
мер 1, поэтому вторая ось тензора не разделяется.
Затем выполняются вычисления с помощью функции, векторизо
ванной с использованием vmap() в обычном стиле, и мы получаем
результат, который также сегментирован. Это объясняется тем, что
на каждом устройстве имеется подмножество векторов, и только для
этого подмножества существует возможность вычисления скаляр
ного произведения. Вычисленные скалярные произведения хранят
ся на тех же устройствах, что и исходные векторы.
Проверить, действительно ли происходит распараллеливание,
можно посредством измерения времени, потребляемого вычисли
тельной функцией.
Листинг 8.2 Измерение времени вычислений с использованием
и без использования сегментирования
%timeit jax.vmap(dot)(v1sp, v2sp).block_until_ready()
❶
>>> 1.7 ms ± 34.3 µs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
%timeit jax.vmap(dot)(v1s, v2s).block_until_ready()
❷
>>> 2.22 ms ± 27.2 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶ Использование сегментированного тензора.
❷ Использование тензора без сегментирования.
Здесь мы проверили одну и ту же функцию, работающую с сег
ментированными тензорами и тензорами без сегментирования.
Сегментированные тензоры распределяются по восьми устрой
ствам, а тензор без сегментирования размещается на одном устрой
стве. Поэтому при первом запуске вычисления выполняются на
всех восьми ядрах TPU, а при втором задействуется только одно
ядро TPU (с индексом 0). В первом случае вычисления выполняются
быстрее. Но почему не в восемь раз? Потому что все доступные ап
паратные средства используются не на полную мощность. Оптими
зация производительности – это отдельная, весьма интересная тема
сама по себе, и модель Roofline (https://crd.lbl.gov/assets/pubs_presos/
parlab08-roofline-talk.pdf) позволяет обнаружить многочисленные
сложности в этой области.
8.1.3
Пример с применением двумерной сетки
В рассматриваемом здесь примере векторы состоят из многих ком
понентов – в данном конкретном случае из 10 000. Обработка век
торов такого размера вполне по силам каждому отдельному устрой
ству, поэтому сегментирование по измерению вектора также может
иметь смысл. Скалярное произведение можно с легкостью сегмен
312
Глава 8
Использование сегментирования тензоров
тировать, поскольку оно представляет собой всего лишь сумму про
изведений соответствующих элементов векторов. На рис. 8.1 пока
зана наглядная схема этой процедуры.
Рис. 8.1 Сегментирование скалярного произведения
Для сегментирования по двум измерениям (самих векторов и их
компонентов) необходимо подготовить двумерную сетку. В приве
денном ниже примере сегментируются два входных тензора ранга 2
по обоим измерениям, а полученный в результате тензор – по свое
му единственному измерению. Процесс может показаться достаточ
но сложным, поэтому начнем с рассмотрения его наглядной схемы,
показанной на рис. 8.2.
Внимательно разберемся во всем, что происходит здесь. В пер
вую очередь мы создаем 4000 пар векторов, каждый из которых со
держит 10 000 элементов. Это гораздо более широкие векторы по
сравнению с используемыми в предыдущих главах. Итак, мы име
ем два двумерных массива размером (4000, 10 000). Воспользуемся
двумерной сеткой аппаратных устройств размером (2, 4). Если бы
мы использовали сегментирование с именованием (о нем немного
позже в текущем разделе), то могли бы получить именованные оси
«x» и «y», или «vectors» и «features».
Оба входных параметра функции dot() разделены по первому
и второму измерениям. Первое измерение входных массивов (раз
мер 4000) распределено по первой оси сетки устройств (размер 2)
с получением фрагментов, имеющих размер 2000 – индекс под
множеств векторов. Второе измерение входных массивов (размер
10 000) распределено по второй оси сетки устройств (размер 4) с по
лучением фрагментов, имеющих размер 2500 – индекс подмножеств
компонентов векторов.
Вычисленный результат распределяется только по одной оси, так
как является тензором ранга 1. Здесь эта ось равнозначна первой оси
входных массивов, по которой считается количество векторов.
Вычисление выполняется следующим способом: каждое устрой
ство может вычислить частичное скалярное произведение на основе
имеющихся в его распоряжении сегментов. Таким образом, на каж
Основы сегментирования тензоров
Случайно
сгенерированный
тензор
Первый набор
векторов
Второй набор
векторов
Оба набора векторов
сегментируются
в соответствии с сеткой
устройств. Каждый сегмент
имеет форму (2000, 2500)
Распределение каждой
пары сегментов по
устройствам TPU.
Каждое устройство TPU
содержит два сегмента
с формой (2000, 2500)
Сетка устройств
Каждое устройство TPU
вычисляет частичное
скалярное произведение.
Форма результата
(2000, 1)
Сетка устройств
Сохранение сегментирования по первому
измерению, но удаление второго
измерения посредством сложения всех
частичных результатов по нему
Теперь не существует частичных скалярных
произведений, и мы получаем два сегмента
для вычисления полных скалярных
произведений, каждое с формой (2000, )
Объединение двух сегментов в конечный
результат скалярных произведений.
Форма результата (4000, )
Рис. 8.2 Вычисление сегментированного скалярного произведения
по двумерной сетке
313
Глава 8
314
Использование сегментирования тензоров
дом устройстве фрагмент с 2500 элементами вектора, состоящего
из 10 000 элементов, из первого массива поэлементно умножается
на фрагмент с 2500 элементами вектора, состоящего из 10 000 эле
ментов, из второго массива, и это выполняется с каждой из 2000 пар
векторов, размещенной на одном конкретном устройстве. Каждое
частичное скалярное произведение выдает одно число, а затем не
обходимо просуммировать все частичные скалярные произведения
в каждом векторе (всего получается четыре таких частичных скаляр
ных произведения). Это делается с помощью коллективной опера
ции, скрытой от нас.
В конце процедуры первая строка сетки устройств содержит ска
лярное произведение первых 2000 векторов из обоих массивов. Во
второй строке сетки устройств находится скалярное произведение
второй половины векторов из обоих массивов.
Для завершения вычисления мы просто объединяем два получен
ных сегмента, получая конечный результат из 4000 скалярных про
изведений. Описанную выше схему достаточно легко преобразовать
в код.
Листинг 8.3
Сегментирование по двум измерениям
from jax.sharding import PartitionSpec as P
from jax.sharding import Mesh
import numpy as np
rng_key = random.PRNGKey(42)
vs = random.normal(rng_key, shape=(8_000,10_000))
v1s = vs[:4_000,:]
v2s = vs[4_000:,:]
v1s.shape, v2s.shape
❶
❶
>>> (4000, 10000), (4000, 10000))
sharding = PositionalSharding(
mesh_utils.create_device_mesh((2,4)))
v1sp = jax.device_put(v1s, sharding)
v2sp = jax.device_put(v2s, sharding)
jax.debug.visualize_array_sharding(v1sp)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
❷
❸
❸
❹
┌─────────────┬─────────────┬─────────────┬─────────────┐
│
│
│
│
│
│
TPU 0
│
TPU 1
│
TPU 2
│
TPU 3
│
│
│
│
│
│
│
│
│
│
│
├─────────────┼─────────────┼─────────────┼─────────────┤
│
│
│
│
│
315
Основы сегментирования тензоров
>>>
>>>
>>>
>>>
│
TPU 6
│
TPU 7
│
TPU 4
│
TPU 5
│
│
│
│
│
│
│
│
│
│
│
└─────────────┴─────────────┴─────────────┴─────────────┘
d = jax.vmap(dot)(v1sp, v2sp)
d.shape
❺
>>> (4000,)
jax.debug.visualize_array_sharding(d)
>>> ┌───────────┬───────────┐
>>> │TPU 0,1,2,3│TPU 4,5,6,7│
>>> └───────────┴───────────┘
❶
❷
❸
❹
❺
❻
❻
Распределение тензоров.
Создание схемы сегментирования для двумерной сетки.
Распределение тензоров.
Визуальное представление сегментирования тензоров.
Применение функции.
Визуальное представление сегментирования полученного результата.
В приведенном выше примере мы создали двумерную сетку
устройств и распределили по ней исходные тензоры. Тензоры сег
ментируются по обоим измерениям сетки: по первой оси с разде
лением всех векторов на две группы и по второй оси, содержащей
элементы векторов.
Вычисление скалярного произведения выполнено успешно, хотя
теперь и результат отдельно сегментируется. Например, для полу
чения скалярного произведения первой пары векторов (с индексом
0) необходимо взять соответствующие сегменты из TPU 0 и вычис
лить на этом устройстве частичное скалярное произведение, взять
сегменты из TPU 1 и вычислить другое частичное скалярное произ
ведение, а также повторить эту процедуру для TPU 2 и 3. Затем вы
численные четыре частичных скалярных произведения суммиру
ются (с использованием коллективной операции), чтобы получить
итоговое скалярное произведение. Аналогичным способом вычис
ляются такие же скалярные произведения на устройствах TPU 0, 1,
2 и 3.
И снова можно проверить, действительно ли произошло распа
раллеливание, измеряя время, потребляемое вычислительной функ
цией.
Листинг 8.4 Измерение времени выполнения вычислений при двумерном
сегментировании
%timeit jax.vmap(dot)(v1sp, v2sp).block_until_ready()
❶
>>> 1.61 ms ± 39.4 µs per loop (mean ± std. dev. of 7 runs, 1,000 loops each)
Глава 8
316
Использование сегментирования тензоров
%timeit jax.vmap(dot)(v1s, v2s).block_until_ready()
❷
>>> 2.32 ms ± 15.7 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶ Использование сегментированного тензора.
❷ Использование тензора без сегментирования.
В рассматриваемом здесь примере мы получаем такое же увели
чение скорости, как и в варианте с одномерным сегментировани
ем, что вполне ожидаемо, поскольку различие заключается всего
лишь в нескольких операциях суммирования. Если посмотреть на
представление кода HLO для этой функции (см. соответствующий
блокнот (Notebook) в репозитории книги), то можно заметить, что
коллективная операция суммирования с накоплением (all-reduce)
автоматически размещена в коде для организации обмена инфор
мацией между группами.
8.1.4
Использование репликации
Иногда не требуется сегментирование тензора по всем измерени
ям. В этом случае можно использовать метод сегментирования
replicate(axis=NUMBER) для копирования срезов тензора на каждое
устройство по заданному измерению. Если ось не задана, то репли
кация выполняется по каждой оси. Если мы решаем не использовать
второе измерение сегментирования для разделения исходных тен
зоров, то можно использовать код, показанный в листинге 8.5.
Листинг 8.5
Использование репликации
sharding = PositionalSharding(
mesh_utils.create_device_mesh((2,4)))
v1sp = jax.device_put(v1s, sharding.replicate(axis=1))
jax.debug.visualize_array_sharding(v1sp)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
❶
❷
❸
┌────────────────────────────────────────────────────────────┐
│
│
│
TPU 0,1,2,3
│
│
│
│
│
├────────────────────────────────────────────────────────────┤
│
│
│
TPU 4,5,6,7
│
│
│
│
│
└────────────────────────────────────────────────────────────┘
❶ Продолжается использование двумерной сетки устройств.
❷ Распределение тензора по первому измерению, но по второму выполняется ре-
пликация.
❸ Визуальное представление сегментации тензора.
317
Основы сегментирования тензоров
В приведенном выше примере мы сообщили JAX о том, что сег
ментирование исходного тензора по второму измерению (с индек
сом 1) не требуется. Устройства по этому измерению сегментиро
вания будут содержать копии тензоров, размещенных по второму
измерению тензора, поэтому фактически каждое устройство в груп
пе (группы: {0, 1, 2, 3} и {4, 5, 6, 7}) получает полное содержимое вто
рого измерения тензора. TPU с 0 по 3 (то же относится и ко второй
группе TPU) содержат одинаковые данные в своей памяти.
Рассмотрим измерения сегментирования: первоначальное сег
ментирование имеет форму (2, 4), а после вызова replicate() форма
меняется на (2, 1). В отличие от методов NumPy, которые уплотняют
измерение с размером 1, и форма приобретает вид (2, ), сокращен
ная ось не уплотняется. Для управления этим поведением можно ис
пользовать дополнительный параметр keepdims=False.
Это вполне естественный способ реализации умножения распре
деленных матриц. Предположим, что умножаются две матрицы:
одна расположена слева, вторая – справа. Для левой матрицы необ
ходимо реплицировать строки так, чтобы каждое устройство полу
чило полную копию некоторой строки. Для правой матрицы требу
ется репликация столбцов для размещения полной копии столбца на
устройстве. При применении функции dot() по таким сегментиро
ванным входным данным мы получаем результат умножения мат
риц, вычисленный в параллельном режиме на нескольких (многих)
устройствах.
Листинг 8.6 Пример умножения распределенных матриц
sharding = PositionalSharding(mesh_utils.create_device_mesh((2,4)))
A = random.normal(rng_key, shape=(10000,2000))
B = random.normal(rng_key, shape=(2000,5000))
Ad = jax.device_put(A, sharding.replicate(1))
Bd = jax.device_put(B, sharding.replicate(0))
jax.debug.visualize_array_sharding(Ad)
jax.debug.visualize_array_sharding(Bd)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
┌───────────┐
│
│
│TPU 0,1,2,3│
│
│
│
│
├───────────┤
│
│
│TPU 4,5,6,7│
│
│
│
│
└───────────┘
❶
❶
❷
❸
❹
❹
318
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
Глава 8
Использование сегментирования тензоров
┌─────────────┬─────────────┬─────────────┬─────────────┐
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
TPU 0,6
│
TPU 1,7
│
TPU 2,4
│
TPU 3,5
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
│
└─────────────┴─────────────┴─────────────┴─────────────┘
Cd = jnp.dot(Ad, Bd)
jax.debug.visualize_array_sharding(Cd)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
┌───────┬───────┬───────┬───────┐
│
│
│
│
│
│ TPU 0 │ TPU 1 │ TPU 2 │ TPU 3 │
│
│
│
│
│
│
│
│
│
│
├───────┼───────┼───────┼───────┤
│
│
│
│
│
│ TPU 6 │ TPU 7 │ TPU 4 │ TPU 5 │
│
│
│
│
│
│
│
│
│
│
└───────┴───────┴───────┴───────┘
%timeit jnp.dot(Ad, Bd).block_until_ready()
❺
❻
❼
>>> 2.12 ms ± 58.8 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit jnp.dot(A, B).block_until_ready()
❽
>>> 10.4 ms ± 10.6 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
%timeit (A@B).block_until_ready()
❾
>>> 10.4 ms ± 21.1 µs per loop (mean ± std. dev. of 7 runs, 100 loops each)
❶
❷
❸
❹
❺
❻
❼
❽
❾
Генерация двух матриц со случайными элементами.
Репликация левой матрицы по строкам.
Репликация правой матрицы по столбцам.
Визуальное представление сегментирования матриц.
Вычисление произведения матриц с использованием функции dot().
Визуальное представление сегментирования вычисленного результата.
Проверка наличия ускорения вычислений благодаря выполнению в параллельном режиме.
Использование несегментированных матриц.
Использование инфиксного оператора Python @ для умножения матриц.
В приведенном выше примере мы организовали данные так, что
бы вычисление скалярного произведения эффективно выполнялось
в параллельном режиме. Мы подготовили сегментированные вход
Основы сегментирования тензоров
319
ные матрицы, выполнили умножение этих матриц с использовани
ем операции скалярного произведения с сегментированными мас
сивами и сравнили полученный результат с результатом обычного
умножения матриц для несегментированных массивов. Мы получи
ли точно такой же результат, как и для обычного умножения матриц,
хотя в нашем учебном примере было получено увеличение скорости
приблизительно в четыре раза.
8.1.5
Ограничения сегментирования
В дополнение к управлению сегментированием входных перемен
ных (которое осуществляется с использованием функции jax.device_put()) можно передать компилятору информацию о том, как
сегментировать промежуточные переменные функции. Такая воз
можность может оказаться полезной, например, если вы не управляе
те входными переменными. Для этого используется функция jax.
lax.with_sharding_constraint(), во многом похожая на функцию
jax.device_put(), но применяемая внутри JIT-трансформируемых
функций. Например, можно реализовать функцию умножения мат
риц, которая принудительно сегментирует свой аргумент так, чтобы
умножение стало распараллеленным.
Листинг 8.7
Использование ограничений сегментирования
from jax import jit
from functools import partial
@partial(jax.jit, static_argnums=2)
def distributed_mul(a, b, sharding):
ad = jax.lax.with_sharding_constraint(a, sharding.replicate(1))
bd = jax.lax.with_sharding_constraint(b, sharding.replicate(0))
return jnp.dot(ad, bd)
❶
❷
❷
sharding = PositionalSharding(mesh_utils.create_device_mesh((2,4)))
jax.debug.visualize_array_sharding(A)
jax.debug.visualize_array_sharding(B)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
┌───────┐
│
│
│
│
│
│
│
│
│ TPU 0 │
│
│
│
│
│
│
│
│
└───────┘
❸
❸
Глава 8
320
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
Использование сегментирования тензоров
┌────────────────────────────────────────────────────────────┐
│
│
│
│
│
│
│
│
│
TPU 0
│
│
│
│
│
│
│
│
│
└────────────────────────────────────────────────────────────┘
d = distributed_mul(A, B, sharding)
jax.debug.visualize_array_sharding(d)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
┌───────┬───────┬───────┬───────┐
│
│
│
│
│
│ TPU 0 │ TPU 1 │ TPU 2 │ TPU 3 │
│
│
│
│
│
│
│
│
│
│
├───────┼───────┼───────┼───────┤
│
│
│
│
│
│ TPU 6 │ TPU 7 │ TPU 4 │ TPU 5 │
│
│
│
│
│
│
│
│
│
│
└───────┴───────┴───────┴───────┘
❹
❶ JIT-аннотированная функция с сегментированием в качестве статического аргу-
мента.
❷ Принудительно заданный способ сегментирования.
❸ Первоначальные входные аргументы не сегментированы.
❹ Вычисленный результат сегментирован.
В приведенном выше примере мы аннотировали вычисления, что
бы сообщить компилятору о том, как должны быть сегментированы
входные аргументы функции. Это тот же паттерн сегментирования,
который использовался ранее для организации распределенного
умножения с применением операции скалярного произведения.
Сегментирование полученного в результате значения свидетель
ствует о том, что ограничения были применены.
8.1.6
Сегментирование с именованием
Также можно использовать сегментирование с именованием для
выражения именованных сегментов вместо позиций в массиве, как
при позиционном сегментировании. Такой подход аналогичен при
меняемому в xmap() и pjit() (подробности см. в приложении D).
Подобная методика реализуется посредством определения ме
неджера контекста сетки аппаратных устройств – n-мерного масси
321
Основы сегментирования тензоров
ва устройств с именованными осями. С технической точки зрения
сетка аппаратных устройств представляет собой двухкомпонентный
объект:
n-мерный массив объектов JAX-устройств – те же самые объ
екты, которые мы получаем при вызове функции jax.devices()
или jax.local_devices(), но представленные как тип np.array.
Будьте внимательны: это чистый массив NumPy (np.array), а не
массив JAX NumPy (jnp.array), поскольку объект устройства не
является допустимым корректным типом массива JAX. Вспомо
гательная функция create_device_mesh() создает выполняемую
надлежащим образом сетку устройств с хорошей коллективной
производительностью;
кортеж ресурсов имен осей – длина кортежа обязательно долж
на совпадать с рангом массива устройств. Например, для трех
мерной сетки кортеж содержит три ресурса имен осей.
JAX предоставляет специализированный менеджер контекста Mesh
(https://docs.jax.dev/en/latest/jax.sharding.html#jax.sharding.Mesh).
Создать объект Mesh можно следующим образом:
from jax.sharding import Mesh
mesh = Mesh(mesh_utils.create_device_mesh((4,2)),
axis_names=('batch', 'features'))
❶
❷
❶ Импорт типа Mesh.
❷ Создание двумерной сетки с двумя именованными осями: «batch» и «features».
Здесь мы создали двумерный массив доступных устройств и ме
неджер контекста Mesh с двумя осями по сетке аппаратных устройств.
Первой оси сетки присвоено имя «batch», вторая получила название
«features». Перепишем пример сегментирования тензора с исполь
зованием этих именованных осей.
Листинг 8.8
Использование сегментирования с именованием
from jax.sharding import Mesh
from jax.sharding import PartitionSpec as P
from jax.sharding import NamedSharding
mesh = Mesh(mesh_utils.create_device_mesh((4,2)),
axis_names=('batch', 'features'))
sharding = NamedSharding(mesh, P('batch', 'features'))
v1sp = jax.device_put(v1s, sharding)
v2sp = jax.device_put(v2s, sharding)
jax.debug.visualize_array_sharding(v1sp)
>>> ┌─────────────────────────────┬─────────────────────────────┐
>>> │
TPU 0
│
TPU 1
│
❶
❶
❶
❷
❸
❹
❹
❺
Глава 8
322
>>>
>>>
>>>
>>>
>>>
>>>
>>>
Использование сегментирования тензоров
├─────────────────────────────┼─────────────────────────────┤
│
TPU 2
│
TPU 3
│
├─────────────────────────────┼─────────────────────────────┤
│
TPU 6
│
TPU 7
│
├─────────────────────────────┼─────────────────────────────┤
│
TPU 4
│
TPU 5
│
└─────────────────────────────┴─────────────────────────────┘
d = jax.vmap(dot)(v1sp, v2sp)
d.shape
❻
>>> (4000,)
❶
❷
❸
❹
❺
❻
Импорт всех необходимых модулей.
Создание двумерной сетки с двумя именованными осями.
Создание именованного сегментирования с двумя именованными осями.
Распределение тензоров.
Визуальное представление сегментирования тензоров.
Вычисление скалярного произведения и проверка формы результата.
Мы создали объект сегментирования типа NamedSharding со специ
фикацией разделения на части. PartitionSpec – это кортеж, элемен
тами которого могут быть значение None, строковое имя оси сетки
или кортеж имен осей сетки. Каждый элемент кортежа описывает
измерение сетки, по которому размещается измерение входного
тензора.
В рассматриваемом здесь примере тензор ранга 2 разделяется по
двум осям. Первое измерение распределено по оси batch сетки, вто
рое – по оси features.
Значения None используются для определения осей, по которым
сегментирование не требуется. По таким осям данные будут репли
цироваться. Значения None можно не указывать для осей, распола
гающихся в конце списка. Поэтому для репликации данных по вто
рой оси допустима следующая форма записи: NamedSharding(mesh,
P('batch', None)) или просто NamedSharding(mesh, P('batch')).
8.1.7
Стратегия размещения устройств и ошибки
Для сегментированных данных применяется обобщенная стратегия
JAX явного размещения устройств, которая описывалась в подраз
деле 3.2.3. Напомню суть этой стратегии: вычисления, связанные
с зафиксированными данными, выполняются на специально выде
ленном для них и зафиксированном устройстве, и результаты также
будут зафиксированы на том же устройстве. Использование функ
ции device_put() с сегментированными данными создает тензоры,
зафиксированные на некотором конкретном устройстве.
Ошибка возникает, если вы вызываете операцию с аргументами,
зафиксированными на разных устройствах (но ошибки не будет,
323
Основы сегментирования тензоров
если некоторые аргументы не являются зафиксированными). При
использовании сегментирования, если два аргумента вычисления
размещены на различных устройствах или порядок соответствую
щих им устройств не является совместимым, то возникает ошибка.
Листинг 8.9 Попытка использования данных, сегментированных
на различных устройствах
sharding_a = PositionalSharding(
np.array(jax.devices()[:4]).reshape(4,1))
sharding_b = PositionalSharding(
np.array(jax.devices()[4:]).reshape(4,1))
sharding_a
❶
❷
>>> PositionalSharding([[{TPU 0}]
>>>
[{TPU 1}]
>>>
[{TPU 2}]
>>>
[{TPU 3}]])
sharding_b
>>> PositionalSharding([[{TPU 4}]
>>>
[{TPU 5}]
>>>
[{TPU 6}]
>>>
[{TPU 7}]])
v1sp = jax.device_put(v1s, sharding_a)
v2sp = jax.device_put(v2s, sharding_b)
d = jax.vmap(dot)(v1sp, v2sp)
>>> …
>>> ValueError: Received incompatible devices for
jitted computation. Got ARG_SHARDING with device ids
[0, 1, 2, 3] on platform TPU and ARG_SHARDING with
device ids [4, 5, 6, 7] on platform TPU
# ValueError: приняты несовместимые устройства для
# JIT-трансформированного вычисления. Получен
# ARG_SHARDING с идентификаторами устройств
# [0, 1, 2, 3] на платформе TPU и ARG_SHARDING
# с идентификаторами устройств [4, 5, 6, 7]
# на платформе TPU
❶
❷
❸
❹
❸
❸
❹
Использование первых четырех устройств.
Использование вторых четырех устройств.
Сегментирование тензоров.
Выводится сообщение об ошибке: входные данные размещены в различных спис
ках устройств.
Мы распределили первый тензор по первым четырем устрой
ствам, а второй – по следующим четырем устройствам, поэтому
Глава 8
324
Использование сегментирования тензоров
они не пересекаются. Затем функция вычисления скалярного про
изведения, которая раньше работала нормально, выдает сообщение
об ошибке, оповещающее о том, что входные данные размещены
в различных списках устройств. Порядок в списке устройств важен,
поэтому возникнет та же ошибка, даже если вы используете те же
устройства для обоих тензоров, но в другом порядке.
Если вы не фиксируете явно тензор на устройстве, то он размеща
ется без фиксации на устройстве по умолчанию (TPU 0 при работе
в среде Cloud TPU). Незафиксированный тензор может быть автома
тически перемещен и сегментирован по-другому, поэтому безопас
ным способом использования незафиксированных тензоров вместе
с зафиксированными является их передача как аргументов вычис
ления.
Листинг 8.10 Использование зафиксированных
и незафиксированных аргументов
sharding_a = PositionalSharding(np.array(jax.devices()).reshape(8,1))❶
sharding_a
>>> (PositionalSharding([[{TPU 0}]
>>>
[{TPU 1}]
>>>
[{TPU 2}]
>>>
[{TPU 3}]
>>>
[{TPU 4}]
>>>
[{TPU 5}]
>>>
[{TPU 6}]
>>>
[{TPU 7}]]),
v1sp = jax.device_put(v1s, sharding_a)
d = jax.vmap(dot)(v1sp, v2s)
jax.debug.visualize_array_sharding(d)
❷
❸
❹
>>> ┌───────┬───────┬───────┬───────┬───────┬───────┬───────┬───────┐
>>> │ TPU 0 │ TPU 1 │ TPU 2 │ TPU 3 │ TPU 4 │ TPU 5 │ TPU 6 │ TPU 7 │
>>> └───────┴───────┴───────┴───────┴───────┴───────┴───────┴───────┘
❶
❷
❸
❹
Сегментирование использует все восемь устройств.
Сегментирование первого тензора.
Использование vmap для сегментированных и несегментированных данных.
Визуальное представление сегментирования итогового результата, чтобы наглядно показать, что он сегментирован по тому же списку устройств, что и первый
аргумент.
В приведенном выше примере сегментируется первый аргумент,
и его сегменты фиксируются на конкретных устройствах, но ничего
Многослойный перцептрон с применением сегментирования тензоров
325
не делается со вторым аргументом. В этом случае вычисления вы
полнены успешно, и итоговый результат сегментирован по тому же
списку устройств, что и первый аргумент.
Теперь вы в основном понимаете, как работает сегментирование
тензоров, и мы можем применить полученные знания к нашему
«подопытному кролику» – классификации MNIST с использованием
многослойного перцептрона (MLP), как это делалось ранее много
кратно.
8.2
Многослойный перцептрон с применением
сегментирования тензоров
Начнем с простого примера тренировки с распараллеливанием по
данным. Наша цель – распараллеливание процесса тренировки по
нескольким устройствам – в данном случае по восьми ядрам TPU. Вы
увидите, что код практически возвращается к виду исходной версии
из главы 2, хотя обладает всеми преимуществами (или даже добав
ляются дополнительные) версий из главы 7 и экспериментальных
методик из приложения D.
8.2.1
Восьмиканальное распараллеливание данных
Пропускаем те части, которые не изменились: загрузку данных
и формирование структуры многослойного перцептрона. Функции
потерь и обновления возвращаются к исходному состоянию, опи
санному в главе 2.
Листинг 8.11
Функции потерь и обновления
INIT_LR = 1.0
DECAY_RATE = 0.95
DECAY_STEPS = 5
NUM_EPOCHS = 20
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits = batched_predict(params, images)
log_preds = logits - logsumexp(logits)
return -jnp.mean(targets*log_preds)
❶
❶
❶
❶
❷
@jit
❸
def update(params, x, y, epoch_number):
loss_value, grads = value_and_grad(loss)(params, x, y)
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
326
Глава 8
Использование сегментирования тензоров
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)], loss_value
❶ Установка метапараметров.
❷ Та же функция потерь, что и ранее.
❸ JIT-аннотированная функция обновления в исходном виде.
Здесь нет ничего нового; мы просто вернулись к исходной версии
кода из главы 2. Самым замечательным является тот факт, что код
в этом месте не изменяется каким-либо образом для применения
распараллеливания. Все изменения по сравнению с кодом из главы 2
будут происходить в цикле тренировки, где они выделены полужир
ным шрифтом в листинге 8.12.
Листинг 8.12
Полный цикл тренировки
@jit
def batch_accuracy(params, images, targets):
images = jnp.reshape(images, (len(images), NUM_PIXELS))
predicted_class = jnp.argmax(batched_predict(params, images), axis=1)
return jnp.mean(predicted_class == targets)
def accuracy(params, data):
accs = []
for images, targets in data:
accs.append(batch_accuracy(params, images, targets))
return jnp.mean(jnp.array(accs))
sharding = PositionalSharding(jax.devices()).reshape(8, 1)
import time
params = init_params
for epoch in range(NUM_EPOCHS):
start_time = time.time()
losses = []
for x, y in train_data:
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = one_hot(y, NUM_LABELS)
x = jax.device_put(x, sharding)
y = jax.device_put(y, sharding)
params = jax.device_put(params, sharding.replicate())
params, loss_value = update(params, x, y, epoch)
losses.append(jnp.sum(loss_value))
epoch_time = time.time() - start_time
train_acc = accuracy(params, train_data)
test_acc = accuracy(params, test_data)
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
print("Training set loss {}".format(jnp.mean(jnp.array(losses))))
print("Training set accuracy {}".format(train_acc))
print("Test set accuracy {}".format(test_acc))
❶
❷
❷
❸
Многослойный перцептрон с применением сегментирования тензоров
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
…
>>>
>>>
>>>
>>>
327
Epoch 0 in 1.66 sec
Training set loss 0.803970217704773
Training set accuracy 0.7943314909934998
Test set accuracy 0.8018037676811218
Epoch 1 in 1.04 sec
Training set loss 0.7146415114402771
Training set accuracy 0.8670936822891235
Test set accuracy 0.875896155834198
Epoch 19
Training
Training
Test set
in 0.88 sec
set loss 0.6565680503845215
set accuracy 0.9357565641403198
accuracy 0.9364372491836548
❶ Подготовка сетки устройств.
❷ Сегментирование тензора данных по восьми устройствам.
❸ Репликация параметров модели на восьми устройствах.
Здесь можно видеть, что изменения минимальны и расположе
ны достаточно компактно. Во-первых, создается конфигурация
сегментирования – сетка (8, 1) из ядер TPU. Затем тензоры данных
назначаются на устройства в соответствии с созданной схемой сег
ментирования. Это означает, что тензоры x и y распределяются по их
первому измерению (измерению пакетов), а их сегменты связывают
ся с конкретными собственными устройствами. Также выполняется
репликация параметров модели на всех устройствах, поэтому каждое
устройство имеет собственную полную копию параметров модели.
Вот и все. Всех описанных выше изменений достаточно, чтобы
создать код для тренировки нейронной сети, которому ничего неиз
вестно о распараллеливании, тем не менее он полностью распреде
лен по восьми ядрам TPU. Это действительно великолепно и намно
го проще, чем тренировка по схеме SPMD с использованием pmap()
или распараллеливания вычислений с помощью pjit(). Возможно,
потребуется более точная настройка гиперпараметров, таких как
размер пакета (мы не изменили его по сравнению с примером без
распараллеливания), чтобы существенно ускорить процедуру тре
нировки. Кроме того, при таком подходе можно с легкостью объеди
нять распараллеливание по данным и распараллеливание модели.
8.2.2
Четырехканальное распараллеливание по данным,
двухканальное распараллеливание тензора
В этом подразделе мы немного подкорректируем код для выпол
нения тренировки с распараллеливанием по данным и распарал
леливанием модели. Здесь используется четырехканальное распа
раллеливание пакетов данных и двухканальное распараллеливание
тензора модели. Это означает, что тренировочные образцы разделя
328
Глава 8
Использование сегментирования тензоров
ются по четырем группам устройств, а параметры модели (некото
рые матрицы весов) разделяются по двум группам устройств.
Распараллеливание тензора и распараллеливание конвейера
Существует два типа распараллеливания модели: распараллеливание
тензора и распараллеливание конвейера.
Распараллеливание (или сегментирование) тензора (tensor parallelism
(sharding)) распределяет вычисление некоторого тензора по различным
устройствам. Например, большой вложенный или полносвязный слой
в нейронной сети можно сегментировать по различным устройствам
так, чтобы каждое устройство вычисляло часть выходного результата
этого слоя. Поскольку тензор может иметь более одного измерения, сегментирование применяется одновременно к различным (многим) измерениям. Для вложенного или полносвязного слоя веса и активации
могут сегментироваться независимо друг от друга.
Распараллеливание конвейера (pipeline parallelism (pipelining)) распределяет вычисление нейронной сети так, чтобы ее различные слои
вычислялись на отдельных устройствах. Например, в трансформере глубокого обучения (deep transformer) блоки трансформера могут размещаться на различных устройствах так, что устройство 1 содержит слои
с 1 по 4, а устройство 2 работает со слоями с 5 по 8.
Эти типы распараллеливания не являются взаимоисключающими, их
можно объединять при необходимости (а также с распараллеливанием
по данным).
Более подробно узнать о распараллеливании по данным, распараллеливании тензора, конвейера и других типах распараллеливания можно здесь: https://huggingface.co/transformers/v4.9.0/parallelism.html
и https://openai.com/index/techniques-for-training-large-neural-net
works/.
Мы модифицировали процедуру сегментирования и изменили
функцию инициализации параметров модели, поэтому некоторые
веса слоев стали сегментированными. Кроме того, мы сделали ней
ронную сеть более глубокой и широкой, чтобы приблизить ее к вари
анту, в котором имеет смысл тренировка с распараллеливанием моде
ли (хотя мы все еще далеко от реально большой нейросети, поскольку
наша модель отлично подходит для единственного акселератора).
Листинг 8.13
Сегментирование весов модели
sharding = PositionalSharding(jax.devices()).reshape(4, 2)
LAYER_SIZES = [28*28, 10000, 10000, 10]
PARAM_SCALE = 0.01
❶
❷
Многослойный перцептрон с применением сегментирования тензоров
def init_network_params(sizes, key=random.PRNGKey(0), scale=1e-2):
"""Initialize all layers for a fully-connected
neural network with given sizes"""
# Инициализация всех слоев для полносвязной нейронной сети
# с заданными размерами.
329
❸
def random_layer_params(m, n, key, scale=1e-2):
"""A helper function to randomly initialize
weights and biases of a dense layer"""
# Вспомогательная функция для случайной инициализации
# весов и отклонений плотного слоя.
w_key, b_key = random.split(key)
return (scale * random.normal(w_key, (n, m)),
scale * random.normal(b_key, (n,)))
keys = random.split(key, len(sizes))
return [random_layer_params(m, n, k, scale)
for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
init_params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
sharded_params = []
for i,(w,b) in enumerate(init_params):
print(i, w.shape, b.shape)
if i==0:
w = jax.device_put(w, sharding.replicate(0))
b = jax.device_put(b, sharding.replicate(0))
elif i==1:
w = jax.device_put(w, sharding.replicate(0))
b = jax.device_put(b, sharding.replicate(0))
elif i==2:
w = jax.device_put(w, sharding.replicate())
b = jax.device_put(b, sharding.replicate())
sharded_params.append((w,b))
❹
❹
❹
❹
❺
❺
>>> 0 (10000, 784) (10000,)
>>> 1 (10000, 10000) (10000,)
>>> 2 (10, 10000) (10,)
for (w,b) in sharded_params:
jax.debug.visualize_array_sharding(w)
jax.debug.visualize_array_sharding(b)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
┌───────────┬───────────┐
│
│
│
│
│
│
│
│
│
│
│
│
│TPU 0,2,4,6│TPU 1,3,5,7│
│
│
│
│
│
│
❻
❻
❻
Глава 8
330
>>>
>>>
>>>
>>>
>>>
>>>
...
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
❶
❷
❸
❹
❺
❻
Использование сегментирования тензоров
│
│
│
│
│
│
└───────────┴───────────┘
┌───────────┬───────────┐
│TPU 0,2,4,6│TPU 1,3,5,7│
└───────────┴───────────┘
┌───────────────────┐
│
│
│
│
│
│
│
│
│TPU 0,1,2,3,4,5,6,7│
│
│
│
│
│
│
│
│
└───────────────────┘
┌───────────────────┐
│TPU 0,1,2,3,4,5,6,7│
└───────────────────┘
Создание двумерной схемы сегментирования.
Углубление и расширение нейронной сети.
В функции инициализации изменений нет.
Репликация весов по оси 0, сегментов по оси 1.
Репликация весов по всем осям.
Визуальное представление сегментирования весов.
В приведенном выше примере мы создали сетку (4, 2). Наша
цель – сегментирование тренировочных элементов по первой оси
сетки и некоторых параметров модели по второй оси.
Функция инициализации весов нейронной сети остается неиз
менной. Мы изменили только структуру нейросети. Теперь она со
держит больше слоев, а в скрытых слоях увеличилось количество
нейронов.
После генерации весов модели мы сегментируем конкретные
матрицы весов и отклонений. Для первых двух слоев выполняется
репликация их весов и отклонения по оси 0 сетки и их сегментиро
вание по оси 1. Это делается с помощью параметра sharding.replicate(0). Последний слой реплицируется по всей сетке полностью,
поскольку он относительно мал.
Визуальное представление сегментирования показывает, что пер
вые два слоя распределены по двум группам устройств. Часть весов
размещена в группе, содержащей TPU {0, 2, 4, 6}. Другая часть весов
находится в группе с TPU {1, 3, 5, 7}. Последний слой реплицирован
на всех восьми TPU.
331
Многослойный перцептрон с применением сегментирования тензоров
Мы подготовили и сегментировали параметры модели и теперь
можем выполнить полный цикл тренировки.
Листинг 8.14
Полный цикл тренировки
params = sharded_params
for epoch in range(NUM_EPOCHS):
start_time = time.time()
losses = []
for x, y in train_data:
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = one_hot(y, NUM_LABELS)
x = jax.device_put(x, sharding.replicate(1))
y = jax.device_put(y, sharding.replicate(1))
params, loss_value = update(params, x, y, epoch)
losses.append(jnp.sum(loss_value))
epoch_time = time.time() - start_time
❶
❶
train_acc = accuracy(params, train_data)
test_acc = accuracy(params, test_data)
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
print("Training set loss {}".format(jnp.mean(jnp.array(losses))))
print("Training set accuracy {}".format(train_acc))
print("Test set accuracy {}".format(test_acc))
>>>
>>>
>>>
>>>
...
Epoch 0 in 14.36 sec
Training set loss 0.7603772282600403
Training set accuracy 0.8505693674087524
Test set accuracy 0.8597139716148376
❶ Репликация тренировочных данных по второй оси, размещение сегментов по
первой оси.
Внесены два изменения. Во-первых, мы изменили схему сегмен
тирования для тренировочных данных. Теперь они реплицируются
по второй оси сетки (с индексом 1), но сегментирование выполняется
по первой оси (в четырех группах устройств). Во-вторых, для пара
метров модели не требуется репликация внутри цикла, так как мы
настроили их сегментирование ранее. Градиенты также будут вычис
ляться локально на тех же устройствах, где размещаются параметры
модели, поэтому после обновления параметров модели они остаются
на том устройстве, на котором находились перед обновлением, и их
локализация не изменяется. Вы сами видите, насколько легко орга
низовать объединение распараллеливания по данным и распаралле
ливания модели с помощью новых API распределенных массивов.
В этом конкретном примере мы не получаем каких-либо замет
ных преимуществ от распараллеливания модели в процессе трени
ровки, потому что модель остается слишком маленькой. Тем не ме
Глава 8
332
Использование сегментирования тензоров
нее этот навык окажется весьма полезным в реальном мире очень
больших моделей.
Резюме
Сегментирование тензоров и новый API распределенных мас
сивов позволяют с легкостью организовать распараллеливание
с минимальными изменениями кода.
С версии JAX 0.4.1 введен тип jax.Array – универсальный тип мас
сива, заменивший типы DeviceArray, ShardedDeviceArray и GlobalDeviceArray. Новый тип jax.Array помогает сделать распарал
леливание одним из основных механизмов ядра JAX.
Тип jax.Array упрощает и унифицирует внутренние механизмы
JAX и позволяет унифицировать трансформации jit() и pjit():
если данные передаются на устройства, то jit() распараллелива
ется автоматически.
В JAX вычисления учитывают размещение данных. Сегментиро
вание тензоров можно представить визуально, как расширение
правил распределения устройств из главы 3.
Функция mesh_utils.create_device_mesh() возвращает наиболее
производительный порядок устройств с учетом заданной формы.
PositionalSharding работает как массив с настройками устройств
в качестве элементов.
Параметр NamedSharding может выражать сегменты с именами
вместо позиций в массиве.
Хорошо знакомая функция jax.device_put() также может при
нимать объект сегментирования вместо конкретного устройства
и размещать данные соответствующим образом.
Можно дать компилятору указание о том, как сегментировать
промежуточные переменные функции. Для этого используется
функция jax.lax.with_sharding_constraint(), которая во многом
похожа на функцию jax.device_put(), но применяется внутри JITдекорированных функций.
Можно воспользоваться методом сегментирования replicate
(axis=NUMBER) для копирования срезов тензора на каждое устрой
ство по заданному измерению. Если оси не заданы, то репликация
выполняется по каждой оси.
При использовании сегментирования, если два аргумента вы
числения размещены на различных устройствах или порядок их
устройств не является совместимым, то возникает ошибка.
Распараллеливание по данным и распараллеливание модели лег
ко организовать с помощью сегментирования тензоров, просто
изменив схему размещения данных и почти не изменяя код, вы
полняющий вычисления.
9
Случайные числа в JAX
Темы главы:
генерация (псевдо)случайных чисел в JAX;
различия между NumPy и использованием ключей для
представления состояния генератора псевдослучайных
чисел;
работа с ключами и генерация случайных чисел
в реальных приложениях.
Работа со случайностью важна, поскольку многие алгоритмы машин
ного обучения применяют стохастичность тем или иным способом.
Случайность используется для создания произвольно выбираемых
разделений и выборок данных, для генерации случайных данных
и для выполнения случайно выбранных приращений данных. Слу
чайность также требуется для специализированных алгоритмов,
связанных с нейронными сетями, например Dropout, или для архи
тектур, подобных вариационным автоматическим кодировщикам
(variational auto-encoders – VAE), или для порождающих (генератив
ных) состязательных сетей (generative adversarial network – GAN),
а также при настройке гиперпараметров поиска для получения бо
334
Глава 9
Случайные числа в JAX
лее точных значений. Не менее важна случайность в процессе, явля
ющемся краеугольным камнем глубокого обучения: инициализации
весов. Разумеется, случайность также играет важную роль в других
областях за пределами машинного обучения. Случайные числа за
ложены в основу имитаций методом Монте-Карло и методов статис
тических выборок, и существует множество приложений для ком
пьютерных имитаций процессов реального мира.
Еще одна важная тема в сфере машинного обучения – воспроиз
водимость (результатов). Достаточно часто задается начальное чис
ло (seed), так что независимо от того, как часто выполняется один
и тот же процесс, ожидается, что будет создана одна и та же после
довательность случайных чисел. Но при определении глобального
начального числа (как это делается в NumPy) воспроизводимость
нарушается, если в код вводится распараллеливание. Кроме того,
XLA может выполнять код в порядке, отличающемся от порядка его
записи, чтобы обеспечить его наиболее эффективное выполнение.
Это означает, что, возможно, мы не достигнем воспроизводимости,
определяя конкретный порядок выполнения в коде, который мы
предполагаем получить.
Процесс генерации случайных чисел в JAX структурирован со
всем по-другому, нежели в NumPy, и основой этих различий явля
ется функциональная сущность JAX. NumPy использует внутреннее
состояние генератора случайных чисел и изменяет его при каждом
вызове функции, поэтому последовательные вызовы создают раз
личные случайные значения. В парадигме функционального про
граммирования наличие внутреннего состояния становится весьма
отрицательным фактором, так как оно нарушает чистоту функций,
и они должны возвращать одинаковые значения для одинаковых
входных данных. Поэтому JAX предоставляет альтернативное ре
шение: функционально чистые генераторы случайных чисел. Такие
генераторы потребляют ключи, представляющие состояние генера
тора, следовательно, вы обязаны ответственно относиться к управ
лению ключами. Не используйте какой-либо ключ дважды (иначе
получите одинаковые «случайные» числа) и создавайте новый ключ
каждый раз, когда необходимо сгенерировать новый набор случай
ных значений.
В этой главе вы узнаете, как работать с (псевдо)случайными чис
лами в JAX. Сначала мы кратко рассмотрим пример генерации слу
чайных чисел для расширения набора данных. Затем более подроб
но обсудим функции для работы со случайными числами и сравним
подход NumPy с методиками JAX. В конце главы мы создадим пару
вариантов для реальной практики, в которых требуются случайные
числа: полный конвейер расширения набора данных и инициализа
ция весов случайными значениями для нейронной сети.
Генерация случайных данных
9.1
335
Генерация случайных данных
Начнем с простого примера генерации случайного шума для расши
рения набора тренировочных данных. Предположим, что имеется
набор данных изображений для тренировки нейронной сети клас
сификатора изображений. Чтобы добавить некоторую устойчивость
к ошибкам этой нейросети (а также увеличить размер набора тре
нировочных данных), необходимо выполнить настройку процеду
ры расширения: добавление случайного шума и создание случай
ных (беспорядочных) вращений. Это общепринятая практика при
классификации изображений – расширение данных посредством
выполнения подобных простых трансформаций, чтобы увеличить
разнообразие в наборе данных, и такой подход будет работать как
неявная регуляризация, следовательно, снижать уровень перепод
гонки (чрезмерного обучения – overfitting). В некотором смысле это
процесс, обратный методике фильтрации изображения, показанной
в примере из раздела 3.1.
Случайные и псевдослучайные числа
Многие компьютерные генераторы случайных чисел (random number
generator – RNG) в действительности представляют собой генераторы
псевдослучайных чисел (pseudo-random number generator – PRNG). Они
используют некоторые детерминированные алгоритмические процессы
для создания последовательностей чисел со свойствами, приблизительно соответствующие свойствам последовательностей истинно случайных чисел из требуемого распределения.
Детерминированные алгоритмические процессы инициализируются
и полностью определяются некоторым начальным числом (seed), которое создает внутреннее состояние PRNG. В дальнейшем это состояние
обновляется между последовательными вызовами либо автоматически
(т. е. скрыто от программиста), либо вручную для получения различных
псевдослучайных значений. Генераторы псевдослучайных чисел, использующие одно и то же внутреннее состояние, создают одинаковые
числа.
Поскольку пространство состояний конечно, PRNG обязательно должен
в конце концов вернуться в состояние, уже существовавшее ранее. Наименьшее количество шагов, после которого происходит возврат к ранее
существовавшему значению, называется периодом (period) генератора.
Широко используемый Mersenne Twister (вихрь Мерсенна) PRNG имеет
период, равный 219937–1 итераций.
Отдельным классом PRNG являются криптографические PRNG со специальными требованиями. В Python можно воспользоваться функцией
os.urandom(), применяющей ОС-специализированный источник слу-
Глава 9
336
Случайные числа в JAX
чайности, который должен соответствовать требованиям криптографических приложений, – эту функцию использует генератор random.
SystemRandom(). Генератор не основан на состоянии программного обес
печения, поэтому последовательности не являются воспроизводимыми.
Аппаратные RNG доступны в некоторых компьютерах и операционных
системах. Обычно это устройства, использующие некоторые физические
процессы для генерации истинно случайных чисел. Процессоры Intel
предоставляют инструкции RDRAND и RDSEED для возврата случайных
чисел из аппаратного RNG на отдельной специальной микросхеме. Источник энтропии таких генераторов использует тепловой шум внутри
самой микросхемы.
Существуют и более экзотические варианты, например квантовый источник случайности, созданный в Национальном университете Австралии
(Australian National University; https://qrng.anu.edu.au/). Случайные
числа генерируются посредством измерения квантовых флуктуаций вакуума. Для Python существует API этого генератора, предоставленный
в модуле quantumrandom.
В этой главе рассматриваются только PRNG, поэтому в тексте не будет
постоянно использоваться слово «псевдослучайный», оно заменено
более коротким термином «случайный». Но в каждом случае следует
помнить о том, что в действительности имеется в виду псевдослучайное
значение.
Мы будем использовать новый набор данных для двоичной клас
сификации кошек и собак. С этим набором мы также будем работать
и в главе 11, где создадим нейронную сеть с более развитыми воз
можностями классификации изображений, применяя библиотеки
поддержки нейросетей высокого уровня.
Процесс расширения данных прост. Для каждого изображения во
входном пакете будет случайно выбрана одна из двух процедур рас
ширения: добавление гауссова шума или зеркальное отражение по
горизонтали. На рис. 9.1 показана схема этого процесса.
Добавление
случайного
шума
Набор
данных
Пакеты
Изображение
Расширенное
изображение
Случайный
выбор
Зеркальное
отображение
по
горизонтали
Рис. 9.1 Схема процесса расширения данных
Генерация случайных данных
337
Реализуем части этого процесса, начиная с загрузки исходного на
бора данных.
9.1.1
Загрузка набора данных
Воспользуемся широко известным набором данных с изображения
ми собак и кошек Dogs vs. Cats с сайта Kaggle (https://www.kaggle.
com/c/dogs-vs-cats). На рис. 9.2 показана небольшая выборка из это
го набора данных.
Рис. 9.2 Примеры изображений из набора данных Dogs vs. Cats
В наборах данных TensorFlow Datasets содержится версия Dogs
vs. Cats с именем cats_vs_dogs, где 1738 искаженных изображений
исключены из исходного набора данных, а общее количество изо
бражений равно 23 262. Этот набор данных не разделен на трениро
вочную и тестовую части, поэтому мы должны самостоятельно вы
полнить требуемое разделение.
Листинг 9.1
Загрузка набора данных
import tensorflow as tf
# Необходимо убедиться в том, что TF не видит GPU
# и не захватывает всю память GPU.
tf.config.set_visible_devices([], device_type='GPU')
import tensorflow_datasets as tfds
data_dir = '/tmp/tfds'
#
#
#
#
Поскольку набор данных не содержит отдельные части
для тренировки и тестирования, мы выполняем
собственную процедуру разделения на
тренировочный и тестовый наборы данных.
Глава 9
338
Случайные числа в JAX
# as_supervised=True дает нам (image, label)
# как кортеж вместо словаря dict
data, info = tfds.load(name="cats_vs_dogs",
data_dir=data_dir,
split=["train[:80%]", "train[80%:]"],
as_supervised=True,
with_info=True)
❶
❷
(cats_dogs_data_train, cats_dogs_data_test) = data
❶ Загрузка набора данных cats_vs_dogs.
❷ Выполнение разделения на тренировочный (80 %) и тестовый (20 %) наборы
данных с применением внутренней функциональности TensorFlow Datasets.
Здесь мы не занимаемся реализацией собственной логики раз
деления, полностью полагаясь на функциональность TensorFlow Da
tasets для разделения на тренировочный (80 %) и тестовый (20 %)
наборы данных.
Можно выборочно проверить отдельные изображения из загру
женного набора данных, чтобы получить изображение, показанное
на рис. 9.2, с помощью кода в листинге 9.2.
Листинг 9.2 Выборочная проверка изображений из загруженного
набора данных
import matplotlib.pyplot as plt
plt.rcParams['figure.figsize'] = [20, 10]
CLASS_NAMES = ['cat', 'dog']
ROWS = 2
COLS = 5
i = 0
fig, ax = plt.subplots(ROWS, COLS)
for image, label in cats_dogs_data_train.take(ROWS*COLS):
ax[int(i/COLS), i%COLS].axis('off')
ax[int(i/COLS), i%COLS].set_title(CLASS_NAMES[label])
ax[int(i/COLS), i%COLS].imshow(image)
i += 1
❶
plt.show()
❶ Вывод изображений в сетке 2×5.
Здесь можно видеть, что размеры изображений различны, поэто
му для их обработки одинаковым образом и для упаковки несколь
ких изображений в пакет потребуется некоторая предварительная
обработка всего входного набора данных.
Предварительная обработка будет включать приведение всех изо
бражений к единому размеру 200×200 пикселов, а также нормализа
339
Генерация случайных данных
цию RGB-значений пикселов с отображением значений из диапазо
на [0, 255] в диапазон [0, 1].
Листинг 9.3 Предварительная обработка входного набора данных
и разделение его на пакеты
HEIGHT = 200
WIDTH = 200
NUM_LABELS = info.features['label'].num_classes
import jax.numpy as jnp
def preprocess(img, label):
"""Resize and preprocess images."""
# Изменение размера и предварительная обработка изображений.
return tf.image.resize(img, [HEIGHT, WIDTH]) / 255.0, label
train_data = tfds.as_numpy(
cats_dogs_data_train.map(preprocess).batch(32).prefetch(1))
test_data = tfds.as_numpy(
cats_dogs_data_test.map(preprocess).batch(32).prefetch(1))
❶
❷
❶ Простая функция предварительной обработки для применения к каждому изо-
бражению.
❷ Изменение размеров и нормализация изображения.
Мы сообщаем загрузчику данных о необходимости применения
функции preprocess к каждому экземпляру данных, упаковки всех
изображений в набор пакетов размером 32 элемента и применении
предварительной выборки нового пакета без ожидания завершения
обработки предыдущего пакета на GPU.
При предварительной обработке будет потеряна некоторая ин
формация, но изображения остаются распознаваемыми (см. рис. 9.3).
Рис. 9.3 Образцы изображений из набора данных Dogs vs. Cats после
предварительной обработки
Теперь займемся расширением набора данных.
Глава 9
340
9.1.2
Случайные числа в JAX
Генерация случайного шума
Первым расширением будет добавление случайного шума в изобра
жение. Таким способом можно увеличить количество изображений,
создавая несколько зашумленных версий из одного исходного изо
бражения.
Начнем с генерации случайного гауссова шума и добавления его
в изображение. Для этого потребуется генерация тензора, содержа
щего шум некоторой формы, соответствующей изображению, с по
следующим суммированием тензоров изображения и шума, воз
можно, с использованием некоторых весов.
JAX предоставляет богатый набор функций для генерации слу
чайных чисел из разнообразных распределений. Полный список
функций доступен здесь: https://docs.jax.dev/en/latest/jax.random.
html#list-of-available-functions.
Одной из этих функций является jax.random.normal() для генера
ции случайных значений из стандартного распределения Гаусса со
средним значением (математическим ожиданием) 0 и стандартным
отклонением 1. Она похожа на функцию numpy.random.normal(), но
с другими параметрами. Немного позже мы рассмотрим функцию
jax.random.normal() более подробно. Сначала загрузим изобра
жения.
Листинг 9.4
Загрузка изображений
import jax
batch_images,batch_labels = next(iter(train_data))
batch_images[0].shape
❶
>>> (200, 200, 3)
❷
image = batch_images[0]
image.min(), image.max()
>>> (0.0, 1.0)
plt.imshow(image)
❶
❷
❸
❹
❸
❹
Прием пакета изображений из загрузчика данных.
Это цветное изображение размером 200×200 пикселов.
Изображение описывается значениями с плавающей точкой в диапазоне [0.0, 1.0].
Вывод (отрисовка) изображения.
Мы получили первое изображение из тренировочного набора
данных (см. рис. 9.4).
Генерация случайных данных
341
Рис. 9.4 Первое изображение из тренировочного набора данных
В коде листинга 9.5 не происходит ничего особенного: мы просто
берем первое изображение из первого пакета тренировочного на
бора данных, чтобы использовать его для расширения. Теперь мы
готовы сформировать тензор шума с учетом формы полученного
изображения.
Листинг 9.5
Генерация тензора случайного шума
seed = 42
key = jax.random.PRNGKey(seed)
std_noise = jax.random.normal(key, image.shape)
std_noise.min(), std_noise.max()
❶
❷
❸
❹
>>> (Array(-4.170905, dtype=float32), Array(4.259979, dtype=float32))
noise = 0.5 + 0.1*std_noise
noise.min(), noise.max()
❺
❻
>>> (Array(0.08290949, dtype=float32), Array(0.92599785, dtype=float32))
plt.imshow(noise)
❶ Использование начального числа (seed) для PRNG.
❷ Инициализация ключа для применения в функции, использующей PRNG.
❸ Генерация тензора, содержащего стандартный гауссов шум (матем. ожидание = 0,
ст. отклонение = 1), с формой, соответствующей изображению.
❹ Устанавливается диапазон случайных значений ≈[–4.2, +4.2].
❺ Преобразование стандартного гауссова шума в гауссов шум с матем. ожиданием
= 0,5 и ст. отклонением = 0,1.
❻ Теперь диапазон случайных значений: ≈[0.08, 0.93].
Глава 9
342
Случайные числа в JAX
В приведенном выше коде необходимо отметить несколько важ
ных вещей. Мы используем некоторое начальное число (seed) для
инициализации PRNG, применяя для этого объект, называемый клю
чом (key). Немного позже мы более подробно рассмотрим внутрен
ний механизм PRNG, но сейчас самое важное заключается в том, что
любая функция (здесь: jax.random.normal()), использующая PRNG,
потребляет ключ. Ключ создается посредством вызова функции
random.PRNGkey(seed). Эта функция принимает 64- или 32-битовые
целочисленные значения в качестве начального числа, чтобы сгене
рировать ключ.
Второй параметр функции jax.random.normal() – форма выходно
го значения. Мы передаем форму, которую имеет изображение, по
этому функция фактически генерирует множество случайных чисел
одновременно.
Если сравнить эту функцию с numpy.random.normal(), то вы замети
те абсолютное различие в параметрах. Во-первых, функция из NumPy
(https://numpy.org/doc/stable/reference/random/generated/numpy.
random.normal.html) принимает в качестве аргументов математиче
ское ожидание (среднее значение) и стандартное отклонение, тогда
как функция JAX (https://docs.jax.dev/en/latest/_autosummary/jax.ran
dom.normal.html#jax.random.normal) принимает значение матема
тического ожидания равным 0, а значение стандартного отклонения
равным 1. Во-вторых, функция NumPy не принимает какое-либо на
чальное число или ключ, а в функцию JAX явно передается ключ. Тем
не менее обе функции позволяют передавать форму для выходного
значения.
Мы сгенерировали значения из стандартного распределения Га
усса, затем, умножив их на 0,1 и прибавив 0,5, получили распреде
ление Гаусса с математическим ожиданием 0,5 и стандартным от
клонением 0,1.
Прямо сейчас важно отметить, что для генерации различных слу
чайных значений необходимо использовать разные ключи. Если вы
повторно используете один и тот же ключ, то получите в точности те
же самые псевдослучайные значения, что и ранее. В следующем раз
деле описаны способы генерации ключей. Сгенерированный шум
показан на рис. 9.5.
После этого мы готовы к созданию зашумленной версии изобра
жения.
Листинг 9.6
Генерация зашумленного изображения
new_image = image + noise
new_image.min(), new_image.max()
❶
343
Генерация случайных данных
>>> (Array(0.2063958, dtype=float32),
➥Array(1.7627949, dtype=float32))
new_image = (new_image - new_image.min())/(
➥ new_image.max() - new_image.min())
new_image.min(), new_image.max()
>>> (Array(0., dtype=float32), Array(1., dtype=float32))
plt.imshow(new_image)
❷
❸
❹
❶ Суммирование изображения и шума.
❷ Теперь мы вышли из диапазона [0.0, 1.0].
❸ Нормализация полученного изображения с целью получения значений в диапа-
зоне [0.0, 1.0].
❹ Проверка, позволяющая убедиться в том, что все выполнено успешно.
Рис. 9.5 Сгенерированный случайный шум
Заключительная часть очевидна: мы выполняем суммирование
(объединение) исходного изображения и сгенерированного шума.
Единственное, о чем следует помнить: изображения содержат значе
ния в диапазоне [0.0, 1.0], но после сложения мы выходим за преде
лы этого диапазона. Поэтому необходимо нормализовать результат,
и в итоге мы получаем зашумленную версию исходного изображе
ния (см. рис. 9.6).
Мы обработали одно изображение, и следующим шагом будет
обеспечение обработки пакета изображений. Этот шаг будет выпол
нен немного позже в текущей главе, после того как мы узнаем боль
ше о ключах.
Глава 9
344
Рис. 9.6
9.1.3
Случайные числа в JAX
Зашумленная версия изображения
Выполнение случайного расширения данных
Добавим другой тип случайного расширения: зеркальное отобра
жение по горизонтали. Такая операция обычно не изменяет смысл
изображения, поэтому вполне безопасно генерировать зеркально
отраженные версии изображений кошек и собак, тогда как зеркаль
ное отображение по вертикали может создавать неестественные
картинки с перевернутыми вверх ногами животными. Зеркальное
отображение тензора по горизонтали осуществляется без каких-ли
бо вызовов специализированных функций простой заменой поряд
ка всех строк на обратный в каждом изображении.
Листинг 9.7
Выполнение зеркального отображения по горизонтали
image_flipped = image[:,::-1,:]
plt.imshow(image_flipped)
❶
❶ Изменение порядка на обратный по второй оси тензора.
Мы реверсировали вторую ось тензора. Эта ось содержит столбцы
изображения. Полученное в итоге изображение стало зеркальным
(см. рис. 9.7). Другой вариант: можно было бы применить функцию
fliplr(), доступную и в NumPy, и в JAX.
Теперь мы готовы написать функцию, случайно выбирающую для
применения процедуру расширения: либо генерацию зашумленной
версии изображения, либо зеркальное отображение по горизонтали.
Здесь мы используем функцию, выбирающую только одну из проце
дур расширения. В вариантах реальной практики, возможно, потребу
ется применение комбинации расширений (если они не противоре
345
Генерация случайных данных
чат друг другу). В качестве дополнительного упражнения попробуйте
реализовать такую возможную комбинацию расширений.
Рис. 9.7 Зеркально
отображенная по горизонтали
версия исходного изображения
Мы уже рассматривали похожую функцию в разделе 3.4 и исполь
зовали базисные элементы управления потоком выполнения из jax.
lax как естественный способ ее реализации. Сейчас мы воспользуем
ся функцией jax.lax.switch(), поскольку она позволяет добавлять
дополнительные процедуры расширения. Здесь важно отметить, что
теперь у нас есть две случайные функции и для каждой требуется
ключ: одна для выбора процедуры расширения, вторая для генера
ции случайного шума (если выбран этот вариант).
Использование одного ключа для обеих функций стало бы ошиб
кой, так как ключ должен применяться только один раз. Иначе мы
получим весьма ограниченную случайность, поскольку вызовы обе
их функций будут связаны и сгенерируют коррелированные резуль
таты.
Поэтому необходимо сгенерировать новый ключ для второго вы
зова. Это делается посредством разделения ключа. Для осущест
вления такого действия применяется функция jax.random.split().
В листинге 9.8 показан код для выполнения случайной процедуры
расширения.
Листинг 9.8
Выполнение случайной процедуры расширения
def add_noise_func(image, rng_key):
noise = 0.5 + 0.1*jax.random.normal(rng_key, image.shape)
new_image = image + noise
new_image = (new_image - new_image.min())/(
new_image.max() - new_image.min())
return new_image
❶
Глава 9
346
Случайные числа в JAX
def horizontal_flip_func(image, rng_key):
return jnp.fliplr(image)
❷
augmentations = [
add_noise_func,
horizontal_flip_func
]
❸
def random_augmentation(image, augmentations, rng_key):
key1, key2 = jax.random.split(rng_key)
augmentation_index = jax.random.randint(
key=key1, minval=0, maxval=len(augmentations), shape=())
augmented_image = jax.lax.switch(
augmentation_index, augmentations, image, key2)
return augmented_image
❹
❺
key = jax.random.PRNGKey(4242)
❽
img = random_augmentation(image, augmentations, key)
plt.imshow(img)
❻
❼
❾
Функция для генерации зашумленного изображения.
Функция для выполнения зеркального отображения по горизонтали.
Список возможных процедур расширения.
Функция для выполнения случайного расширения.
Разделение ключа на два ключа.
Случайный выбор одной из функций расширения с использованием первого
ключа.
❼ Применение выбранной функции к изображению и ко второму ключу.
❽ Инициализация ключа некоторым начальным числом.
❾ Вызов всего конвейера для выполнения случайного расширения.
❶
❷
❸
❹
❺
❻
В приведенном выше коде все предыдущие разработки собраны
в конвейер, работающий с одним изображением. Сначала мы реор
ганизовали код из предыдущих подразделов в две функции Python
и список процедур расширения. Каждая функция непременно долж
на иметь одну и ту же сигнатуру (одинаковый список параметров
в неизменном порядке), иначе невозможно будет вызывать их еди
нообразным способом из функции jax.lax.switch(). Поэтому функ
ция для выполнения зеркального отображения по горизонтали так
же принимает ключ для PRNG, хотя не использует его.
Затем мы подготовили функцию для выбора и применения слу
чайного расширения. Эта функция почти та же самая, что и в листин
ге 3.27, с единственным добавлением операции разделения ключа.
Функция потребляет ключ для PRNG и разделяет его на два клю
ча. Первый ключ потребуется для генерации случайного значения
(целочисленного индекса) для выбора одной из доступных проце
дур расширения. Этот ключ передается в функцию jax.random.randint(). Функция randint() с параметрами minval=0 и maxval=2 возвра
Отличия от NumPy
347
щает 0 или 1, но не 2. Второй ключ необходим для генерации тензора
случайного шума, если соответствующая функция была выбрана на
предыдущем шаге. Если вместо нее выбрана функция зеркального
отображения по горизонтали, то второй ключ фактически отбрасы
вается, так как он не нужен для такой операции. Но мы вынуждены
передавать его, чтобы обеспечить возможность вызова различных
функций из jax.lax.switch(). Вы можете поэкспериментировать
с различными начальными значениями ключа, чтобы увидеть, как
работает эта функция. Также можно модифицировать функцию, до
бавляя другие процедуры расширения.
Теперь полученных знаний достаточно, чтобы подробнее рассмот
реть внутренние механизмы PRNG и понять, почему в JAX работа со
случайными числами организована именно таким образом.
9.2
Отличия от NumPy
Прежде чем начать обсуждение JAX, взглянем на NumPy, так как это
поможет нам лучше понять различия между тем, как NumPy и JAX
работают со случайными числами.
9.2.1
Как работает NumPy
Для генерации случайных чисел в NumPy вы просто вызываете функ
ции, чтобы получить случайные значения, выбираемые из различ
ных распределений, и задаете требуемую форму выходных данных.
Важно отметить, что последовательные вызовы возвращают разные
значения. Процесс генерации случайного числа в NumPy схематиче
ски показан на рис. 9.8.
Состояние
Вызов
функции
генерации
случайного
значения
Генератор
случайных
чисел
Случайное
значение
Рис. 9.8 Генерация случайного числа в NumPy;
состояние обновляется внутри самого RNG
Глава 9
348
Случайные числа в JAX
В старых версиях NumPy пользователь вызывал методы из модуля
numpy.random. В новых версиях, начиная с 1.17, вызываются методы
объекта RNG. В листинге 9.9 генерируются случайные значения из
нормального распределения с использованием нового и устаревше
го NumPy API.
Листинг 9.9 Генерация случайных значений в NumPy
from numpy import random
vals = random.normal(loc=0.5, scale=0.1, size=(3,5))
more_vals = random.normal(loc=0.5, scale=0.1, size=(3,5))
vals
>>> array([[0.41442475, 0.66198785, 0.42058724, 0.54404004, 0.43972029],
>>>
[0.41518162, 0.36498766, 0.54380958, 0.63188696, 0.40681366],
>>>
[0.49568808, 0.45086299, 0.49072887, 0.40336257, 0.41021533]])
more_vals
>>> array([[0.37960004, 0.48809215, 0.59210434, 0.53701619, 0.60282681],
>>>
[0.60296999, 0.42870291, 0.61688912, 0.43114899, 0.41913782],
>>>
[0.45572322, 0.48780771, 0.482078 , 0.59348996, 0.41206967]])
from numpy.random import default_rng
rng = default_rng()
vals = rng.normal(loc=0.5, scale=0.1, size=(3,5))
more_vals = rng.normal(loc=0.5, scale=0.1, size=(3,5))
vals
>>> array([[0.49732516, 0.41752651, 0.54104826, 0.60639913, 0.49745545],
>>>
[0.39052283, 0.57229021, 0.54367553, 0.70409461, 0.44481841],
>>>
[0.4184092 , 0.48017174, 0.32490981, 0.30408382, 0.45733146]])
more_vals
>>> array([[0.39426541, 0.6461455 , 0.38793849, 0.50340449, 0.62198861],
>>>
[0.4760281 , 0.43383763, 0.41066168, 0.57226022, 0.36438518],
>>>
[0.71569246, 0.52847295, 0.61811126, 0.45912844, 0.59835265]])
❶
❷
❸
❹
❺
❻
❼
❽
❶
❷
❸
❹
❹
❺
❺
❻
❼
❽
❽
Использование устаревшего механизма генерации случайных значений.
Генерация случайных значений из нормального распределения.
Еще одна операция генерации случайных значений из нормального распределения.
Последовательные вызовы генерируют различные значения.
Использование новой методики генерации случайных значений.
Генерация случайных значений из нормального распределения.
Еще одна операция генерации случайных значений из нормального распределения.
Последовательные вызовы генерируют различные значения.
Мы использовали устаревший механизм генерации случайных
чисел, применив функцию random.normal(), и новую методику, ис
пользующую отдельный объект RNG и метод генератора normal().
349
Отличия от NumPy
Следует отметить, что последовательные вызовы одной и той же
функции с одинаковыми параметрами генерируют различные чис
ла. Это может выглядеть удобным приемом для программиста, ге
нерирующего больше чисел, чем необходимо, но подобный подход
противоречит функциональному стилю, поскольку такая функция
не является функционально чистой. Функционально чистые функ
ции обязаны возвращать одни и те же значения при заданных оди
наковых параметрах. Именно по этой причине в JAX требуется дру
гая методика, которую мы рассмотрим немного позже.
Теперь рассмотрим подробнее процесс генерации случайных чи
сел в NumPy, чтобы лучше понять, как это работает и что нужно из
менить для адаптации процесса в JAX.
9.2.2
Начальное число и состояние в NumPy
Сначала разберемся с понятием начального числа (seed) для ини
циализации PRNG. Время от времени мы использовали начальное
число в примерах кода, где требовалась некоторая случайность.
Применение одного и того же начального числа для PRNG помога
ет обеспечить воспроизводимость. Повторение последовательности
операций генерации случайных чисел будет выдавать абсолютно
одинаковые числа, если вы повторно инициализируете PRNG од
ним и тем же начальным числом. Такое поведение демонстрируется
в листинге 9.10.
Листинг 9.10 Использование начального числа для воспроизведения
случайных значений
random.seed(42)
vals = random.normal(loc=0.5, scale=0.1, size=(3,5))
random.seed(42)
more_vals = random.normal(loc=0.5, scale=0.1, size=(3,5))
vals
❶
❷
❶
❷
❸
>>> array([[0.54967142, 0.48617357, 0.56476885, 0.65230299, 0.47658466],
>>>
[0.4765863 , 0.65792128, 0.57674347, 0.45305256, 0.554256 ],
>>>
[0.45365823, 0.45342702, 0.52419623, 0.30867198, 0.32750822]])
more_vals
❸
>>> array([[0.54967142, 0.48617357, 0.56476885, 0.65230299, 0.47658466],
>>>
[0.4765863 , 0.65792128, 0.57674347, 0.45305256, 0.554256 ],
>>>
[0.45365823, 0.45342702, 0.52419623, 0.30867198, 0.32750822]])
❶ Использование одинакового начального числа для обоих вызовов.
❷ Вызовы абсолютно одинаковы.
❸ Генерируются одни и те же числа.
Глава 9
350
Случайные числа в JAX
Здесь можно видеть, что идентичные вызовы с одним и тем же
значением начального числа генерируют одинаковые числа.
NumPy обеспечивает еще одно интересное свойство – гарантию
последовательной равнозначности (sequential equivalent guarantee).
Это означает, что независимо от того, генерируется ли N отдельных
чисел или массив из N элементов, полученная в результате последо
вательность случайных чисел будет одинаковой. Это можно наблю
дать в листинге 9.11.
Листинг 9.11 Демонстрация гарантии последовательной равнозначности
random.seed(42)
even_more_vals = np.array(
[random.normal(loc=0.5, scale=0.1) for i in range(3*5)]
).reshape((3,5))
❶
even_more_vals
❸
❷
>>> array([[0.54967142, 0.48617357, 0.56476885, 0.65230299, 0.47658466],
>>>
[0.4765863 , 0.65792128, 0.57674347, 0.45305256, 0.554256 ],
>>>
[0.45365823, 0.45342702, 0.52419623, 0.30867198, 0.32750822]])
❶ Использование того же начального числа, что и в предыдущем примере.
❷ Последовательная генерация того же количества (отдельных) случайных чисел.
❸ Сгенерированы те же числа.
В приведенном выше примере мы последовательно сгенерирова
ли случайные числа из нормального распределения, получая значе
ние за значением. Мы создали массив NumPy и преобразовали его
в форму (3, 5), используемую ранее. Полученная в итоге матрица
оказалась точно такой же, как матрица результатов в листинге 9.10.
Кроме того, генераторы PRNG в NumPy сохраняют состояние,
т. е. генератор обладает собственным внутренним состоянием. Со
стояние детерминированно инициализируется с использованием
предоставленного начального числа (или некоторого значения по
умолчанию). Затем состояние обновляется после каждого вызова ге
нератора, поэтому последовательные вызовы одной функции с оди
наковыми параметрами используют различные значения состояния
и возвращают разные значения. Узнать значение состояния можно
с помощью функции random.get_state().
Листинг 9.12
Просмотр состояния PRNG для устаревшей методики
random.seed(42)
random.get_state()
>>> ('MT19937', array([
>>> ...
❶
❷
42, 3107752595, 1895908407, 3900362577,
351
Отличия от NumPy
>>>
>>>
2783561793, 1329389532,
624, 0, 0.0)
836540831,
26719530], dtype=uint32),
vals = random.normal(loc=0.5, scale=0.1, size=(3,5))
random.get_state()
❸
❹
>>> ('MT19937', array([ 723970371, 1229153189, 4170412009, 2042542564,
>>> ...
>>>
3446775024, 1857191784, 1432291794, 4088152671], dtype=uint32),
>>>
40, 1, -0.5622875292409727)
❶ Используется то же начальное число, что и в предыдущих примерах.
❷ Получение состояния.
❸ Генерация некоторых случайных значений.
❹ Получение обновленного состояния.
Здесь можно видеть, что состояние представлено кортежем из
пяти элементов. Первый элемент кортежа содержит информацию
об алгоритме, на котором основан PRNG; в данном случае это ал
горитм MT19937. Далее следует массив из 624 целочисленных зна
чений и еще три числа. Мы не будем подробно рассматривать внут
ренние механизмы и состояние алгоритма MT19937, так как это не
относится к теме книги.
Состояния после инициализации и после вызова функции random.
normal() различны. Если для начального числа снова установить
то же значение 42, то мы вернемся к состоянию, существовавшему
в начале примера. Новая методика NumPy для RNG использует со
четание BitGenerator для создания последовательностей случайных
чисел (обычно беззнаковых целых слов, заполненных последова
тельностями из 32 или 64 случайных битов) и Generator для преоб
разования последовательностей случайных битов из BitGenerator
в последовательности чисел, соответствующих конкретному рас
пределению вероятностей.
BitGenerator управляет состоянием, а начальные числа переда
ются в объекты BitGenerator. По умолчанию используется генератор
PCG64 (он может быть заменен в будущих версиях) с улучшенными
статистическими свойствами по сравнению с устаревшим генерато
ром MT19937. При необходимости можно использовать старый гене
ратор битов MT19937, хотя это не рекомендуется.
MT19937, PCG64, Threefry и другие PRNG
Знание внутреннего устройства механизма генерации случайных чисел
не требуется для использования JAX, но для некоторых читателей эта
тема, возможно, окажется интересной.
352
Глава 9
Случайные числа в JAX
Вихрь Мерсенна (Mersenn Twister) – это широко известный PRNG общего назначения (но не для криптографии), который является генератором
случайных чисел по умолчанию во многих программах. Он был разработан Макото Мацумото (Makoto Matsumoto (松本 眞)) и Такуджи Нишимурой (Takuji Nishimura (西村 拓士) в 1997 г. Длина его периода выбирается равной простому числу Мерсенна, отсюда и название. Период
равен 219937–1, т. е. показатель степени включен в название MT19937 для
стандартной 32-битовой реализации. Состояние включает 624 32-битовых беззнаковых целых числа плюс отдельное целое значение, индексирующее текущую позицию в предшествующем массиве. Такой PRNG
обладает низкой пропускной способностью и большим состоянием, но
не является устойчивым к критическим сбоям (Crush-resistant), так как
проходит многие, но не все тесты на статистическую случайность из биб
лиотеки TestU01.
PRNG называют устойчивым к критическим сбоям (Crush-resistant),
если он прошел все тесты из специального набора Crush (Small
Crush, Crush и Big Crush) из библиотеки TestU01 (https://dl.acm.org/
doi/10.1145/1268776.1268777).
Конгруэнтный генератор с перестановками (permuted congruential generator – PCG) – это PRNG, разработанный М. Е. О’Нилом (Dr. M. E. O’Neill)
в 2014 г. Он имеет состояние небольшого размера, обеспечивает высокую скорость и демонстрирует превосходную статистическую производительность. Более подробно об этом алгоритме можно узнать на
странице автора: https://www.pcg-random.org/. PCG64 представляет
собой реализацию по умолчанию в NumPy. Также существует обновленная более продвинутая версия для вариантов использования с высокой
степенью распараллеливания под названием PCG64DXSM.
Существуют и другие PRNG. Подробности и рекомендации можно найти
на странице: https://numpy.org/doc/stable/reference/random/perfor
mance.html#recommendation.
В JAX используется генератор псевдослучайных чисел Threefry на основе счетчика, описанный в документе «Parallel Random Numbers:
As Easy as 1, 2, 3» (https://www.thesalmons.org/john/random123/pa
pers/random123sc11.pdf) с функциональной моделью разделения,
ориентированной на массивы, описанной в другом документе «Splittable Pseudorandom Number Generators Using Cryptographic Hashing»
(https://dl.acm.org/doi/10.1145/2578854.2503784). Это быстрый PRNG
со сниженной криптографической стойкостью. Threefry является самым
быстрым Crush-resistant PRNG на CPU без поддержки аппаратных AES
и входит в группу самых быстрых PRNG на GPU.
Более подробно о проектном решении PRNG в JAX можно узнать здесь:
https://docs.jax.dev/en/latest/jep/263-prng.html.
Получить доступ к состоянию, хранящемуся внутри объекта BitGenerator, можно с помощью его атрибута state.
Отличия от NumPy
353
Листинг 9.13 Просмотр состояния PRNG при использовании новой
методики
from numpy.random import default_rng
rng = default_rng(42)
rng.bit_generator.state
❶
❷
>>> {'bit_generator': 'PCG64',
>>> 'state': {'state': 274674114334540486603088602300644985544,
>>>
'inc': 332724090758049132448979897138935081983},
>>> 'has_uint32': 0,
>>> 'uinteger': 0}
❸
❸
❸
❸
❸
vals = rng.normal(loc=0.5, scale=0.1, size=(3,5))
rng.bit_generator.state
>>> {'bit_generator': 'PCG64',
>>> 'state': {'state': 45342206775459714514805635519735197061,
>>>
'inc': 332724090758049132448979897138935081983},
>>> 'has_uint32': 0,
>>> 'uinteger': 0}
❶
❷
❸
❹
❺
Импорт PRNG по умолчанию.
Использование того же начального числа, что и ранее.
Получение состояния.
Генерация некоторых случайных значений.
Получение состояния – оно изменилось.
❹
❺
❺
❺
❺
❺
И в этом случае можно видеть, что существует некоторое внут
реннее состояние RNG, изменяющееся между последовательными
вызовами, и это приводит к получению различных результатов вы
зовов функции с одинаковыми явными аргументами. При этом со
стояние является неявным аргументом, невидимым извне, поэтому
пользователь не работает с состоянием в явном виде.
9.2.3
PRNG в JAX
На рис. 9.9 показана схема процесса генерации случайного числа
в JAX. Как мы видим, генераторы PRNG в NumPy используют неко
торое глобальное состояние. Это противоречит функциональному
подходу, применяемому в JAX, и может создавать проблемы с вос
производимостью результатов и распараллеливанием. Если мы
используем любую многопоточную, многопроцессную или муль
тихост-логику в NumPy (например, для реализации вычислений
с распараллеливанием по данным), то не сможем обойтись всего
лишь одним начальным числом и, вероятнее всего, даже не полу
чим возможности управлять изменением внутреннего состояния
в параллельном режиме. Таким образом, случайность приводит
Глава 9
354
Случайные числа в JAX
к состоянию гонки. Но в JAX подобная проблема исключена в прин
ципе. Кроме того, как вы помните из главы 5, JIT может работать
некорректно с функциями, не являющимися чистыми.
Состояние
Вызов
функции
генерации
случайного
значения
Генератор
случайных
чисел
Случайное
значение
Рис. 9.9 Генерация случайного числа в JAX; состояние передается в явной
форме в генератор случайных чисел и внутри него не обновляется
JAX не использует глобальное состояние, вместо этого функции ге
нерации случайных чисел JAX принимают состояние в явной форме.
В JAX вводится концепция ключа, явно представляющего состояние
PRNG. Ключ создается при вызове функции random.PRNGkey(seed),
которая принимает 64- или 32-битовое целочисленное значение как
начальное число для генерации ключа.
Функции генерации случайных чисел принимают ключ, но их по
ведение отличается от NumPy. Они не обновляют ключ каким-либо
образом, а просто используют его как внешнее состояние. Поэтому
если вы многократно передаете одинаковый ключ в функцию, то
каждый раз будете получать один и тот же результат.
Листинг 9.14 Генерация случайных значений в JAX
from jax import random
key = random.PRNGKey(42)
vals = random.normal(key, shape=(3,5))
more_vals = random.normal(key, shape=(3,5))
vals
>>> Array([[-0.716899 , -0.20865498, -2.5713923 ,
>>>
[ 1.3873734 , -0.8396519 , 0.3010434 ,
>>>
[-1.6755073 , 0.31390068, 0.5912831 ,
>>>
dtype=float32)
more_vals
❶
❷
❸
❹
❺
1.0337092 , -1.4035789 ],
0.1421263 , -1.7631724 ],
0.5325395 , -0.9133108 ]],
❺
355
Отличия от NumPy
>>> Array([[-0.716899 , -0.20865498, -2.5713923 ,
>>>
[ 1.3873734 , -0.8396519 , 0.3010434 ,
>>>
[-1.6755073 , 0.31390068, 0.5912831 ,
>>>
dtype=float32)
❶
❷
❸
❹
❺
1.0337092 , -1.4035789 ],
0.1421263 , -1.7631724 ],
0.5325395 , -0.9133108 ]],
Импорт JAX PRNG.
Генерация ключа по некоторому начальному числу.
Генерация случайных значений из нормального распределения.
Повторная генерация случайных значений из нормального распределения.
Получились одинаковые числа.
Были выполнены два последовательных вызова с одинаковым
ключом. Полученные значения аналогичны результату в листин
ге 9.10, где мы преднамеренно переустанавливали начальное чис
ло перед каждым вызовом. Использование одного и того же ключа
в JAX равнозначно повторной установке начального числа перед
каждым вызовом в NumPy. Это абсолютно отличается от результа
тов в листинге 9.9, где между вызовами функции NumPy изменяли
скрытое состояние PRNG. Если посмотреть на состояние, то можно
видеть, что оно представляет собой простой массив из двух 32-би
товых беззнаковых целых значений.
Листинг 9.15. Просмотр ключа
key = random.PRNGKey(42)
type(key)
>>> jaxlib.xla_extension.Array
key
>>> Array([ 0, 42], dtype=uint32)
❶
❷
❸
❶ Генерация ключа по некоторому начальному числу.
❷ Тип ключа – массив.
❸ Ключ содержит два значения типа uint32.
Здесь можно видеть, что состояние, представленное ключом, про
ще, чем состояние, используемое в NumPy. Именно поэтому JAX
применяет другой PRNG.
Если необходимо получить несколько случайных значений, то для
этого предоставляются два основных варианта:
можно запросить более одного случайного числа с одним клю
чом, передавая параметр формы. Именно так мы генерировали
тензоры случайных значений ранее в этой главе. В рассматри
ваемом здесь примере можно было бы запросить генерацию
вдвое большей матрицы (скажем, 6×5), а затем разделить ее на
две матрицы по 3×5;
Глава 9
356
Случайные числа в JAX
можно разделить ключ на два и более ключей, чтобы исполь
зовать разные ключи в последовательности вызовов функции,
генерирующей случайные значения.
Первый способ удобен, когда заранее известно, сколько потребу
ется случайных значений, и для их хранения имеется достаточный
объем памяти. Иногда необходимо очень много случайных чисел, по
этому нерационально генерировать и сохранять их заранее, или не
известно количество требуемых случайных значений. В этом случае
разделение ключа является более предпочтительным вариантом.
Разделение ключа выполняется просто, это делается с помощью
функции random.split(). Функция принимает ключ и целое число,
определяющее, сколько новых ключей нужно создать (по умолчанию
задано значение 2). Функция возвращает объект, подобный массиву,
содержащий заданное количество ключей PRNG.
Функция split() детерминированно преобразует один ключ в не
сколько ключей, которые могут сгенерировать независимые значения.
При необходимости можно разделять и новые полученные ключи.
ПРИМЕЧАНИЕ Ключ не должен использоваться более одно
го раза. Если ключ передан как аргумент в функцию split(),
то вы должны отбросить (не использовать) старый ключ пос
ле разделения.
Иногда результаты функции split() называют ключом (key)
и подключом (subkey), или подключами (subkeys). По соглашению
ключи, используемые для генерации других ключей, называют прос
то ключами, а ключи, применяемые для генерации случайных зна
чений, – подключами. Но эти названия не представляют какую-либо
иерархию отношений ключей; все ключи имеют одинаковый статус.
Абсолютно не важно, какой из ключей будет использоваться для раз
деления, а какой – для генерации случайных значений.
В листинге 9.16 генерируются последовательно две матрицы слу
чайных значений из нормального распределения, равнозначные ре
зультату примера с применением NumPy в листинге 9.9.
Листинг 9.16 Генерация различных случайных значений в JAX
from jax import random
key = random.PRNGKey(42)
key
❶
>>> Array([ 0, 42], dtype=uint32)
key1, key2 = random.split(key, num=2)
key1
❷
❸
357
Отличия от NumPy
>>> Array([2465931498, 3679230171], dtype=uint32)
key2
>>> Array([255383827, 267815257], dtype=uint32)
vals = random.normal(key1, shape=(3,5))
more_vals = random.normal(key2, shape=(3,5))
vals
>>> Array([[-0.95032114,
>>>
[ 0.6062024 ,
>>>
[-0.95166105,
>>>
dtype=float32)
0.89362663, -1.9382219 ,
0.37990445, 0.30284515,
1.2611355 , 0.41334143,
more_vals
>>> Array([[ 0.7679137 , 0.46966743, -1.4884446 ,
>>>
[-0.66447204, -1.3314192 , 1.5208852 ,
>>>
[ 1.1168208 , -0.4216216 , 0.32398054,
>>>
dtype=float32)
❶
❷
❸
❹
❺
❻
❸
❹
❺
❻
-0.9676806 , -0.3920417 ],
1.3282853 , 1.1882905 ],
-0.7831721 , 0.09786294]],
❻
-1.155719 , -1.2574353 ],
-0.55124223, -0.3213504 ],
1.3500887 , -0.22909231]],
Генерация ключа по некоторому начальному числу.
Разделение ключа на два.
Полученные два ключа различны.
Генерация случайных значений из нормального распределения с помощью первого ключа.
Генерация случайных значений из нормального распределения с помощью второго ключа.
Получены различные числа.
Мы разделили ключ на два различных ключа и использовали
функцию random.normal() дважды с разными ключами. В этом слу
чае функция сгенерировала различные результаты.
ПРИМЕЧАНИЕ Никогда не используйте один и тот же ключ
повторно. Единственным исключением является случай, ког
да необходимо получить одинаковые результаты.
Сгенерировать нужное количество ключей легко. В приведенном
ниже примере (листинг 9.17) генерируется 100 ключей. Далее пер
вый ключ будет использоваться для разделения, а все остальные –
для генерации случайных чисел.
Листинг 9.17
Генерация 100 ключей
key = random.PRNGKey(42)
key, *subkeys = random.split(key, num=100)
key
>>> Array([ 825763528, 3327736007], dtype=uint32)
❶
❷
❸
Глава 9
358
Случайные числа в JAX
len(subkeys)
>>> 99
❹
❶ Генерация ключа по некоторому начальному числу.
❷ Разделение полученного ключа на 100 ключей.
❸ Этот ключ будет использоваться для дальнейшего разделения при необходи
мости.
❹ Эти ключи предназначены для генерации случайных чисел.
После разделения исходного ключа мы решили использовать пер
вый ключ для дальнейшего разделения, а все остальные ключи – для
генерации случайных чисел. Но при этом все ключи имеют равный
статус, и вы можете использовать любой из них, как пожелаете.
Другой способ создания новых ключей на основе существующе
го ключа и некоторых данных – использование функции fold_in().
Функция принимает существующий ключ и некоторое целое число,
добавляет данные в ключ и возвращает новый ключ, являющийся
статистически безопасным для генерации потока новых псевдослу
чайных значений. Вариант использования из реальной практики:
имеется цикл, на каждой итерации которого необходимо иметь от
дельный ключ. Можно заранее подготовить список ключей и индек
сировать его в цикле. Но такой подход потребует большого объема
памяти для длинных циклов, а кроме того, не всегда есть возмож
ность узнать заранее количество итераций.
Листинг 9.18 Использование функции fold_in() для генерации
новых ключей
key = random.PRNGKey(42)
for i in range(5):
new_key = random.fold_in(key, i)
print(new_key)
vals = random.normal(new_key, shape=(3,5))
# Здесь со значениями выполняются некоторые операции.
>>>
>>>
>>>
>>>
>>>
❶
❷
❸
[1832780943 270669613]
[ 64467757 2916123636]
[2465931498 255383827]
[3134548294 894150801]
[2954079971 3276725750]
❶ Генерация ключа по некоторому начальному числу.
❷ Генерация новых ключей на основе существующего ключа и номера итерации.
❸ Генерация данных с новым ключом.
На каждой итерации цикла генерируется новый ключ. Нет необ
ходимости заранее составлять список используемых ключей. По вы
веденным результатам понятно, что каждый ключ является новым.
Отличия от NumPy
359
Функция fold_in() принимает только целое число как дополни
тельные данные. Если требуется добавлять в ключ данные других
типов, то необходимо преобразовать их в целые числа. Например,
имеются некоторые строки, скажем имена слоев нейронной сети.
В этом случае можно воспользоваться некоторой детерминированной
хеш-функцией для преобразования строковых данных в целое число.
Библиотека Python hashlib предоставляет разнообразные возможно
сти выполнения такой операции. В листинге 9.19 используется хешалгоритм SHA-1 из этой библиотеки для формирования целого числа.
Листинг 9.19
Использование функции fold_in() со строками
import hashlib
def my_hash(s):
return int(hashlib.sha1(s.encode()).hexdigest()[:8], base=16)
some_string = 'layer7_2'
some_int = my_hash(some_string)
some_int
❶
❷
>>> 2649017889
key = random.PRNGKey(42)
new_key = random.fold_in(key, some_int)
new_key
❸
>>> Array([3110527424, 3716265121], dtype=uint32)
❶ Функция для генерации целых чисел из строк.
❷ Генерация 32-битового целого числа из строки.
❸ Генерация нового ключа из существующего ключа и целого числа.
Мы берем строку, вычисляем по ней SHA-1 хеш-сумму, преобра
зовываем вычисленную хеш-сумму в строку, состоящую из шест
надцатеричных цифр, и отрезаем первые 8 цифр. Поскольку каждая
шестнадцатеричная цифра кодируется 4 битами, для кодирования
8 шестнадцатеричных цифр требуется 32 бита, которые легко преоб
разовать в 32-битовое целое число. С помощью этого числа без ка
ких-либо затруднений генерируется новый ключ при вызове функ
ции fold_in().
ПРИМЕЧАНИЕ Использование встроенной в Python функ
ции hash() для получения целых чисел из строк может при
вести к невоспроизводимым результатам, так как функция
hash() по умолчанию рандомизирована. Более подробно
об этом см. здесь: https://jimmycallin.com/2016/01/03/makeyour-python3-code-reproducible/.
Глава 9
360
Случайные числа в JAX
Важное различие между NumPy и JAX состоит в том, что JAX не
обеспечивает гарантию последовательной равнозначности (se
quential equivalent guarantee). Отсутствие такой гарантии связано
с тем, что она может помешать векторизации на аппаратном обо
рудовании SIMD (а кроме того, потому, что не существует известных
пользователей или примеров, для которых это свойство является
важным).
В листинге 9.20 демонстрируется генерация массива случайных
значений 3×5 с применением единственного вызова функции с од
ним ключом и сравнение полученного результата с массивом значе
ний, последовательно сгенерированных с использованием ключей,
полученных из исходного ключа.
Листинг 9.20 Демонстрация отсутствия гарантии последовательной
равнозначности в JAX
import jax.numpy as jnp
key = random.PRNGKey(42)
subkeys = random.split(key, num=3*5)
vals = random.normal(key, shape=(3,5))
more_vals = jnp.array(
[random.normal(key) for key in subkeys]
).reshape((3,5))
vals
>>> Array([[-0.716899 , -0.20865498, -2.5713923 ,
>>>
[ 1.3873734 , -0.8396519 , 0.3010434 ,
>>>
[-1.6755073 , 0.31390068, 0.5912831 ,
>>>
dtype=float32)
more_vals
❶
❷
❸
❹
❺
1.0337092 , -1.4035789 ],
0.1421263 , -1.7631724 ],
0.5325395 , -0.9133108 ]],
❺
>>> Array([[-0.46220607, -0.33953536, -1.1038666 , 0.34873134, -0.577846 ],
>>>
[ 1.9925295 , 0.02132202, 0.7312121 , -0.40033886, -0.9915146 ],
>>>
[-0.03269025, 1.1624117 , 0.15050009, -1.3577023 , -1.7751262 ]],
>>>
dtype=float32)
❶
❷
❸
❹
❺
Генерация ключа по некоторому начальному числу.
Разделение ключа на 15 новых ключей.
Генерация матрицы 3×5 с использованием единственного вызова функции.
Последовательная генерация значений матрицы 3×5.
Получены различные результаты.
Мы сгенерировали две матрицы разными способами и убедились
в том, что они различны. Кроме того, не существует простого спо
соба последовательной генерации множества значений по одному
361
Отличия от NumPy
ключу. Если для каждой генерации использовать один и тот же ключ,
то все создаваемые значения будут равными.
9.2.4
Более подробная конфигурация JAX PRNG
NumPy предоставляет несколько реализаций PRNG, но и про JAX
можно сказать то же самое, хотя это другой набор PRNG. Как уже
отмечалось ранее, по умолчанию JAX использует генератор псевдо
случайных чисел Threefry на основе счетчика. Также имеются экс
периментальные PRNG с применением XLA RngBitGenerator (https://
openxla.org/xla/operation_semantics#rngbitgenerator).
Преимущество Threefry PRNG заключается в том, что он остается
одинаковым во всех сегментах, CPU, GPU, TPU и в различных верси
ях JAX и XLA. Но, возможно, вы не захотите использовать принятый
по умолчанию PRNG, если он медленно компилируется / выполняет
ся на TPU или если требуется эффективное сегментирование.
ПРИМЕЧАНИЕ Экспериментальные реализации пока еще
не были полностью протестированы опытным путем (напри
мер, с использованием Big Crush), поэтому нет никаких га
рантий того, что они не изменятся в следующих версиях JAX.
В дополнение к принятому по умолчанию PRNG Threefry, имею
щему внутреннее название threefry2x32, существуют два экспери
ментальных PRNG под названиями rbg и unsafe_rbg.
PRNG rbg использует тот же Threefry как стандартный PRNG для
разделения, но XLA RBG для генерации данных. PRNG unsafe_rbg
служит только для демонстрационных целей и использует RBG для
разделения и генерации. Оба PRNG rbg и unsafe_rbg работают быст
рее на TPU, а также остаются одинаковыми в сегментах. XLA RBG
не гарантирует детерминированность между различными внутрен
ними компонентами и разными версиями компилятора. Экспери
ментальные PRNG можно подключить, установив флаг jax_default_
prng_impl при инициализации JAX, как показано в листинге 9.21.
Листинг 9.21
Использование экспериментальных PRNG
from jax.config import config
config.update("jax_default_prng_impl", "rbg")
import jax
from jax import random
key = random.PRNGKey(42)
key
>>> Array([ 0, 42,
0, 42], dtype=uint32)
❶
❷
❸
Глава 9
362
Случайные числа в JAX
❶ Установка флага для использования RBG PRNG.
❷ Генерация ключа по некоторому начальному числу.
❸ Сгенерированный ключ отличается от создаваемого по умолчанию.
В приведенном выше примере мы установили флаг jax_default_
prng_impl, чтобы воспользоваться PRNG rbg и сгенерировать ключ,
отличающийся от ключа, созданного PRNG, принятого по умолчанию.
В настоящее время ведется разработка подключаемых реализа
ций PRNG (https://docs.jax.dev/en/latest/jep/9263-typed-keys.html), но
она пока еще не завершена.
Есть еще один флаг jax_threefry_partitionable, подключающий
новую реализацию PRNG Threefry, которая сегментируется более
эффективно. При использовании обычного PRNG Threefry проблема
заключается в том, что по историческим причинам он не является
автоматически сегментируемым. Поэтому требуются некоторые
операции обмена данными между устройствами для получения сег
ментированного итогового результата.
Установка для флага jax_threefry_partitionable значения True
(по умолчанию установлено значение False) приводит к переклю
чению на новую реализацию (которая пока еще находится в разра
ботке). Новая реализация исключает накладные расходы на обмен
информацией, хотя сгенерированные случайные значения могут от
личаться от значений, полученных при сброшенном флаге (со зна
чением False). Они остаются детерминированными и будут одина
ковыми в конкретной версии JAX, но могут оказаться различными
в разных релизах.
Мы завершили работу с генераторами псевдослучайных чисел
в JAX, и необходимо запомнить самый важный факт: PRNG JAX тре
бует передачи состояния PRNG, представленного ключом. Это глав
ное отличие от NumPy, где состояние PRNG скрыто внутри. Теперь
нам известно все необходимое для реализации полного конвейе
ра расширения данных и других операций, полезных в реальной
практике.
9.3
Генерация случайных чисел в реальных
приложениях
Мы будем рассматривать два варианта:
создание полного конвейера для расширения данных, который
может работать с пакетами изображений;
реализация инициализации нейронной сети.
Начнем с варианта из раздела 9.1 и окончательно завершим соз
дание конвейера расширения данных, чтобы он работал не только
363
Генерация случайных чисел в реальных приложениях
с одним экземпляром, но непрерывно обрабатывал каждый прихо
дящий пакет.
9.3.1
Создание полного конвейера расширения данных
Мы оставили без внимания весьма важную функцию создаваемого
конвейера расширения данных изображений – управление ключа
ми. Необходимо так организовать цикл по загрузчику данных, что
бы генерировалось достаточное количество ключей для обработки
всех изображений при работе с каждым пакетом. При наличии зна
ний, полученных в предыдущих разделах, мы можем полностью за
вершить решение этой задачи.
Пропустим бóльшую часть уже написанного ранее кода для за
грузки, предварительной обработки и визуализации данных. Все это
можно найти в репозитории книги. Здесь мы будем рассматривать
только основной цикл, особо выделяя работу с ключами.
Листинг 9.22
Цикл расширения данных изображений
key = random.PRNGKey(42)
❶
for x, y in train_data:
display_batch(x, y, 4, 8, CLASS_NAMES)
❷
batch_size = len(x)
key, *subkeys = random.split(key, num=batch_size+1)
❸
aug_x = jax.vmap(
random_augmentation, in_axes=(0,None,0)
❹
)(x, augmentations, jnp.array(subkeys))
❺
display_batch(aug_x, y, 4, 8, CLASS_NAMES)
❻
# ...
# Выполнение некоторой другой обработки и тренировка нейросети.
# ...
Генерация ключа по некоторому начальному числу.
Вывод неизмененных изображений.
Разделение ключа на достаточное количество подключей и обновленный ключ.
Векторизация функции по ее первому и третьему аргументам.
Использование ключей для случайно выбранного расширения данных изображений в векторизованной функции.
❻ Вывод изображений с расширенными данными.
❶
❷
❸
❹
❺
Здесь главное действие – генерация достаточного количества под
ключей для обработки каждого изображения в пакете. Кроме того,
следует помнить про обновление ключей, чтобы можно было повто
рять процесс при последовательной обработке пакетов.
Вместо поочередной обработки изображений в пакете мы при
меняем знания, полученные в главе 6, для векторизации функции
обработки одного изображения random_augmentation() и применяем
364
Глава 9
Случайные числа в JAX
ее ко всему пакету изображений. Необходимо предоставить массив
подключей, поэтому выполняется преобразование списка (list) Py
thon в массив Array JAX.
Я преднамеренно пропустил часть, относящуюся к тренировке
нейронной сети. Вы можете адаптировать код, чтобы использовать
его для нейронной сети, аналогичной той, что применяется для
классификации изображений MNIST. В главе 11 будет использовать
ся высокоуровневая библиотека поддержки нейросетей для реали
зации на современном уровне сверточной (конволюционной) ней
ронной сети. Изображения с расширенными данными показаны на
рис. 9.10.
Рис. 9.10
Пакет изображений с расширенными данными
Вы можете проверить, действительно ли некоторые изображения
зеркально отображены по горизонтали, а остальные стали зашум
ленными версиями исходных изображений.
9.3.2
Генерация случайных инициализаций
для нейронной сети
Другой вариант использования случайных значений в реальной
практике – инициализация нейронной сети. Для тренировки ней
ронной сети с нуля необходимо начать с некоторой случайно вы
бранной инициализации. Правильная инициализация – это важная
тема в глубоком обучении, и во многих документах предлагаются
методики, позволяющие улучшить результаты. Мы не будем отвле
каться на выяснение того, какая методика инициализации является
самой лучшей, а просто реализуем обобщенный подход, который вы
Генерация случайных чисел в реальных приложениях
365
самостоятельно сможете адаптировать и модифицировать для сво
их потребностей.
На самом деле мы уже использовали этот код в предыдущих гла
вах, но не разбирались детально в том, как он работает. Здесь мы
особо выделяем части, отвечающие за генерацию случайных чисел
и разделение ключей. Мы используем простой многослойный пер
цептрон с той же структурой, что и в предыдущих примерах класси
фикации изображений.
Предположим, что имеется входной слой, принимающий цветные
изображения размером 200×200 пикселов. Такой полностью связан
ный слой содержит 2048 нейронов. Следующий полностью связан
ный слой состоит из 1024 нейронов, и последний слой, тоже полно
стью связанный, предназначен для классификации по двум классам
(двоичной классификации).
Для каждого полностью связанного слоя существует матрица ве
сов размером (количество входных элементов данных) × (количест
во нейронов). Также имеется характеристика отклонения размером
(количество нейронов). Эти переменные обычно с именами w (от
weights – веса) и b (от bias – отклонение) должны быть инициализи
рованы случайным образом. Воспользуемся стандартным нормаль
ным распределением с средним значением 0 и стандартным откло
нением 0,01.
Листинг 9.23
Случайно выбранная инициализация нейронной сети
import jax.numpy as jnp
from jax import random
LAYER_SIZES = [200*200*3, 2048, 1024, 2]
PARAM_SCALE = 0.01
def random_layer_params(m, n, key, scale=1e-2):
w_key, b_key = random.split(key)
return (scale * random.normal(w_key, (n, m)),
scale * random.normal(b_key, (n,)))
def init_network_params(sizes,
key=random.PRNGKey(0), scale=0.01):
keys = random.split(key, len(sizes)-1)
return [random_layer_params(m, n, k, scale)
for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
key = random.PRNGKey(42)
params = init_network_params(LAYER_SIZES, key, scale=PARAM_SCALE)
❶ Описание структуры нейронной сети.
❷ Стандартное отклонение для весов нейронной сети.
❸ Функция для инициализации весов и отклонений плотного слоя.
❶
❷
❸
❹
❺
❻
❼
❽
❾
❿
Глава 9
366
Случайные числа в JAX
❹ Разделение ключа на ключ для весов и ключ для отклонений.
❺ Использование обоих ключей для генерации весов и отклонений.
❻ Функция инициализации всех слоев для полностью связанной нейронной сети
❼
❽
❾
❿
с заданными размерами.
Разделение ключа на ключи, количество которых равно числу слоев.
Передача каждого ключа в функцию для инициализации слоев.
Генерация ключа инициализации с некоторым начальным числом.
Инициализация всех слоев.
Мы начинаем с одного ключа, затем разделяем этот ключ на мно
жество ключей, количество которых равно числу слоев, и передаем
каждый ключ в отдельный вызов функции для инициализации слоя.
Внутри функции инициализации слоя переданный ключ снова
разделяется на два: для весов и для отклонений. Далее каждый но
вый ключ используется для генерации случайных значений.
Резюме
Генераторы псевдослучайных чисел (PRNG) используют некото
рый детерминированный алгоритмический процесс для создания
последовательностей чисел со свойствами, приблизительно соот
ветствующими свойствам последовательностей истинно случай
ных чисел из заданного распределения.
Модуль jax.random предоставляет богатый набор функций для ге
нерации случайных чисел из различных распределений.
Начальное число (seed) используется для инициализации PRNG.
В NumPy PRNG сохраняет состояние, т. е. генератор обладает внут
ренним состоянием.
В NumPy состояние детерминированно инициализируется с ис
пользованием предоставленного начального числа. Затем состоя
ние обновляется после каждого вызова генератора, поэтому по
следовательные вызовы одной и той же функции с одинаковыми
параметрами используют различные значения состояния и воз
вращают разные значения.
NumPy обеспечивает гарантию последовательной равнозначно
сти (sequential equivalent guarantee). Это означает, что независимо
от того, генерируется ли N отдельных чисел или массив из N эле
ментов, полученная в результате последовательность случайных
чисел будет одинаковой.
Использование глобального состояния противоречит функцио
нальному подходу, применяемому в JAX, и может создавать проб
лемы с воспроизводимостью результатов и распараллеливанием.
Как было отмечено в главе 5, JIT может работать некорректно
с функциями, не являющимися чистыми.
Резюме
367
JAX не использует глобальное состояние. Вместо этого функции
JAX для генерации случайных значений явно принимают состоя
ние как аргумент.
JAX вводит концепцию ключа, который явно представляет состо
яние PRNG. Ключ создается с помощью вызова функции random.
PRNGKey(seed).
В JAX функции для генерации случайных значений принимают
ключ, но, в отличие от NumPy, они не обновляют ключ каким-либо
способом, а просто используют его как внешнее состояние. Поэто
му если вы многократно передаете один и тот же ключ в функцию,
то каждый раз будете получать одинаковый результат.
Если необходимо получить много случайных значений, то для
этого существуют два основных варианта:
– можно запросить более одного случайного числа с одним клю
чом, передавая параметр формы;
– можно разделить ключ на два и более ключей, чтобы исполь
зовать разные ключи в последовательности вызовов функции,
генерирующей случайные значения.
Функция split() выполняет детерминированное преобразование
одного ключа в несколько ключей, которые могут генерировать
независимые значения.
Другой способ создания новых ключей на основе существующе
го ключа и некоторого элемента данных: применение функции
fold_in(). Функция принимает существующий ключ и некоторое
целое число, добавляет данные в ключ и возвращает новый ключ,
являющий статистически безопасным для генерации потока но
вых псевдослучайных значений.
Ключ не должен использоваться более одного раза. Если ключ
передан как аргумент в функцию split(), fold_in() или любую
другую функцию генерации случайных значений, то вы должны
отбросить (не использовать) старый ключ после разделения.
Важное различие между NumPy и JAX состоит в том, что JAX не
обеспечивает гарантию последовательной равнозначности (se
quential equivalent guarantee). Отсутствие такой гарантии связано
с тем, что она может помешать векторизации на аппаратном обо
рудовании SIMD.
Существуют экспериментальные PRNG, использующие алгоритм
XLA RBG.
Флаг jax_threefry_partitionable позволяет подключить новую
реализацию PRNG Threefry, обеспечивающую более эффективное
сегментирование.
10
Работа с pytree
Темы главы:
представление сложных структур данных как деревьев
pytree;
использование функций для работы с деревьями pytree;
создание специализированных узлов деревьев pytree.
В предыдущих главах мы в основном использовали тензоры для
представления данных и параметров моделей. Этого достаточно для
простых вариантов, но не всегда удобно, если модели и наборы дан
ных становятся более сложными.
Решение задач машинного обучения часто требует работы с объ
ектами, представленными как списки словарей, списки массивов,
словари массивов и т. п. Например, элементы набора данных могут
быть представлены подобными способами, а веса нейронной сети
обычно организованы в форме некоторой иерархии с весами и от
клонениями, хранящимися для каждого слоя. Если продолжать рабо
тать с таким уровнем сложности, используя низкоуровневые инстру
менты, такие как тензоры и простые операции с ними, то код быстро
станет менее понятным и его объем существенно увеличится. В кон
це концов вам придется изобретать собственные абстракции более
Работа с pytree
369
высокого уровня и более сложные структуры данных. В других сфе
рах деятельности, таких как астрофизика, биоинформатика, моде
лирование погоды и т. д., могут существовать собственные удобные
структуры данных.
К счастью, в JAX имеются некоторые полезные инструменталь
ные средства, и иные из них доступны непосредственно в ядре JAX.
В среде JAX такие древовидные структуры данных, сформирован
ные на основе контейнерных объектов языка Python, называются
pytrees, и эти структуры поддерживаются во многих библиотечных
функциях.
Например, для хранения всех весов и отклонений для многослой
ной нейронной сети наличие единой иерархической структуры дан
ных очень удобно. Такая структура содержит список, представляю
щий слои нейросети, и в каждом элементе списка хранится другой
список или словарь с весами и отклонениями слоя в виде матрицы
и вектора соответственно. Подобный список с вложениями пред
ставляет собой пример дерева pytree.
Структуру pytree очень удобно передавать во все функции, требу
ющие параметров модели, чтобы не предоставлять их по отдельно
сти. Также полезна возможность работы с подобной структурой при
выполнении трансформаций JAX, таких как вычисление градиентов
или автоматическая векторизация функции. К счастью, это осущест
вимо: функции JAX без проблем работают с pytree. Кроме того, JAX
предоставляет комплект инструментов для работы с деревьями py
tree, набор функций для трансформации этих структур с помощью
map-подобной операции, функции для преобразования структуры
в плоскую форму и в обратном направлении, функции для выполне
ния reduce-подобных операций, широко применяемых в функцио
нальном программировании, и т. д.
Вы также можете определять свои классы контейнеров для исполь
зования в pytree. Это удобно при создании собственных абстракций
высокого уровня для слоев нейронных сетей, и при этом требуется
представление pytree для понимания их внутренней структуры. Та
ким образом, все функции обработки pytree будут корректно рабо
тать с вашими классами.
В этой главе рассматриваются все подробности работы с pytree.
В последних главах книги вы узнаете о еще более мощных средствах
из экосистемы JAX, которые вооружат вас абстракциями высокого
уровня для слоев нейронной сети, оптимизаторов и т. п.
Начнем с варианта хранения параметров нейронной сети для
многослойного перцептрона, чтобы продемонстрировать основные
принципы организации pytree, способы и приемы работы с этой
структурой. В следующем разделе рассматриваются более продви
нутые операции с pytree, а в разделе 10.3 вы узнаете, как создавать
специализированные деревья pytree.
370
Глава 10
Работа с pytree
10.1 Представление сложных структур данных
в форме pytree
В многослойной нейронной сети с прямым распространением, или
в многослойном перцептроне, каждый слой состоит из матрицы ве
сов и вектора отклонений. В предыдущих главах мы использовали
достаточно сложные структуры Python для хранения параметров
модели. Структурой, содержащей веса и отклонения, является спи
сок кортежей, состоящих из двух элементов: матрицы весов как пер
вого элемента кортежа и вектора отклонений как второго элемента.
На рис. 10.1 показана схема структуры для варианта с входным и вы
ходным слоями и скрытым слоем между ними из примера в листин
ге 9.23. Фактически это вложенная древовидная структура, которая
в JAX называется pytree.
Рис. 10.1 Структура pytree с весами и отклонениями модели
Термин pytree обозначает древовидную структуру, сформиро
ванную из списков, кортежей, словарей и т. п. Этот набор можно
расширить, зарегистрировав собственный класс контейнера. Мы
сделаем это в разделе 10.3. В листинге 10.1 воспроизводится код
из листинга 9.23 для инициализации соответствующих весов и от
клонений.
Представление сложных структур данных в форме pytree
371
Листинг 10.1 Инициализация нейронной сети случайными
значениями (из листинга 9.23)
import jax.numpy as jnp
from jax import random
LAYER_SIZES = [200*200*3, 2048, 1024, 2]
PARAM_SCALE = 0.01
def random_layer_params(m, n, key, scale=1e-2):
w_key, b_key = random.split(key)
return (scale * random.normal(w_key, (n, m)),
scale * random.normal(b_key, (n,)))
❶
❷
❸
def init_network_params(sizes, key=random.PRNGKey(0), scale=0.01):
keys = random.split(key, len(sizes)-1)
return [random_layer_params(m, n, k, scale)
for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
❹
key = random.PRNGKey(42)
❺
params = init_network_params(LAYER_SIZES, key, scale=PARAM_SCALE)
❻
Стандартное отклонение для весов нейронной сети.
Описание структуры нейронной сети.
Функция для инициализации весов и отклонений плотного слоя.
Функция для инициализации всех слоев для полностью связанной нейронной
сети с заданными размерами.
❺ Генерация ключа инициализации с некоторым начальным числом.
❻ Инициализация всех слоев.
❶
❷
❸
❹
По этой структуре дерева можно пройти и показать (вывести на
экран) все ее свойства. В листинге 10.2 показан вывод форм для ве
сов и отклонений из каждого слоя.
Листинг 10.2
Проход по структуре данных
for i,layer in enumerate(params):
w,b = layer
print(i, w.shape, b.shape)
>>> 0 (2048, 120000) (2048,)
>>> 1 (1024, 2048) (1024,)
>>> 2 (2, 1024) (2,)
❶ Итерация по слоям.
❷ Вывод форм для весов и отклонений.
❶
❷
Проблема этой функции заключается в том, что она должна зара
нее знать структуру, чтобы осмысленно выполнить ее парсинг (син
Глава 10
372
Работа с pytree
таксический анализ). Можно написать более хитроумный код для
работы с различными структурами, но в этом нет необходимости:
существует пакет jax.tree_util, содержащий множество полезных
функций, которые мы более подробно рассмотрим в разделе 10.2,
а сейчас просто воспользуемся функцией jax.tree_util.tree_map().
Функция jax.tree_util.tree_map() отображает заданную функ
цию на аргументы pytree и создает новое дерево pytree. Она похожа
на функцию Python map(), но предназначена для pytree. Мы можем
с легкостью отображать любую структуру в другую структуру, содер
жащую формы для получения почти того же результата, что и в лис
тинге 10.2.
Листинг 10.3 Проход по структуре данных с использованием
tree_map()
shapes = jax.tree_util.tree_map(lambda p: p.shape, params)
❶
for i,shape in enumerate(shapes):
print(i, shape)
❷
❸
>>> 0 ((2048, 120000), (2048,))
>>> 1 ((1024, 2048), (1024,))
>>> 2 ((2, 1024), (2,))
❶ Отображение дерева pytree в дерево pytree форм.
❷ Итерация по слоям.
❸ Вывод форм для весов и отклонений.
Дерево pytree состоит из листов (leaves) и узлов (nodes). Узел – это
само дерево pytree (рекурсивное определение), которое может быть
представлено контейнерным объектом Python. Следующие контей
нерные типы регистрируются в реестре контейнеров pytree: list,
tuple, dict, namedtuple, OrderedDict и None по умолчанию. Следует
особо отметить, что тип None интерпретируется как узел без потом
ков, а не как лист. Все прочие типы по умолчанию считаются листья
ми, включая числовые значения, классы данных (dataclasses), типы
массивов Array и ndarray. Функция jax.tree_util.tree_leaves() воз
вращает листья дерева pytree.
Для демонстрации описанного выше поведения воспользуемся
учебным примером с различными типами вместо pytree с больши
ми массивами.
Листинг 10.4
import
import
import
import
Извлечение листьев из pytree
jax
numpy as np
jax.numpy as jnp
collections
❶
❶
❶
373
Представление сложных структур данных в форме pytree
Point = collections.namedtuple('Point', ['x', 'y'])
example_pytree = [
{
'a': [1, 2, 3],
'b': jnp.array([1, 2, 3]),
'c': np.array([1, 2, 3])
},
[42, [44, 46], None],
31337,
(50, (60, 70)),
Point(640, 480),
collections.OrderedDict([('a', 100), ('b', 200)]),
'some string'
]
❶
❷
❸
❹
❺
❻
❼
❽
❾
❿
⓫
jax.tree_util.tree_leaves(example_pytree)
>>> [1,
>>> 2,
>>> 3,
>>> Array([1, 2, 3], dtype=int32),
>>> array([1, 2, 3]),
>>> 42,
>>> 44,
>>> 46,
>>> 31337,
>>> 50,
>>> 60,
>>> 70,
>>> 640,
>>> 480,
>>> 100,
>>> 200,
>>> 'some string']
❶
❷
❸
❹
❺
❻
❼
❽
❾
❿
⓫
⓬
⓭
⓮
⓯
Необходимые подготовительные операции импорта.
Создание pytree, содержащего различные типы.
Листья из списка внутри словаря.
Массив JAX является листом.
Массив NumPy является листом.
Листья из вложенного списка. None – не лист.
Число является листом.
Листья из кортежа.
Листья из именованного кортежа.
Листья из упорядоченного словаря OrderedDict.
Строка является листом.
Листья из списка внутри словаря.
Массив JAX является листом.
Массив NumPy является листом.
Листья из вложенного списка. None – не лист.
⓬
⓬
⓬
⓭
⓮
⓯
⓯
⓯
⓰
⓱
⓱
⓱
⓲
⓲
⓳
⓳
⓴
Глава 10
374
⓰
⓱
⓲
⓳
⓴
Работа с pytree
Число является листом.
Листья из кортежа.
Листья из именованного кортежа.
Листья из упорядоченного словаря OrderedDict.
Строка является листом.
В приведенном выше примере можно видеть, что типы list, tup
le, dict, namedtuple, OrderedDict и None работают как контейнеры,
а строки, числовые значения и типы массивов Array и ndarray явля
ются листьями.
Многие функции JAX, например jax.lax.scan() или jax.lax.map(),
работают с деревьями pytree. Функции трансформации JAX также
можно применять к функциям, принимающим в качестве входных
данных и возвращающим как результат деревья pytree с массивами
(но не массивы деревьев pytree – это очень важно).
В предыдущих примерах с нейронными сетями для классифи
кации изображений MNIST и Cats vs. Dogs использовалась похожая
структура для хранения параметров модели. Теперь мы знаем, что
это были деревья pytree. Они передаются в функцию для вычисле
ния градиентов, и возвращаемое значение также является pytree
с той же структурой.
Для выделения множества вариантов использования pytree здесь
мы воспроизводим код из листингов с 2.5 по 2.13, работающий с py
tree, содержащим параметры модели. Это неполный пример, по
этому необходимо использовать соответствующий блокнот, если вы
хотите выполнить код на своем компьютере.
Листинг 10.5 Повторение кода тренировки модели из листингов
с 2.5 по 2.13
init_params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
def predict(params, image):
"""Function for per-example predictions."""
# Функция для прогнозов по одному элементу.
activations = image
for w, b in params[:-1]:
outputs = jnp.dot(w, activations) + b
activations = swish(outputs)
final_w, final_b = params[-1]
logits = jnp.dot(final_w, activations) + final_b
return logits
batched_predict = vmap(predict, in_axes=(None, 0))
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
❶
❷
❸
❹
Представление сложных структур данных в форме pytree
375
logits = batched_predict(params, images)
log_preds = logits - logsumexp(logits)
return -jnp.mean(targets*log_preds)
@jax.jit
def update(params, x, y, epoch_number):
shapes = jax.tree_util.tree_map(lambda p: p.shape, params)
print(f"Params shapes: {shapes}")
loss_value, grads = value_and_grad(loss)(params, x, y)
grad_shapes = jax.tree_util.tree_map(lambda p: p.shape, grads)
print(f"Grads shapes: {grad_shapes}")
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)], loss_value
❺
params, loss_value = update(init_params, x, y, 0)
>>> Params shapes: [((512, 784), (512,)), ((10, 512), (10,))]
>>> Grads shapes: [((512, 784), (512,)), ((10, 512), (10,))]
❶
❷
❸
❹
❺
❻
❼
❻
❼
Генерация начальных значений параметров.
Функция для прогноза по одному элементу с использованием параметров модели.
Функция для пакетного прогнозирования.
Функция потерь, использующая пакетное прогнозирование.
Функция для выполнения одного шага обновления градиента.
Параметр pytree перед обновлением.
Дерево pytree с градиентами.
Дерево pytree с параметрами модели передавалось в функции
после вызовов трех функций трансформации. Первой была транс
формация vmap() для создания пакетной версии функции predict().
Затем сработала трансформация value_and_grad() для вычисления
градиентов. Последней выполнялась трансформация jit(). Ни в од
ной из трансформаций не возникло никаких проблем при работе
с pytree в качестве параметра функции.
Вспомним о смысле параметра in_axes для vmap() и pmap(). В лис
тинге 6.12 отмечалось, что этот параметр может работать с вложен
ными контейнерами Python. Теперь вы знаете, что это означало
использование pytree. Параметр in_axes может принимать pytree,
если соответствующий параметр отображаемой функции также
представляет собой pytree. В приведенном выше примере мы ис
пользовали просто None, поскольку не требовалось отображение по
параметрам модели. Но такое свойство позволяет делать некоторые
более сложные вещи, например формирование отдельного набора
весов для каждого элемента в пакете или предоставление отдель
ных отклонений для каждого элемента пакета при использовании
одинаковых весов. Здесь мы не будем рассматривать подобные эк
зотические варианты, но следует помнить о возможности весьма де
тализированного управления параметрами.
Глава 10
376
Работа с pytree
В дополнение к поддержке pytree в трансформациях JAX и вызо
вах стандартной библиотеки существует еще и полезный специа
лизированный пакет для работы с pytree, который мы рассмотрим
в следующем разделе.
10.2 Функции для работы с pytree
Пакет jax.tree_util предоставляет много полезных функций для рабо
ты с pytree, упрощая работу программиста. Например, часто требуется
итеративный проход по pytree с выполнением одной и той же опера
ции для каждого элемента, скажем изменение формы. Для этой цели
существует map-подобная трансформация. Или можно преобразовать
pytree в неструктурированную последовательность для сериализации
параметров модели, а затем выполнить десериализацию и восстано
вить структуру pytree. Функции пакета позволяют проделать каждую
из таких операций и многое другое. В этом разделе описаны наиболее
важные и полезные функции, помогающие работать с pytree.
ПРИМЕЧАНИЕ Ранее ко многим подпрограммам jax.tree_
util можно было получить доступ из пакета верхнего уровня
JAX, например jax.tree_leaves() вместо jax.tree_util.tree_
leaves(). Такие варианты импорта объявлены устаревшими
и нерекомендуемыми к использованию; они будут удалены
в будущем релизе.
10.2.1 Использование tree_map()
Мы уже знакомы с функцией jax.tree_util.tree_leaves(), возвраща
ющей листья дерева pytree. Нам также известна функция jax.tree_
util.tree_map(), которая отображает функцию по аргументам pytree
и создает новое дерево pytree. Мы использовали ее для обхода дере
ва и фиксации форм всех массивов в нем, но аналогичным способом
можно создавать измененные деревья. Например, не составляет ни
какого труда масштабировать любой массив, скажем умножить каж
дый тензор на 10. Но при этом следует всегда помнить: невозможно
вносить изменения в исходный существующий массив (объект); всег
да создается новая измененная версия существующей структуры.
Листинг 10.6
Изменение структуры данных
params = init_network_params(LAYER_SIZES, key, scale=PARAM_SCALE)
scaled_params = jax.tree_util.tree_map(lambda p: 10*p, params)
❶
❷
❶ Генерация некоторых весов нейронной сети.
❷ Масштабирование весов модели посредством умножения каждого веса на 10.
377
Функции для работы с pytree
Так же легко использовать функцию tree_map() для репликации
параметров модели на нескольких устройствах в варианте трени
ровки с распараллеливанием по данным. Нужно лишь немного из
менить код примера тренировки с распараллеливанием по данным
MNIST из раздела 7.3. В исходном варианте на каждом шаге трени
ровки параметры модели в широковещательном режиме распро
странялись по всем устройствам, а затем каждое устройство вы
числяло глобальные градиенты, используя коллективные операции,
поэтому требовались две потенциально крупномасштабные опера
ции обмена информацией между устройствами, связанные с весами
модели (см. рис. 10.2).
update() получает
обновленные параметры
модели с первого
устройства
pmap()
распространяет
параметры по всем
устройствам
Параметры
Параметры
psum() обеспечивает
обмен данными между
устройствами
Загрузка
тензоров
параметров
в каждый TPU
Шаг
вычисления
градиентов
Шаг
обновления
параметров
Рис. 10.2 Схема тренировки с распараллеливанием по данным из главы 7. На каждом шаге
параметры модели передаются в широковещательном режиме на все TPU
Как было отмечено в главе 7, постоянное хранение отдельной ко
пии параметров модели на каждом устройстве и локальное их об
новление с использованием агрегированных градиентов было бы
более эффективным. При таком подходе нужно просто подготовить
pytree с параметрами модели, затем при вызове pmap() указать, что
параметры модели тоже должны отображаться. Эта методика ис
пользует параметр in_axes (см. рис. 10.3).
Для соответствующей модификации параметров модели необхо
димо добавить дополнительные основные оси в каждую матрицу
весов и в каждый вектор отклонений и обеспечить репликацию (ко
пирование) значений весов и отклонений в новое измерение (это
делается только один раз в начале тренировки). Трансформация
pmap() будет отображать новые оси так, что каждое устройство полу
чит собственную полную копию параметров модели.
Глава 10
378
Работа с pytree
pmap() выполняет
отображение
по новой введенной оси
в параметрах
Параметры
update() возвращает
обновленные
реплицированные
параметры
Реплици
рованные
параметры
Реплици
рованные
параметры
Параметры
Параметры
Параметры
Репликация
вручную
каждого
параметра
в pytree
Параметры
psum()
обеспечивает обмен
данными между
устройствами
Параметры
Теперь каждый
тензор содержит
дополнительную ось
с размером, равным
количеству TPU
Параметры
Каждый TPU
содержит
собственную часть
реплицированных
параметров pytree
Шаг
вычисления
градиентов
Шаг
обновления
параметров
Теперь каждый
тензор содержит
дополнительную
ось с размером,
равным
количеству TPU
Рис. 10.3 Схема оптимизированного процесса тренировки с распараллеливанием
по данным. Параметры модели вручную реплицируются перед тренировкой, поэтому
в широковещательном распространении их на каждом шаге нет необходимости
Для репликации значений по новым созданным осям воспользу
емся функцией jnp.broadcast_to() (https://docs.jax.dev/en/latest/_au
tosummary/jax.numpy.broadcast_to.html). В листинге 10.7 представлен
неполный вариант, включающий только самые важные части кода,
которые отличаются от листингов с 7.26 по 7.29.
Листинг 10.7
Внесение изменений в процесс тренировки из главы 7
from jax.tree_util import tree_map
init_params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
replicate_array = lambda x: jnp.broadcast_to(
x, (NUM_DEVICES,) + x.shape)
replicated_params = tree_map(replicate_array, init_params)
❶
❷
❸
❸
@partial(jax.pmap, axis_name='devices', in_axes=(0, 0, 0, None))
❹
def update(params, x, y, epoch_number):
loss_value, grads = value_and_grad(loss)(params, x, y)
grads = [(jax.lax.psum(dw, 'devices'), jax.lax.psum(db, 'devices'))
for dw, db in grads]
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)], loss_value
for epoch in range(NUM_EPOCHS):
Функции для работы с pytree
379
start_time = time.time()
losses = []
for x, y in train_data:
num_elements = len(y)
x = jnp.reshape(x, (NUM_DEVICES, num_elements//NUM_DEVICES, NUM_PIXELS))
y = jnp.reshape(
one_hot(y, NUM_LABELS),
(NUM_DEVICES, num_elements//NUM_DEVICES, NUM_LABELS))
replicated_params, loss_value = update(
❺
replicated_params, x, y, epoch)
losses.append(jnp.sum(loss_value))
epoch_time = time.time() - start_time
❻
params = tree_map(lambda x: x[0], replicated_params)
train_acc = accuracy(params, train_data)
test_acc = accuracy(params, test_data)
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
print("Training set loss {}".format(jnp.mean(jnp.array(losses))))
print("Training set accuracy {}".format(train_acc))
print("Test set accuracy {}".format(test_acc))
❶ Импорт tree_map().
❷ Генерация начальных параметров модели.
❸ Многократная репликация параметров нейронной сети, чтобы каждое устройство
получило собственную копию с одинаковой структурой.
❹ Использование pmap() для распараллеливания функции и отображения ее пер-
вых трех параметров.
❺ Использование реплицированных вручную параметров.
❻ Извлечение параметров только из первого устройства для вычисления точности
на одном устройстве.
Обратите внимание на то, как здесь реплицируются параметры
модели. Мы вводим новую основную ось с размером, равным ко
личеству устройств, и копируем (или распространяем в широко
вещательном режиме) параметры модели в новое измерение. Это
делается для каждого листа в pytree, хранящего параметры модели,
с использованием функции tree_map().
Весьма существенное, но малозаметное изменение заключа
ется в том, что мы используем in_axes=(0, 0, 0, None) вместо in_
axes=(None, 0, 0, None) из листинга 7.27. Это означает, что мы не
полагаемся на автоматическую репликацию, поскольку вручную
реплицировали параметры модели и теперь их необходимо отобра
зить. Кроме того, мы отказываемся от out_axes=(None, 0), потому
что нет необходимости в возврате обновленных параметров модели
только из первого устройства. Сейчас нам вообще не нужен возврат
параметров модели, поскольку они хранятся на каждом устройстве
отдельно (благодаря сегментированию массива). Единственное мес
то, где извлекаются параметры модели из первого устройства, – вы
числение точности. Это делается с использованием только одного
380
Глава 10
Работа с pytree
устройства, но вы можете переписать эту часть, чтобы использовать
несколько устройств.
10.2.2 Преобразование pytree в плоскую структуру
и восстановление древовидной формы
В реальной практике необходим обмен информацией с внешними
библиотеками. Во многих библиотеках допускается использова
ние только простых списков или одномерных массивов в качестве
входных данных. Кроме того, может потребоваться применение вы
сокопроизводительных функций, оптимизированных для работы
с одномерными массивами и списками. Возможно, возникнет не
обходимость в сохранении параметров модели во внешнем храни
лище или в файле. Во всех перечисленных случаях, вероятнее всего,
станет невозможной работа непосредственно с pytree и потребует
ся способ преобразования древовидной структуры в более простую
с возможностью восстановления ее исходной формы.
Существует набор полезных функций для преобразования pytree
в плоскую структуру и восстановления древовидной формы: jax.
tree_util.tree_flatten(), jax.tree_util.tree_unflatten(), jax.
flatten_util.ravel_pytree(), а также дополняющая комплект функ
ция jax.tree_util.tree_structure().
Функция jax.tree_util.tree_structure() возвращает структуру
pytree, представленную типом PyTreeDef. Функция jax.tree_util.
tree_flatten() возвращает плоскую (спрямленную) версию pytree,
используя детерминированную процедуру, соответствующую алго
ритму обхода дерева слева направо и преимущественно в глубину.
Функция также возвращает тип PyTreeDef, представляющий струк
туру уплощенного дерева. Такой же результат позволяет получить
функция tree_structure().
В обеих функциях разрешен необязательный параметр is_leaf,
определяющий функцию с возвращаемым логическим значением,
которая будет вызываться на каждом шаге преобразования pytree
в плоскую структуру. Если функция возвращает True, то обход дере
ва останавливается и все поддерево в целом интерпретируется как
лист. При возврате значения False процедура уплощения продолжа
ет обход текущего объекта.
Функция jax.tree_util.tree_unflatten() выполняет обратную
процедуру. При наличии структуры дерева в PyTreeDef и итериру
емых листьев (что было обеспечено функцией tree_flatten()) она
восстанавливает исходное дерево pytree.
Можно имитировать работу tree_map(), сначала преобразовав де
рево в плоскую структуру, а затем выполнить обычное отображение
Python по итерируемым листьям, после чего восстановить дерево
с помощью функции tree_unflatten().
381
Функции для работы с pytree
Листинг 10.8 Преобразование pytree в плоскую форму
и восстановление исходной структуры
some_pytree = [
[1,1,1],
[
[10,10,10], [20, 20]
]
]
❶
jax.tree_util.tree_map(lambda p: p+1, some_pytree)
❷
>>> [[2, 2, 2], [[11, 11, 11], [21, 21]]]
leaves, struct = jax.tree_util.tree_flatten(some_pytree)
leaves
>>> [1, 1, 1, 10, 10, 10, 20, 20]
struct
>>> PyTreeDef([[*, *, *], [[*, *, *], [*, *]]])
updated_leaves = map(lambda x: x+1, leaves)
jax.tree_util.tree_unflatten(struct, updated_leaves)
>>> [[2, 2, 2], [[11, 11, 11], [21, 21]]]
❶
❷
❸
❹
❺
❻
❼
❸
❹
❺
❻
❼
Создание pytree.
Обработка pytree с использованием tree_map().
Преобразование pytree в плоскую структуру.
Значения, содержащиеся в полученной структуре.
Объект содержит структуру pytree.
Отображение значений.
Восстановление pytree по заданной структуре.
Обычно нет необходимости в использовании комбинации про
цедур flatten/unflatten для обработки дерева, так как функция
tree_map() великолепно выполняет эту работу. Тем не менее может
потребоваться применение функций tree_flatten()/tree_unflatten() при сериализации/десериализации сложных структур для со
хранения в некотором внешнем хранилище или для обмена между
различными процессами.
В приведенном выше примере (листинг 10.8) можно заметить,
что функция tree_flatten() создает обычный список Python. Если
вы работаете с массивами, а не со списками, то, вероятно, предпо
чтете получить одномерный массив после процедуры уплощения.
Это обычный вариант для разнообразных библиотек поддержки ма
шинного обучения, поскольку они в основном используют тензоры
Глава 10
382
Работа с pytree
и массивы, а не списки. Для подобных случаев существует отдельная
функция jax.flatten_util.ravel_pytree(). Она возвращает одно
мерный массив, представляющий уплощенные и объединенные
значения листьев, а кроме того, вместо структуры PyTreeDef воз
вращает функцию для восстановления одномерного вектора той же
длины обратно в исходную структуру pytree.
Листинг 10.9 Уплощение и восстановление структуры pytree с использованием
одномерного массива
from jax.flatten_util import ravel_pytree
leaves, unflatten_func = ravel_pytree(some_pytree)
leaves
>>> Array([ 1,
1,
1, 10, 10, 10, 20, 20], dtype=int32)
unflatten_func
>>> <function jax._src.flatten_util.ravel_pytree.
➥<locals>.<lambda>(flat)>
unflatten_func(leaves)
❶
❷
❸
❹
>>> [[Array(1, dtype=int32), Array(1, dtype=int32), Array(1, dtype=int32)],
>>> [[Array(10, dtype=int32), Array(10, dtype=int32), Array(10, dtype=int32)],
>>>
[Array(20, dtype=int32), Array(20, dtype=int32)]]]
❶
❷
❸
❹
Импорт необходимых функций.
Преобразование pytree в одномерный массив.
В одномерном массиве содержатся значения листьев pytree.
Функция для восстановления pytree.
В приведенном выше примере мы начали со структуры на осно
ве списка, т. е. с той же структуры, которая использовалась в приме
ре уплощения/восстановления. Поскольку функция ravel_pytree()
ориентирована на одномерные массивы, она преобразовывает типы
входных элементов в подходящий одномерный массив (подроб
ности преобразования (приведения) типов см. здесь: https://docs.
jax.dev/en/latest/type_promotion.html). Другими словами, комби
нация unravel(ravel(x)) не является идемпотентной, в отличие от
unflatten(flatten(x)). Поэтому требуется особая внимательность
при использовании действительно сложных структур, которые объ
единены в различные типы. Все они будут приведены к единому
типу. Но в приложениях машинного обучения такая ситуация воз
никает чрезвычайно редко.
Функции для работы с pytree
383
10.2.3 Использование tree_reduce()
Иногда необходимо применить функцию ко всем листьям pytree
и вернуть единственное значение, например при вычислении сум
мы pytree. Один из возможных вариантов использования – вычис
ление некоторой агрегации весов нейросети и добавление получен
ного значения как штрафного члена (слагаемого) в функцию потерь.
В такой ситуации может оказаться полезной функция jax.tree_
util.tree_reduce(). Она в общих чертах похожа на обычную функ
цию редукции reduce, широко применяемую в функциональном
программировании. Внутри используется стандартная функция Py
thon functools.reduce() (https://docs.python.org/3/library/functools.
html#functools.reduce) для работы с листьями pytree.
Функция jax.tree_util.tree_reduce() применяет другую функ
цию с двумя аргументами кумулятивно (с нарастающим итогом)
для итеративного прохода по списку листьев слева направо, начи
ная с необязательного начального значения (или первого элемента
итерации). В конце процедуры функция сводит (редуцирует) листья
к одному значению. Например, таким способом легко вычисляется
сумма всех значений в pytree.
Листинг 10.10
Редукция pytree
some_pytree = [
[1,1,1],
[
[10,10,10], [20, 20]
]
]
❶
jax.tree_util.tree_reduce(lambda acc,value: acc+value, some_pytree,
❷
initializer=0)
>>> 73
❶ Создание pytree.
❷ Редукция pytree с использованием функции для суммирования значений с на-
коплением.
В приведенном выше примере передается функция суммирования
значений с накоплением по всем итерируемым элементам, начиная
с некоторого начального значения, в данном случае с 0. Функция
проходит по всем листьям и для каждого листа вызывается с теку
щим накапливаемым значением (изначально с нулем) и со значе
нием листа. Оба значения суммируются, затем результат становится
накапливаемым значением для следующего вызова функции.
Глава 10
384
Работа с pytree
10.2.4 Транспонирование pytree
Иногда требуется транспонирование pytree, т. е. преобразование py
tree с некоторой иерархической структурой типа (outer, inner) в py
tree с транспонированной структурой типа (inner, outer). Например,
список деревьев pytree можно преобразовать в дерево pytree спис
ков. Рассмотрим конкретный пример.
Предположим, что имеется структура, представляющая точку
в двумерном пространстве. Пусть это будет именованный кортеж
с полями x и y. Также имеется массив таких точек и требуется при
менить функцию для поворота относительно начала координат каж
дой точки в массиве.
Листинг 10.11
Работа с точками в двумерном пространстве
import math
from collections import namedtuple
Point = namedtuple('Point', ['x', 'y'])
points = [
Point(0.0, 0.0),
Point(3.0, 0.0),
Point(0.0, 4.0)
]
❶
❷
def rotate_point(p, theta):
x = p.x * math.cos(theta) - p.y * math.sin(theta)
y = p.x * math.sin(theta) + p.y * math.cos(theta)
return Point(x,y)
❸
rotate_point(points[1], math.pi)
❹
>>> Point(x=-3.0, y=3.6739403974420594e-16)
❶
❷
❸
❹
Создание структуры типа именованный кортеж.
Создание списка точек.
Функция для поворота точки.
Поворот одной точки.
Самый простой способ – обработка каждой точки вручную, но мы
помним, что в JAX имеется функция автоматической векторизации,
поэтому воспользуемся ее мощью.
Кажущийся очевидным подход с использованием vmap() работать
не будет, потому что vmap() поддерживает деревья pytree в форме
«структура массивов», а не «массив структур». Причина в том, что
форму «массив структур» невозможно обработать эффективно, изза чего, в свою очередь, XLA поддерживает только числовые массивы
типов dtype. Поэтому невозможно применить vmap() для вектори
385
Функции для работы с pytree
зации по массиву структур. Необходимо реорганизовать точки так,
чтобы получить структуру массивов, т. е. должна получиться одна
структура, содержащая массивы координат всех точек.
Листинг 10.12 Попытка применения vmap() к списку двумерных точек
jax.vmap(rotate_point, in_axes=(0, None))(points, math.pi)
❶
>>> ...
>>> ValueError: vmap was requested to map its argument along axis 0, which
implies that its rank should be at least 1, but is only 0 (its shape is ())
# ValueError: функция vmap получила запрос на отображение ее аргумента
# по оси 0 с предположением, что ее ранг должен быть не менее 1, но
# ранг равен только 0 (ее форма ())
❶ Предполагаем, что необходимо отображение по массиву, но этот способ не рабо-
тает.
Необходимо преобразовать имеющийся список структур в струк
туру массивов. Для решения этой задачи используется функция jax.
tree_util.tree_transpose(). Здесь внешней структурой является
список, а внутренней – структура, содержащая одну точку.
Листинг 10.13 Просмотр внутренней и внешней структур
jax.tree_util.tree_structure(points)
❶
>>> PyTreeDef([CustomNode(namedtuple[Point], [*, *]),
CustomNode(namedtuple[Point], [*, *]), CustomNode(namedtuple[Point], [*, *])])
jax.tree_util.tree_structure(points[0])
❷
>>> PyTreeDef(CustomNode(namedtuple[Point], [*, *]))
❶ Возврат внешней структуры дерева.
❷ Возврат внутренней структуры дерева.
Уточнение: когда мы говорим, что возвращается внешняя струк
тура в рассматриваемом здесь примере, то в действительности воз
вращаются обе структуры, поскольку выводится подробное описа
ние как внешнего списка, так и внутреннего именованного кортежа.
Для извлечения только внешней структуры необходимо реплици
ровать этот список без подробностей о его элементах.
Листинг 10.14
Просмотр только внешней структуры
jax.tree_util.tree_structure([0 for p in points])
>>> PyTreeDef([*, *, *])
❶ Возврат внешней структуры дерева.
❶
Глава 10
386
Работа с pytree
Предоставив функции tree_transpose() все необходимые пара
метры, мы можем «транспонировать» существующую исходную
структуру в требуемую форму.
Листинг 10.15
Транспонирование pytree
points_t = jax.tree_util.tree_transpose(
outer_treedef = jax.tree_util.tree_structure(
[0 for p in points]),
inner_treedef = jax.tree_util.tree_structure(points[0]),
pytree_to_transpose=points
)
❶
❷
points_t
>>> Point(x=[0.0, 3.0, 0.0], y=[0.0, 0.0, 4.0])
❶ Передача внешней структуры дерева.
❷ Передача внутренней структуры дерева
❸ Получение транспонированной структуры.
❸
Еще раз обратите внимание на то, как мы передали структуру де
рева для внешнего уровня. Передается упрощенный список без ка
ких-либо деталей внутренней структуры элементов. Если бы пере
давался полный исходный список, то он содержал бы и внешние,
и внутренние уровни, и вызов функции привел бы к ошибке.
Осталось сделать только один шаг. Мы преобразовали «список
структур» в «структуру из списков», но vmap() не работает со списка
ми – только с массивами. Поэтому последний шаг очень прост – пре
образование списков в массивы.
Листинг 10.16
Преобразование списков в массивы
points_t_array = Point(jnp.array(points_t.x),jnp.array(points_t.y))
points_t_array
❶
❷
>>> Point(x=Array([0., 3., 0.], dtype=float32), y=Array([0., 0., 4.],
dtype=float32))
❶ Преобразование двух списков в два массива.
❷ Вывод структуры, теперь содержащей массивы.
Теперь можно применить vmap() к преобразованной структуре
(более правильно: vmap() применяется к функции, обрабатывающей
один элемент, а затем трансформированная функция применяется
к преобразованной структуре).
Создание специализированных узлов pytree
Листинг 10.17
387
Успешное применение векторизованной функции
jax.vmap(rotate_point, in_axes=(0, None))(points_t_array, math.pi)
>>> Point(x=Array([-0.0000000e+00, -3.0000000e+00,
➥-4.8985874e-16], dtype=float32),
➥y=Array([ 0.0000000e+00, 3.6739406e-16,
➥-4.0000000e+00], dtype=float32))
❶ Векторизация и применение функции к преобразованной структуре.
❷ Вывод результата успешного применения функции.
❶
❷
Мы успешно завершили работу, и можно видеть, что точки по
вернулись на 180 градусов, а их координаты поменяли знаки на
противоположные. Такое перемещение может оказаться весьма по
лезным.
Мы рассмотрели работу со многими полезными функциями из па
кета jax.tree_util. Единственным важным набором функций, с ко
торым мы пока еще не познакомились, является комплект средств
для создания специализированных контейнеров. Это тема следую
щего раздела.
10.3 Создание специализированных узлов pytree
Возможно, у вас имеются собственные классы, содержащие данные,
с которыми желательно работать, как с контейнерами, а не листья
ми. Превосходными примером являются классы, представляющие
слои нейронной сети.
Следует отметить, что это учебный пример. В реальной практике
высококачественные библиотеки уже решили эту задачу. Поэтому
в том случае, если вы используете классы данных (dataclasses) (ко
торые являются листьями по умолчанию) и хотите сделать их узла
ми pytree, вам поможет превосходная библиотека Chex (https://chex.
readthedocs.io/en/latest/api.html#dataclasses) (с добавлением всего
лишь одной аннотации). Если вы пользуетесь Flax (тема главы 11)
для создания нейронных сетей, то для аналогичных целей имеет
ся аннотация flax.struct.dataclass (https://flax.readthedocs.io/en/
latest/api_reference/flax.struct.html). В качестве альтернативного ре
шения также предлагается equinox.Module (https://docs.kidger.site/
equinox/api/module/module/) из небольшой, удобной и чрезвычайно
полезной библиотеки Equinox. Но здесь я покажу вам, как сделать
это самостоятельно.
Предположим, что имеется класс, представляющий линейный
слой нейронной сети, хранящий имя слоя вместе с матрицей весов
и вектором отклонений.
Глава 10
388
Работа с pytree
Листинг 10.18 Специализированный класс для линейного уровня
нейронной сети
class Layer:
def __init__(self, name, w, b):
self.w = w
self.b = b
self.name = 'name'
❶ Матрица с весами.
❷ Вектор с отклонениями.
❸ Имя слоя.
❶
❷
❸
Ничего особенного – простая структура, содержащая некоторые
взаимосвязанные данные.
Можно создать pytree, включающее новый созданный слой и дру
гие листья с некоторыми данными.
Листинг 10.19 Дерево pytree со специализированным классом
h1 = Layer('hidden1', jnp.zeros((100,20)), jnp.zeros((20,)))
pt = [
jnp.ones(50),
h1
]
jax.tree_util.tree_leaves(pt)
❶
❷
❸
>>> [Array([1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1.,
>>>
1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1.,
>>>
1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1., 1.],
>>>
dtype=float32),
>>> <__main__.Layer at 0x7fc87db04490>]
❶ Создание слоя.
❷ Создание pytree.
❸ Вывод листьев дерева pytree.
В приведенном выше примере можно видеть, что объект слоя
интерпретируется как отдельный лист. Возможно, это правиль
ный подход, но нам известно, что в действительности он является
контейнером для некоторых других данных, и мы предполагаем
применение некоторых функций к этим данным, используя tree_
map(). Если мы попытаемся, например, инкрементировать или вы
полнить операцию умножения для всех значений в таком pytree, то
возникнет ошибка, так как наш класс не поддерживает требуемую
операцию.
389
Создание специализированных узлов pytree
Листинг 10.20 Применение tree_map() к pytree
со специализированным классом
jax.tree_map(lambda x: x*10, pt)
❶
>>> ...
>>> TypeError: unsupported operand type(s) for *: 'Layer' and 'int' ❷
# TypeError: неподдерживаемый тип(ы) операнда для *: 'Layer' и 'int'
❶ Попытка умножения всех значений на 10.
❷ Возвращается ошибка.
Проблему можно решить, реализовав все требуемые операции для
созданного класса. Но есть и другой вариант. Можно зарегистриро
вать класс как контейнер и сообщить JAX, что находится внутри кон
тейнера и как работать с этими данными.
Для этого необходимо предоставить JAX информацию о том, как
преобразовать контейнер в плоскую форму и восстановить его ис
ходную форму, предоставив две соответствующие функции.
Функция уплощения возвращает (1) итерируемый объект для
потомков, подлежащих уплощению, и (2) необязательные вспомо
гательные данные, сохраняемые в определении дерева, но не явля
ющиеся потомками (фактически это некоторые метаданные). В на
шем случае желательно представить веса и отклонения как листья,
но не хотелось бы видеть имя слоя как лист, поэтому оно становится
явным кандидатом во вспомогательные данные.
Функция восстановления формы принимает два аргумента –
вспомогательные данные и уплощенные потомки – и восстанавли
вает исходный объект. Напишем обе функции для класса Layer.
Листинг 10.21 Функции уплощения и восстановления формы
для специализированного класса
def flatten_layer(container):
flat_contents = [container.w, container.b]
aux_data = container.name
return flat_contents, aux_data
def unflatten_layer(aux_data, flat_contents):
return Layer(aux_data, *flat_contents)
❶ Упаковка элементов данных в плоский список.
❷ Возврат имени слоя как вспомогательных данных.
❸ Восстановление исходного объекта.
❶
❷
❸
В рассматриваемом здесь примере мы упаковали веса и отклоне
ния в плоский (одномерный) список. Имя слоя не включено в этот
список, так как нежелательно видеть его как отдельный лист в дере
Глава 10
390
Работа с pytree
ве pytree, поэтому оно возвращается как вспомогательные данные.
Функция восстановления формы выполняет обратную операцию:
принимает вспомогательные данные и плоский список потомков
и восстанавливает исходный объект.
Осталось только зарегистрировать этот контейнер в реестре кон
тейнеров pytree. Для этого используется функция jax.tree_util.
register_pytree_node(), принимающая тип Python для интерпрета
ции его как внутреннего узла pytree и две функции для уплощения
и восстановления формы (см. листинг 10.22).
Листинг 10.22 Регистрация созданного контейнера в реестре контейнеров
pytree
jax.tree_util.register_pytree_node(
Layer, flatten_layer, unflatten_layer)
❶
jax.tree_util.tree_leaves(pt)
❷
>>> [Array([1., 1., 1., 1., 1., 1., 1., 1.,
>>>
1., 1., 1., 1., 1., 1., 1., 1.,
>>>
1., 1., 1., 1., 1., 1., 1., 1.,
>>>
dtype=float32),
>>> Array([[0., 0., 0., ..., 0., 0., 0.],
>>>
[0., 0., 0., ..., 0., 0., 0.],
>>>
[0., 0., 0., ..., 0., 0., 0.],
>>>
...,
>>>
[0., 0., 0., ..., 0., 0., 0.],
>>>
[0., 0., 0., ..., 0., 0., 0.],
>>>
[0., 0., 0., ..., 0., 0., 0.]],
>>> Array([0., 0., 0., 0., 0., 0., 0., 0.,
>>>
0., 0., 0.], dtype=float32)]
1., 1., 1., 1., 1., 1., 1., 1., 1.,
1., 1., 1., 1., 1., 1., 1., 1., 1.,
1., 1., 1., 1., 1., 1., 1., 1.],
dtype=float32),
0., 0., 0., 0., 0., 0., 0., 0., 0.,
pt2 = jax.tree_map(lambda x: x+1, pt)
>>> [Array([2., 2., 2., 2., 2., 2., 2., 2.,
>>>
2., 2., 2., 2., 2., 2., 2., 2.,
>>>
2., 2., 2., 2., 2., 2., 2., 2.,
>>>
dtype=float32),
>>> Array([[1., 1., 1., ..., 1., 1., 1.],
>>>
[1., 1., 1., ..., 1., 1., 1.],
>>>
[1., 1., 1., ..., 1., 1., 1.],
>>>
...,
>>>
[1., 1., 1., ..., 1., 1., 1.],
>>>
[1., 1., 1., ..., 1., 1., 1.],
>>>
[1., 1., 1., ..., 1., 1., 1.]],
>>> Array([1., 1., 1., 1., 1., 1., 1., 1.,
>>>
1., 1., 1.], dtype=float32)]
❸
2., 2., 2., 2., 2., 2., 2., 2., 2.,
2., 2., 2., 2., 2., 2., 2., 2., 2.,
2., 2., 2., 2., 2., 2., 2., 2.],
dtype=float32),
1., 1., 1., 1., 1., 1., 1., 1., 1.,
❶ Регистрация контейнера.
❷ Теперь мы видим новые листья в pytree.
❸ Вызов tree_map() успешно изменяет данные внутри контейнера.
Резюме
391
Вот и все. После регистрации контейнера можно применять tree_
map() к данным внутри него, а листья pytree теперь содержат эле
менты данных из внутренней части контейнера, а не просто объект
класса.
Мы закончили изучение работы с pytree. Вы будете часто поль
зоваться этими деревьями, поскольку для многих реальных задач
требуются сложные иерархические структуры хранения данных.
В области глубокого обучения почти все параметры модели хранятся
именно таким способом.
Резюме
Pytree – это вложенная древовидная структура, формируемая из
контейнерных объектов Python.
Pytree состоит из листьев и узлов.
Узел – это само дерево pytree, и его можно представить контейнер
ным объектом Python. Такие контейнерные типы регистрируются
в реестре контейнеров pytree, и к ним относятся list, tuple, dict,
namedtuple, OrderedDict и значение None, принятое по умолчанию.
Все прочие типы по умолчанию являются листьями и включают
числовые значения, а также типы Array и ndarray.
Многие библиотечные функции JAX и функции трансформации
могут работать с деревьями pytree.
Пакет jax.tree_util предоставляет множество полезных функций
для работы с pytree.
Функция jax.tree_util.tree_map() отображает некоторую другую
функцию на аргументы pytree и создает новое дерево pytree.
Функция jax.tree_util.tree_leaves() возвращает листья дерева
pytree.
Функция jax.tree_util.tree_structure() возвращает структуру
дерева pytree, представленную типом PyTreeDef.
Функция jax.tree_util.tree_flatten() возвращает уплощенную
версию pytree в списке, а PyTreeDef представляет структуру упло
щенного дерева.
Функция jax.tree_util.tree_unflatten() выполняет обратную
операцию: с помощью структуры дерева, хранящейся в PyTreeDef,
и итерируемого списка листьев (созданного функцией tree_flatten()) восстанавливает структуру исходного дерева pytree.
Функция jax.flatten_util.ravel_pytree() возвращает одномер
ный массив, представляющий уплощенные и объединенные зна
чения листьев, а кроме того, вместо структуры PyTreeDef возвра
щает функцию для восстановления одномерного вектора.
Функция jax.tree_util.tree_reduce() похожа на обычную функ
цию редуцирования (reduce), широко применяемую в функцио
392
Глава 10
Работа с pytree
нальном программировании. Она применяет функцию с двумя
аргументами в накопительном режиме к итерируемому списку
листьев слева направо, редуцируя данные листьев до единствен
ного значения.
Функция jax.tree_util.tree_transpose() выполняет преобразо
вание pytree с некоторой иерархической структурой типа (outer,
inner) в pytree с транспонированной структурой типа (inner, outer).
Можно
зарегистрировать собственный специализированный
класс контейнера в реестре контейнеров JAX, чтобы сообщить JAX
о том, что находится внутри этого контейнера и как работать с его
данными.
Для регистрации специализированного контейнера pytree не
обходимо сообщить JAX, как выполнять уплощение и восстанов
ление его исходной формы, предоставив две соответствующие
функции и воспользовавшись функцией jax.tree_util.register_
pytree_node().
Часть III
Экосистема
Ч
асть III представляет чрезвычайно активную экосистему, окру
жающую JAX, распространяющую библиотеки и инструментальные
средства, расширяющие функциональность этого фреймворка в раз
личных областях глубокого обучения и других сферах деятельности.
Здесь наглядно демонстрируется, как JAX вписывается в более ши
рокий контекст, позволяя использовать его мощь в совокупности
с другими специализированными библиотеками. В двух следующих
главах вы узнаете о высокоуровневых библиотеках поддержки ней
ронных сетей, которые упрощают создание и тренировку модели,
а также о других членах экосистемы JAX, удовлетворяющих широ
кий спектр потребностей в области научных исследований и вычис
лений.
В главе 11 рассматриваются библиотеки поддержки нейронных
сетей и особое внимание уделено Flax и Linen API в сочетании с Op
tax для трансформаций градиентов. Вы научитесь создавать модели
более осмысленно, управлять состояниями тренировки и взаимо
действовать с экосистемой Hugging Face для доступа к самым со
временным моделям. Эта глава устраняет разрыв между основными
функциональными средствами JAX и реальными практическими
потребностями проектов глубокого обучения, предоставляя поль
зователю инструменты для эффективного создания, тренировки
и развертывания сложных моделей.
Глава 12 дает более широкое представление об экосистеме JAX
с описанием библиотек для разнообразных задач машинного обуче
ния, включая обучение с подкреплением и эволюционные вычисле
394
Экосистема
ния. Также рассматриваются модули JAX для других областей науки,
таких как физика, химия и т. д.
После изучения части III вы получите исчерпывающее представ
ление об экосистеме JAX и будете готовы к практическому примене
нию JAX для решения широкого диапазона интересных и сложных
задач.
11
Высокоуровневые
библиотеки поддержки
нейронных сетей
Темы главы:
создание многослойного перцептрона (MLP) для
классификации рукописных цифр MNIST с использованием
библиотеки Flax и ее Linen API;
использование библиотеки трансформации градиентов
Optax для тренировки модели;
использование TrainState dataclass для представления
состояния тренировки и хранения метрик;
создание остаточной нейронной сети для классификации
изображений и работы с переменными состояния модели;
использование библиотек Hugging Face в сочетании
с трансформерами и диффузорами JAX/Flax.
Ядро JAX – это мощная, но относительно низкоуровневая библиоте
ка. Точно так же, как редко вы будете создавать сложную нейронную
сеть, применяя исключительно библиотеку NumPy или базисные
элементы TensorFlow, вы не ограничитесь только лишь функцио
нальными возможностями JAX. Существуют высокоуровневые биб
лиотеки поддержки нейронных сетей для TensorFlow (Keras, Sonnet)
и для PyTorch (torch.nn, Pytorch Lightning, fast.ai), но, разумеется, по
добные библиотеки имеются и для JAX.
396
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
Одной из наиболее широко известных высокоуровневых библио
тек поддержки нейронных сетей является Flax компании Google.
Ранее компания DeepMind предлагала библиотеку Haiku, но в на
стоящее время эта библиотека находится в режиме сопровождения,
поэтому компания Google DeepMind рекомендует применять в но
вых проектах Flax вместо Haiku. Кроме того, существует абсолютно
новая версия Keras 3.0 с поддержкой множественных внутренних
компонентов и Equinox. В этой главе мы будем использовать Flax,
но другие библиотеки имеют с ней много общего, поэтому освоение
любой другой библиотеки не составит особого труда после близко
го знакомства с Flax. Все библиотеки предоставляют базисные эле
менты высокого уровня, помогающие создавать нейронные сети из
существующих блоков, таких как плотные, конволюционные (свер
точные) или LSTM1 слои, слои самовнимания со множественными
путями, функции активации и т. п.
Если вы не нуждаетесь в чем-то особенном и чрезвычайно спе
циализированном, то почти все ваши основные потребности будут
покрыты этими стандартными блоками. Вы сэкономите огромное
количество времени, многократно применяя такие тщательно про
тестированные и оптимизированные решения, а кроме того, умень
шите потенциальную опасность внесения ошибок при собственных
разработках.
Зачем вообще связываться с высокоуровневыми JAX-библиоте
ками поддержки нейронных сетей, если уже имеются TensorFlow
и PyTorch? Ответ: методика компонуемых трансформаций функций
в JAX помогает сделать код более удобным в сопровождении, мас
штабируемым и высокопроизводительным.
Главный принцип Flax – предоставить API, знакомый всем, кто
имеет опыт работы с Keras, PyTorch или Sonnet. API в Flax называет
ся Linen. В своей основе это функциональная система для определе
ния нейронных сетей в JAX, отличающаяся от большинства объект
но ориентированных методик, принятых в экосистемах TensorFlow
и PyTorch.
Flax поддерживает общую философию экосистемы JAX с ис
пользованием отдельных независимых библиотек с надежным со
провождением. Библиотека Flax спроектирована так, чтобы с лег
костью интегрироваться в другие части экосистемы. Например,
она использует Optax, отдельную библиотеку для оптимизаторов,
а в более широком смысле – для компонуемых трансформаций
градиентов.
1
LSTM – long short-term memory – сеть с долговременной и кратковремен
ной памятью; разновидность архитектуры рекуррентных нейронных се
тей, предложенная в 1997 году Зеппом Хохрайтером и Юргеном Шмидху
бером. – Прим. перев.
Классификация изображений MNIST
397
В этой главе мы начнем с рассмотрения уже знакомого примера
классификации изображений MNIST и постепенно перепишем его
с использованием Flax, чтобы лучше понять его основы. В следую
щем разделе мы создадим более современный и сложный пример
остаточной нейронной сети и изучим более продвинутые функцио
нальные средства Flax.
Кроме всего прочего, Flax является третьей по распространенно
сти библиотекой (после PyTorch и TensorFlow) среди моделей, публи
куемых на сайте компании Hugging Face (https://huggingface.co/mode
ls?library=jax&sort=downloads). Hugging Face предоставляет хаб с са
мыми современными моделями для разнообразных задач обработ
ки естественного языка (NLP) и обработки изображений, поэтому,
вооружившись Flax, вы сможете быстро начать работу с наилучши
ми доступными большими языковыми моделями (LLM) и диффузи
онными моделями, используя предварительно натренированные
модели, настраивая их более детально и тренируя собственные мо
дели с нуля. В последнем разделе мы используем модели Hugging
Face для диффузии (размытия) изображений и генерации текста.
11.1 Классификация изображений MNIST
с использованием многослойного
перцептрона
Начнем с классификации изображений MNIST с использованием
многослойного перцептрона (MLP). Мы потратили уже достаточно
много времени на разработку различных версий этого решения, на
чиная с простого примера в главе 2, и дошли до распараллеленных
версий в главах 7 и 8. Теперь пришло время узнать, как будет вы
глядеть та же нейросеть, если воспользоваться Flax вместо чистого
ядра JAX.
Сначала рассмотрим самый простой способ перехода к Flax – за
мена определения нейронной сети и использование функции predict(), применяющей нейронную сеть к некоторым входным дан
ным. На следующем этапе мы увеличим степень использования Flax
в примере, применив класс данных dataclass, представляющий пол
ное состояние тренировки, и внешний оптимизатор из библиотеки
Optax.
11.1.1 Многослойный перцептрон в Flax
Проще всего начать с замены определения нейросети слоями, пре
доставленными библиотекой Flax, или более точно – ее интерфей
сом Linen API.
398
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
Linen API
Linen, или flax.linen, – это API нейронной сети второго поколения.
Предыдущим API был flax.nn.
Linen предоставляет абстракцию Module, поведение которой очень похоже на поведение чистых объектов языка Python. Linen Module API является стабильным интерфейсом и рекомендуется к применению в новых проектах. Авторы стремятся предоставить API, хорошо знакомый
всем, кто имеет опыт работы с Keras/Sonnet/PyTorch. В то же время Linen
в своей основе является функциональной системой для определения
нейронных сетей в JAX.
Описание основных принципов («философии») Flax: https://flax-linen.
readthedocs.io/en/latest/philosophy.html. Информацию о целях проектного решения Linen можно найти здесь: https://github.com/google/
flax/blob/main/flax/linen/README.md.
Сначала воспроизведем чистый код JAX из главы 2, где вручную
инициализировались все тензоры весов и отклонений (листинг 2.4)
и была написана функция predict() (листинг 2.5), выполняющая все
необходимые операции умножения матриц, функции активации
и т. п.
Листинг 11.1 Инициализация и применение MLP в чистом коде JAX
(воспроизведение листингов 2.4 и 2.5)
from jax import random
import jax.numpy as jnp
from jax.nn import swish
LAYER_SIZES = [28*28, 512, 10]
PARAM_SCALE = 0.01
❶
❷
def init_network_params(sizes, key=random.PRNGKey(0), scale=1e-2):
"""Initialize all layers for a fully-connected neural network
with given sizes"""
# Инициализация всех слоев для полностью связанной нейронной сети
# с заданными размерами.
def random_layer_params(m, n, key, scale=1e-2):
"""A helper function to randomly initialize
weights and biases of a dense layer"""
# Вспомогательная функция для случайной инициализации
# весов и отклонений плотного слоя.
w_key, b_key = random.split(key)
return (scale * random.normal(w_key, (n, m)),
scale * random.normal(b_key, (n,)))
keys = random.split(key, len(sizes))
❸
399
Классификация изображений MNIST
return [random_layer_params(m, n, k, scale)
for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
def predict(params, image):
"""Function for per-example predictions."""
# Функция для прогноза по одному образцу.
activations = image
for w, b in params[:-1]:
outputs = jnp.dot(w, activations) + b
activations = swish(outputs)
final_w, final_b = params[-1]
logits = jnp.dot(final_w, activations) + final_b
return logits
❶
❷
❸
❹
❺
❻
❼
❽
❹
❺
❻
❼
❼
❽
❽
Список размеров слоев.
Параметр для масштабирования случайных значений.
Генерация случайных значений для параметров слоя w и b.
Генерация случайных значений для всех слоев.
Инициализация активаций с пикселами входного изображения.
Циклы от первого до предпоследнего слоя.
Успешное обновление активаций выходными данными каждого слоя.
Для последнего слоя функция активации не применяется.
Код инициализации параметров и прямого прохода по нейрон
ной сети достаточно прост и понятен, но является низкоуровневым,
поэтому слегка многословен. Здесь нет единой локации, где описы
вается структура нейронной сети. Некоторые описания нейросети
содержатся в списках Python с размерами слоев, другие части вклю
чены в функцию predict(), поэтому приходится вручную анализи
ровать код, чтобы понять структуру нейросети, если только она не
документирована в комментариях. Но даже если структура докумен
тирована (это правильная практическая методика), вы вынуждены
сопровождать ее отдельно от кода и всегда помнить о необходимо
сти ее обновления при внесении изменений в код. При этом остается
потенциальная опасность того, что описание структуры нейросети
и код в какой-то момент времени перестанут соответствовать друг
другу.
Flax предоставляет пользователю абстракцию Module, знакомую
разработчикам, применяющим PyTorch или TensorFlow. Абстракция
сохраняет состояние внутри, и вы пишете код для своих нейросетей
в объектно ориентированном стиле с сохранением состояния. Но для
использования трансформаций JAX Module требует предоставления
чистых функций. Flax создает чистые функции из абстракций Mod
ule, поэтому они функциональны во внешней среде. Эти функции не
сохраняют состояние: они принимают (потребляют) и возвращают
400
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
его. Это тот же самый паттерн, который мы наблюдали в генерато
рах случайных чисел JAX (в главе 9), отличающийся от подхода, при
нятого в аналогичных механизмах NumPy, тем, что не имеет внут
реннего состояния.
Flax предоставляет простой способ описания структуры нейрон
ной сети в стиле самодокументирования. На рис. 11.1 показана на
глядная схема этого процесса.
Переменные
Module
(параметры
и состояние)
Определение
класса нейросети
MLP(nn.Module)
Создание
экземпляра
нейросети
model=MLP()
Инициализация
нейросети model.init()
Случайные
входные данные
PRNGKey
Выходные
данные
(необязательное)
Обновленное
состояние
Применение нейросети
model.apply(params)
Входные
данные
PRNGKey
Рис. 11.1 Схема процесса инициализации и применения нейронной сети в Flax
Чтобы написать код нейронной сети с применением Flax, сначала
необходимо описать нейросеть как последовательность слоев спо
собом, очень похожим на Keras или torch.nn. Этому шагу не соот
ветствует ни одна процедура в примере с чистым кодом JAX. Для
создания описания подкласс класса flax.linen.Module с использова
нием аннотации @nn.compact определяет прямое вычисление внутри
метода __call__, которым мы скоро воспользуемся.
Аннотация @nn.compact представляет собой простой и компакт
ный способ определения нейронной сети. Она позволяет писать ло
гику нейросети прямо внутри одного метода «прямого прохода» (for
ward-pass). Более эффективным способом определения нейронной
сети является использование отдельного метода setup(). Он может
быть востребован для более сложных конфигураций с несколькими
методами прямого прохода (например, если имеется одна функция
для кодирующей части нейросети и другая для декодирующей ча
сти), но в нашем простом примере нет необходимости в примене
нии этого метода. Более подробную информацию о методе setup()
можно получить здесь: https://flax-linen.readthedocs.io/en/latest/
guides/flax_fundamentals/setup_or_nncompact.html. Далее создается
экземпляр объекта нового определенного класса.
401
Классификация изображений MNIST
После этого генерируются фиктивные входные данные для ини
циализации переменных Module, включающие параметры Module
(веса и отклонения) и любые переменные состояния (например,
статистические характеристики для пакетной нормализации сло
ев). Переменные (или параметры) модели инициализируются мето
дом init() созданного экземпляра Module. Для вызова model.init()
требуется ключ PRNGKey и фиктивные входные данные. Это равно
значно вызову init_network_params() в примере чистого кода JAX,
за исключением того, что в текущем примере нет необходимости
в передаче фиктивных данных для вывода форм тензоров, посколь
ку они были определены явно. Flax обеспечивает логический вывод
формы, поэтому от вас требуется только объявление количества ней
ронов, а не размеров входных данных, а затем Flax автоматически
определяет формы матриц весов. Если по каким-либо причинам
в дополнение к инициализированным параметрам необходим вы
вод результатов прямого прохода по фиктивным данным, то можно
воспользоваться функцией init_with_output() вместо init(). Функ
ция принимает те же параметры и возвращает кортеж выходных ре
зультатов и инициализированных параметров.
На завершающем этапе для управления прямым проходом моде
ли с использованием заданного набора параметров вызывается ме
тод apply() созданного экземпляра Module. Этот метод принимает
инициализированные переменные и входные данные. Последний
этап полностью замещает функцию predict(), используемую ра
нее, и вы можете просто заменить все вызовы predict() на вызовы
model.apply(). Если модель имеет некоторое внутреннее состояние,
оно также будет обновляться во время этого вызова. Пример MNIST
не имеет собственного состояния, но мы рассмотрим реализацию
более сложного примера с собственным состоянием в разделе 11.2.
ПРИМЕЧАНИЕ Параметры никогда не хранятся внутри мо
дели. Функции init() и apply() возвращают состояние, а не
поддерживают его.
Полученный в итоге код для определения и инициализации моде
ли показан в листинге 11.2.
Листинг 11.2 Определение и применение MLP в библиотеке Flax
from jax import random
from flax import linen as nn
class MLP(nn.Module):
"""A simple MLP model."""
# Простая модель MLP.
@nn.compact
❶
❶
❷
❸
402
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
def __call__(self, x):
x = nn.Dense(features=512)(x)
x = nn.activation.swish(x)
x = nn.Dense(features=10)(x)
return x
❹
❹
❹
❹
model = MLP()
key1, key2 = random.split(random.PRNGKey(0))
random_flattened_image = random.normal(key1, (28*28*1,))
params = model.init(key2, random_flattened_image)
jax.tree_util.tree_map(lambda x: x.shape, params)
>>> dict({
>>>
params: {
>>>
Dense_0: {
>>>
bias: (512,),
>>>
kernel: (784, 512),
>>>
},
>>>
Dense_1: {
>>>
bias: (10,),
>>>
kernel: (512, 10),
>>>
},
>>>
},
>>> })
model.apply(params, random_flattened_image)
>>> Array([ 0.43500084, 0.20896661, -0.9779186 ,
>>>
0.29811084, -0.24472891, -0.30921277,
>>>
dtype=float32)
❶
❷
❸
❹
❺
❻
❼
❽
❾
❺
❻
❼
❽
❾
0.3447541 , 0.05475511,
0.2945502 , -0.34386003],
Импорт необходимых компонентов.
Класс для многослойного перцептрона (MLP) – подкласс класса flax.linen.Module.
Использование режима объявления (аннотации) nn.compact.
Определение нейросети непосредственно внутри единственного метода прямого прохода.
Создание экземпляра объекта для MLP.
Подготовка фиктивных входных данных (также можно использовать массив нулей).
Инициализация параметров модели.
Проверка выходных форм для параметров модели.
Применение модели к входным данным.
В коде тренировки MLP имеются два места, в которых вызовы predict() должны обновляться с помощью model.apply(): в функции
потерь и в функции точности.
В рассматриваемом здесь примере нет необходимости в измене
нии вычислений градиентов. Функция grad() просто работает с об
новленной функцией потерь, как ранее. Нужно всего лишь немного
изменить функцию update() для применения градиентов к парамет
рам модели. Поскольку параметры модели теперь структурированы
по-другому, с использованием словаря dict, проще изменять их
403
Классификация изображений MNIST
с помощью функции jax.tree_util.tree_map(), с которой мы позна
комились в предыдущей главе.
Это все, что нужно сделать. Код тренировки с функциями loss()
и update() показан в листинге 11.3.
Листинг 11.3 Выполнение шага обновления градиентов для MLP
в Flax
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits = model.apply(params, images)
log_preds = logits - jax.nn.logsumexp(logits)
return -jnp.mean(targets*log_preds)
@jax.jit
def update(params, x, y, epoch_number):
loss_value, grads = jax.value_and_grad(loss)(params, x, y)
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return jax.tree_util.tree_map(
lambda p, g: p - lr * g, params, grads), loss_value
❶
❷
❶ Вызов model.apply() вместо predict().
❷ Использование функции tree_map(), которая больше подходит для pytree.
Это неполный пример, в котором специально выделены толь
ко измененные части. Полноценный работающий пример Chap
ter_11.1_MNIST_MLP_Flax_Simple размещен в репозитории GitHub
этой книги.
Flax соблюдает функциональные соглашения JAX по хранению
данных в деревьях pytree. Исследователям часто необходимо взаи
модействовать вручную с такими данными, поэтому Flax использует
вложенные словари с осмысленными ключами по умолчанию и пре
доставляет несколько утилит для их прямой обработки (например,
для последовательных обходов).
ПРИМЕЧАНИЕ Если вы работаете с версией Flax, более ран
ней, чем 0.7.1, то, возможно, заметите, что Flax использует
структуру данных FrozenDict для хранения параметров мо
дели (https://flax.readthedocs.io/en/latest/api_reference/flax.
core.frozen_dict.html#flax.core.frozen_dict.FrozenDict). Это не
изменяемый вариант словаря Python dict, помогающий ра
ботать с учетом функциональной сущности JAX посредством
запрещения любых изменений внутреннего словаря dict
и предупреждения пользователя об этом. Дополнительное
преимущество структуры данных FrozenDict заключается
в том, что это ускоренная версия неизменяемого словаря Py
thon, кеширующая его JAX-уплощенную форму для уменьше
404
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
ния издержек на вызов JIT-трансформированной функции.
Но после длительных обсуждений (https://github.com/google/
fax/issues/1223) группа разработчиков переключилась на ис
пользование обычных словарей Python при вызове методов
Module init(), init_with_output() и apply(). Внутри Flax про
должает использовать FrozenDict для защиты переменных
типа dict от случайного изменения.
Преимущества такого кода заключаются в двух аспектах. Во-пер
вых, определение нейронной сети представлено в понятной фор
ме, и труднее сделать ошибку, выполняя все операции с тензорами.
Структура данных, содержащая параметры модели, также является
самодокументированной. Вы сразу видите, какие части словаря dict
отвечают за соответствующие слои. Функция tabulate() возвращает
описание модели в форме строк. Она имеет ту же сигнатуру и внут
ренние вызовы, что и метод Module.init(), но вместо переменных
возвращает строки с кратким описанием Module в таблице, которую
можно без затруднений вывести на экран или на печать для отладки:
print(model.tabulate(key2, random_flattened_image))
❶
MLP Summary
┌─────────┬────────┬──────────────┬──────────────┬──────────────────────────┐
│ path
│ module │ inputs
│ outputs
│ params
│
├─────────┼────────┼──────────────┼──────────────┼──────────────────────────┤
│
│ MLP
│ float32[784] │ float32[10] │
│
├─────────┼────────┼──────────────┼──────────────┼──────────────────────────┤
│ Dense_0 │ Dense │ float32[784] │ float32[512] │ bias: float32[512]
│
│
│
│
│
│ kernel: float32[784,512] │
│
│
│
│
│
│
│
│
│
│
│ 401,920 (1.6 MB)
│
├─────────┼────────┼──────────────┼──────────────┼──────────────────────────┤
│ Dense_1 │ Dense │ float32[512] │ float32[10] │ bias: float32[10]
│
│
│
│
│
│ kernel: float32[512,10] │
│
│
│
│
│
│
│
│
│
│
│ 5,130 (20.5 KB)
│
├─────────┼────────┼──────────────┼──────────────┼──────────────────────────┤
│
│
│
│
Total │ 407,050 (1.6 MB)
│
└─────────┴────────┴──────────────┴──────────────┴──────────────────────────┘
Total Parameters: 407,050 (1.6 MB)
❶ Генерация строки с описанием модели.
Кроме того, при наличии такой структуры данных гораздо проще
выполнять интроспекцию, перестройку нейросети, если требуются
какие-либо изменения в уже натренированной нейросети, или пи
сать процедуру преобразования при необходимости импорта или
экспорта нейросети. Например, пользователи Flax уже применяли
Классификация изображений MNIST
405
эту библиотеку для отображения контрольных точек TensorFlow
и PyTorch в Flax.
Во-вторых, все функции трансформации JAX работают превос
ходно, и нет необходимости в существенном изменении кода при
выполнении тех же трансформаций в Flax. Более того, не требуется
даже применение jax.vmap(). В Flax пишете модели как код «для од
ного экземпляра», а фреймворк автоматически исполняет процеду
ру пакетирования.
То, что мы сделали, было не самым типичным вариантом исполь
зования Flax. Существует обобщенный паттерн применения Flax,
описывающий, как структурировать полное состояние тренировки,
включая параметры модели, состояние оптимизатора, количество
итераций и т. п. Но прежде чем заняться этим, необходимо ближе
познакомиться с библиотекой трансформации градиентов Optax.
11.1.2 Библиотека трансформации градиентов Optax
Изначально в Flax существовал собственный API flax.optim для оп
тимизации. Он содержал некоторые оптимизаторы и классы данных
(dataclass), полезные для процесса тренировки модели. Но список
оптимизаторов был далеко не полным, а паттерн, применяемый для
тренировки, являлся относительно сложным и весьма многослов
ным. Группа сопровождения предложила (https://github.com/google/
flax/blob/main/docs/flip/1009-optimizer-api.md) перейти на специа
лизированную библиотеку Optax, уже разработанную компанией
DeepMind.
Optax (https://github.com/google-deepmind/optax) предлагает реа
лизацию широкого спектра оптимизаторов и предоставляет фрейм
ворк для компоновки новых оптимизаторов из многократно ис
пользуемых трансформаций градиентов.
На первый взгляд, Optax предлагает обширный список предвари
тельно определенных современных оптимизаторов, которые можно
использовать прямо «из коробки». Но в действительности это гораз
до больше, чем просто набор оптимизаторов. Точно так же, как сам
по себе JAX представляет собой нечто большее, нежели простая биб
лиотека обработки многомерных массивов, – это библиотека ком
понуемых функций трансформации, так и Optax в большей степени
является библиотекой трансформации градиентов. Optax спроекти
рован для обеспечения исследований посредством предоставления
конструктивных блоков, которые можно с легкостью комбинировать
в различных требуемых сочетаниях.
Первоначальный прототип Optax был доступен в эксперимен
тальном каталоге JAX как jax.experimental.optix, но в итоге из экс
периментального компонента превратился в отдельную библиотеку
с открытым исходным кодом и с новым именем optax. Мы будем ис
406
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
пользовать Optax только для тщательно протестированных и эффек
тивных реализаций оптимизаторов, но следует помнить о том, что
это далеко не все функциональные возможности Optax.
Использование оптимизаторов Optax
Обобщенная схема высокого уровня использования оптимизатора
Optax показана на рис. 11.2.
Состояние
оптимизатора
Создание объекта
оптимизатора
optimizer =
optax.adam(lr)
Инициализация
состояния
оптимизатора
opt_state =
optimizer.
init(params)
Новое
состояние
оптимизатора
Обновления
Получение обновлений
параметров updates,
opt_state = optimizer.
update(grads, opt_state)
Обновленные
параметры
модели
Применение
обновлений параметров:
params = optax.apply_
updates(params, updates)
Градиенты
Параметры
модели
Получение градиентов
grads = jax.grad(loss)
(params, x, y)
Входные
данные
Выходные
данные
Рис. 11.2 Процесс использования оптимизатора Optax (показан только один шаг
обновления градиента)
Схема может показаться сложной, но если привести описание
в виде последовательности шагов, то процесс становится понятным:
1 создание
объекта оптимизатора, например оптимизатора
Adam: optimizer = optax.adam(learning_rate);
2 инициализация состояния оптимизатора (например, векто
ра импульса) с использованием функции init(), вызываемой
с параметрами модели: opt_state = optimizer.init(params);
3 в цикле обновления модели имеется функция потерь, которую
можно продифференцировать с помощью JAX, чтобы получить
градиенты: grads = jax.grad(loss_function)(params, x, y);
4 затем выполняется преобразование градиентов с помощью
вызова optimizer.update(), чтобы получить обновления пара
Классификация изображений MNIST
407
метров. Эта функция также принимает и обновляет состояние
оптимизатора: updates, opt_state = optimizer.update(grads,
opt_state);
5 обновления параметров должны применяться к текущим па
раметрам для получения их новых значений. Для этого мож
но воспользоваться удобной функцией optax.apply_updates():
params = optax.apply_updates(params, updates).
Как можно видеть, Optax использует паттерн, аналогичный Flax,
подразумевающий наличие одной функции для инициализации
(в обоих случаях init()) и другой функции для применения (apply()
в Flax, update() в Optax).
С технической точки зрения каждый оптимизатор реализует ин
терфейс GradientTransformation. Функция init() инициализирует
набор статистических характеристик (или состояние оптимизато
ра), а функция update() преобразовывает потенциально пригодный
градиент с учетом некоторых статистических характеристик (состо
яния) и (необязательно) текущего значения параметров.
Другие компоненты Optax
Трансформация градиентов возвращает не обновленные параметры
модели, а обработанные градиенты. Это создает возможность объ
единения произвольных трансформаций в специализированный
оптимизатор и комбинирования трансформаций для различных
градиентов, работающих с совместно используемым набором пере
менных. Одним из обобщенных приложений является создание це
почки трансформаций градиентов, подразумевающей вычисление
градиентов некоторым оптимизатором, масштабирование их в со
ответствии с некоторым планом скорости обучения, а затем усече
ние для соответствия заданному диапазону.
Optax предоставляет классы и обертки для компоновки транс
формаций градиентов: chain() для применения списка последова
тельных (цепочечных) обновляющих трансформаций и multi_transform() для разделения на части параметров модели и применения
различных трансформаций к каждому подмножеству.
Также существуют обертки, принимающие GradientTransformation как входные данные и возвращающие новую GradientTransformation, изменяющую поведение внутренней трансформации
конкретным заданным способом. Например, одна обертка уплоща
ет градиенты в один вектор перед применением трансформации
градиентов, а затем восстанавливает форму результата. Это можно
использовать для уменьшения издержек при выполнении большо
го объема вычислений с многочисленными небольшими перемен
ными за счет немного увеличенного объема используемой памяти.
Другая обертка делает оптимизатор устойчивым к появлению NaN
408
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
(Not a Number – не чисел) и Inf (бесконечностей). Есть обертка, по
могающая маскировать некоторые параметры, отменяя их обнов
ление, например пропуск уменьшения весов для масштабирования
BatchNorm и всех параметров отклонения.
Планировщики используются для создания в оптимизаторе ком
понентов, зависимых от времени. Общим примером такого подхо
да является алгоритм имитации отжига для скорости обучения или
некоторого другого гиперпараметра. Кроме того, существует груп
па стандартных функций потерь, используемых в глубоком обуче
нии: 12_loss, softmax_cross_entropy, cosine_distance, kl_divergence
и ctc_loss.
Мы рассмотрели самые важные элементы библиотеки Optax, хотя
в ней имеются и другие полезные функциональные средства, опи
санные в официальной документации (https://optax.readthedocs.io/
en/latest/index.html). В следующем подразделе я покажу вам, как ис
пользовать Optax в сочетании с Flax.
11.1.3 Тренировка нейронной сети с применением Flax
Здесь мы включим оптимизатор Optax в процедуру тренировки, до
бавим вычисление метрик и воспользуемся специализированной
структурой данных для хранения состояния тренировки.
Существует специальный класс flax.training.train_state.TrainState (https://flax.readthedocs.io/en/latest/api_reference/flax.training.
html#train-state) для представления простого состояния тренировки
в общем варианте с одним оптимизатором Optax. Если необходимо
что-то добавить в состояние тренировки, скажем метрики качества,
потребуется создать соответствующий подкласс. Начнем с самого прос
того варианта с добавлением единственного оптимизатора Optax.
Добавление TrainState
В функцию TrainState.create() передаются следующие параметры:
apply_fn – функция для применения нейронной сети, обычно
передается model.apply;
params – параметры модели, используемые функцией apply_fn
и обновляемые оптимизатором (см. следующий параметр tx);
tx – оптимизатор, реализующий интерфейс Optax GradientTransformation.
Также существуют параметры для количества шагов тренировки
и состояния оптимизатора, но для наших целей вполне достаточно
трех вышеперечисленных параметров. TrainState упрощает процесс
применения оптимизатора Optax. Из пяти шагов, описанных в пре
дыдущем подразделе, мы будем наблюдать шаги 1, 3 и объединен
ные шаги 4 и 5.
409
Классификация изображений MNIST
Мы создаем состояние тренировки с помощью простого опти
мизатора: метода стохастического градиентного спуска (SGD) с им
пульсным параметром возмущения. Таким образом, мы реализуем
шаг 1 из процедуры использования Optax (создаем объект оптими
затора) и шаг 2 (инициализация состояния оптимизатора), но вто
рой шаг скрыт внутри вызова метода TrainState.create().
Листинг 11.4
Создание состояния тренировки
from flax.training import train_state
import optax
state = train_state.TrainState.create(
apply_fn=model.apply,
params=params,
tx=optax.sgd(learning_rate=1.0, momentum=0.9))
❶
❷
❸
❹
Импорт класса TrainState.
Импорт библиотеки оптимизации Optax.
Создание состояния тренировки.
Использование оптимизатора SGD из библиотеки optax.
❶
❷
❸
❹
Теперь у нас есть состояние тренировки и можно организовать
цикл тренировки.
Внутри цикла тренировки используется почти та же функция по
терь с единственным изменением: вместо прямого вызова model.
apply() вызывается применяемая функция из объекта TrainState,
т. е. train_state.apply_fn().
В функции update() мы получаем градиенты тем же способом,
что и ранее (шаг 3 вышеописанной процедуры использования оп
тимизатора Optax), но дальнейшее применение градиентов к весам
модели упрощается и выполняется с помощью вызова train_state.
apply_gradients(grads). Внутри этой функции вызывается несколь
ко функций для обновления состояния оптимизатора, формирова
ния обновлений для параметров модели (шаг 4) и применения этих
обновлений (шаг 5). Полученный в результате цикл тренировки вы
глядит приблизительно так, как показано в листинге 11.5. (Полный
листинг см. в репозитории GitHub этой книги.)
Листинг 11.5
Обновленный цикл тренировки
@jax.jit
def update(train_state, x, y):
"""A single training step"""
# Один шаг тренировки.
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits = train_state.apply_fn(params, images)
❶
410
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
log_preds = logits - jax.nn.logsumexp(logits)
return -jnp.mean(targets*log_preds)
loss_value, grads = jax.value_and_grad(loss)(
train_state.params, x, y)
train_state = train_state.apply_gradients(grads=grads)
return train_state, loss_value
❷
❸
…
for epoch in range(num_epochs):
start_time = time.time()
losses = []
for x, y in train_data:
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = jax.nn.one_hot(y, NUM_LABELS)
state, loss_value = update(state, x, y)
losses.append(loss_value)
epoch_time = time.time() - start_time
❶
❷
❸
❹
❹
Применение модели к входным изображениям.
Вычисление градиентов остается неизменным.
Обновления параметров модели и состояния оптимизатора.
Обновления состояния тренировки.
Код стал еще более высокоуровневым и скрывающим некоторые
подробности низкого уровня, касающиеся обновления параметров.
Добавление вычисления метрик
с использованием библиотеки CLU
Для увеличения степени вовлеченности Flax в процесс тренировки
можно добавить вычисление метрик. Существует общий паттерн
для добавления метрик в состояние тренировки.
Для вычисления метрик воспользуемся еще одной библиотекой
из экосистемы JAX – Common Loop Utils (CLU) (https://github.com/
google/CommonLoopUtils). Она содержит общую функциональность
для написания циклов тренировки машинного обучения, позволя
ющую сделать их более короткими и удобными для чтения без сни
жения степени гибкости, требуемой для исследований. CLU предна
значена для удобной работы с JAX и Flax.
Помимо всего прочего, библиотека CLU определяет функциональ
ный интерфейс вычисления метрик Metric, основанный на метри
ках, накапливающих промежуточные значения, а затем использую
щих эти промежуточные значения для вычисления окончательного
значения метрики. Такой подход работает для агрегируемых мет
рик. Ниже приведено краткое описание этой методики.
1 Для каждого пакета вычисляются локальные пакетные метрики
по выходным данным модели. «Выходные данные модели» – это
словарь значений с неповторяющимися ключами с конкретно
411
Классификация изображений MNIST
определенным смыслом (например, «loss» (потери), «logits»
(логиты) и «labels» (метки)). Каждая метрика зависит как мини
мум от одного такого комплекта выходных данных модели в со
ответствии с именем. Интерфейс Metrics предоставляет функ
цию from_output() для определения имени выходных данных
модели, которое будет использоваться для вычисления метрик.
2 Локальные, или промежуточные, метрики из различных паке
тов объединяются с помощью функции merge().
3 Итоговые метрики вычисляются по объединенным промежу
точным значениям с использованием функции compute().
Библиотека CLU также предоставляет интерфейс с именем metrics.Collection для вычисления коллекции (набора) метрик по всем
выходным данным модели одновременно. Интерфейс использует
функцию single_from_model_output() для конфигурации с нерас
пределенными вычислениями и функцию gather_from_model_output() для конфигурации с распределенными вычислениями.
Чтобы добавить метрики в состояние тренировки, мы сначала
объявляем класс данных dataclass для хранения метрик, потерь
и точности. Затем создается подкласс TrainState для включения
метрик. После этого создается экземпляр нового определенного
класса TrainState и инициализируются все связанные с ним поля,
включая метрики. Код показан в листинге 11.6.
Листинг 11.6
Добавление метрик в состояние тренировки
from flax.training import train_state
from clu import metrics
import flax
import optax
@flax.struct.dataclass
class Metrics(metrics.Collection):
accuracy: metrics.Accuracy
loss: metrics.Average.from_output('loss')
class TrainState(train_state.TrainState):
metrics: Metrics
state = TrainState.create(
apply_fn=model.apply,
params=params,
tx=optax.sgd(learning_rate=0.01, momentum=0.9),
metrics=Metrics.empty())
❶
❷
❸
❹
❺
❶
❶
❷
❸
❹
❺
Определение собственного специализированного класса для метрик.
Добавление метрики точности.
Добавление метрики потерь.
Добавление метрик в объект TrainState.
Инициализация метрик пустой коллекцией.
412
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
После этого мы обязательно должны изменить цикл тренировки
для включения вычисления метрик. Мы воспользуемся функцией
для вычисления всех метрик и будем вызывать ее для каждого тре
нировочного пакета. Внутри этой функции применяется функция из
библиотеки Optax для вычисления потерь перекрестной энтропии
(что позволит сократить код на несколько строк). Для метрик точно
сти используется существующая метрика из библиотеки CLU, и этой
библиотеке известно, как накапливать (агрегировать) состояние
и вычислять итоговую метрику по промежуточным значениям. Если
необходима другая метрика, скажем мера точности (верности; preci
sion) или отклик модели (доля от истинно положительных результа
тов; recall), то, возможно, потребуется ее независимая реализация,
что не составляет особого труда.
Для вычисления тестового набора метрик мы клонируем TrainState с пустыми метриками и вычисляем все метрики тем же спосо
бом, который применяется во время тренировки.
Листинг 11.7
Вычисление метрик внутри цикла тренировки
@jax.jit
def compute_metrics(state, x, y):
logits = state.apply_fn(state.params, x)
loss = optax.softmax_cross_entropy_with_integer_labels(
logits=logits, labels=y).mean()
metric_updates = state.metrics.single_from_model_output(
logits=logits, labels=y, loss=loss)
metrics = state.metrics.merge(metric_updates)
state = state.replace(metrics=metrics)
return state
❶
❷
❸
❹
❺
❻
for epoch in range(num_epochs):
start_time = time.time()
for x, y in train_data:
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = y.astype(jnp.int32) # По умолчанию это int64, а clu ожидает int32.
state, loss_value = update(state, x, y)
❼
state = compute_metrics(state, x, y)
epoch_time = time.time() - start_time
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
for metric,value in state.metrics.compute().items():
print(f"Training set {metric} {value}")
state = state.replace(metrics=state.metrics.empty())
❽
test_state = state
for x, y in test_data:
❿
❾
413
Классификация изображений с использованием ResNet
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = y.astype(jnp.int32)
test_state = compute_metrics(test_state, x, y)
for metric,value in test_state.metrics.compute().items():
print(f"Test set {metric} {value}")
⓫
⓫
❶ Эта функция выполняет вычисления всех метрик и обновляет TrainState новыми
значениями метрик.
❷ Получение выходных данных модели.
❸ Использование функции потерь softmax из библиотеки Optax.
❹ Предоставление выходного словаря модели для вычисления всех промежуточ❺
❻
❼
❽
❾
❿
⓫
ных метрик в коллекции.
Объединение промежуточных метрик.
Обновление TrainState с использованием объединенных метрик.
Обновление метрик во время выполнения цикла тренировки.
Вычисление итоговых метрик.
Переустановка метрик после каждой эпохи.
Клонирование состояния с пустыми метриками для вычисления.
Вычисление метрик по тестовому набору после каждой эпохи тренировки.
Мы натренировали нейронную сеть с использованием Flax, но мо
дель продолжает оставаться чрезвычайно простым многослойным
перцептроном, не соответствующим самым современным требова
ниям. В следующем разделе мы реализуем более продвинутую со
временную остаточную нейронную сеть для классификации изобра
жений.
11.2 Классификация изображений
с использованием ResNet
Вероятно, вы понимаете, что использование многослойных пер
цептронов для задач обработки изображений обычно не является
оптимальным вариантом, поскольку существуют более специали
зированные и эффективные решения, такие как конволюционные
(сверточные) нейронные сети (convolutional neural networks – CNN),
развивающиеся в течение уже многих лет. Обычно пользователи вы
бирают для применения остаточные нейросети (residual neural net
works – ResNets).
Здесь мы снова воспользуемся набором данных Dogs vs. Cats из
главы 9 и реализуем относительно простую остаточную нейросеть
(ResNet), которая, возможно, является правильным вариантом вы
бора для подобной задачи.
414
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
Современное компьютерное зрение
В реальной практике обычным решением задачи классификации изображений чаще всего является использование конволюционных (сверточных) нейросетей (CNN). CNN – это тип нейронной сети, хорошо
подходящий для работы с изображениями. Такая нейросеть обладает
замечательными свойствами, например эквивариантностью параллельных перемещений (translation equivariance) и способностью изучать
(понимать) локальные характеристики.
Нейросети CNN уже давно считаются наилучшим типом нейронных сетей для работы с изображениями, и многочисленные их усовершенствования раздвинули границы еще шире. Остаточные нейросети (ResNets)
(https://arxiv.org/abs/1512.03385) представляли собой один из вариантов CNN. Поэтому именно CNN являлись естественным выбором для
решения задач компьютерного зрения. Хотя за последние несколько лет
ситуация изменилась.
Трансформеры (https://arxiv.org/abs/1706.03762) появились в 2017 г.
и впервые были применены для решения задач обработки естественного языка (NLP), таких как машинный перевод. В 2018 г. была разработана модель BERT (https://arxiv.org/abs/1810.04805), применимая почти
ко всем задачам NLP.
В 2020 г. нейросеть Vision Transformer (ViT) (https://arxiv.org/abs/
2010.11929) продемонстрировала превосходную производительность
при обработке изображений. Эта работа положила начало широкому
потоку приложений-трансформеров для изображений и видео.
В 2021 г. появилась модель MLP-Mixer (https://arxiv.org/abs/2105.01601)
вместе с несколькими аналогичными работами, которые окончательно
определили способ применения старых добрых многослойных перцептронов (MLP) в задачах компьютерного зрения и показали результаты,
сравнимые с трансформерами.
Архитектура CNN не оставалась неизменной в течение всего периода
их существования и развития, в нее вносились многочисленные усовершенствования. ConvNeXt (https://arxiv.org/abs/2201.03545) и EffcientNetV2 (https://arxiv.org/abs/2104.00298) заняли место среди лидеров
выполнения классификации изображений в 2022 г.
Достаточно интересный факт: ViT и MLP-Mixer были созданы с использованием JAX.
Полный код процедур загрузки набора данных и его предвари
тельной обработки размещен в репозитории GitHub этой книги. Эта
часть кода не так важна в рассматриваемом здесь примере, поэтому
сразу переходим к созданию ResNet.
Классификация изображений с использованием ResNet
415
11.2.1 Управление состоянием в Flax
Ранее мы использовали слои простой нейронной сети без какоголибо внутреннего состояния. Но некоторые широко применяемые
типы слоев имеют внутреннее состояние, и самым известным при
мером является слой нормализации пакетов (BatchNorm).
BatchNorm нормализует активации слоя. Это помогает быстрее
тренировать модели и достигать более высокой производительно
сти. Существуют многочисленные типы нормализации: нормализа
ция слоя, нормализация групп, нормализация экземпляров и т. д.
BatchNorm
Методика BatchNorm была представлена в 2015 г. Сергеем Иоффе (Sergey Ioffe) и Кристианом Сегеди (Christian Szegedy) в документе «Batch
Normalization: Accelerating Deep Network Training by Reducing Internal
Covariate Shift» (https://arxiv.org/abs/1502.03167).
Наглядное визуальное описание работы BatchNorm можно найти здесь:
https://www.youtube.com/watch?v=yXOMHOpbon8.
BatchNorm – это слой с различным поведением во время тренировки
и логического вывода результата. Ниже приведено описание его работы.
Во время тренировки BatchNorm вычисляет статистические характерис
тики пакета, а именно среднее значение и (среднеквадратичное) отклонение по пакету для каждой активации. Затем он нормализует каждую
активацию посредством вычитания среднего значения и деления на
стандартное отклонение (квадратный корень из среднеквадратичного
отклонения). Потом применяется обученная линейная трансформация,
выполняющая масштабирование (по другому стандартному отклонению) и сдвиг (по другому среднему значению). Как и многие другие веса
слоев нейронной сети, эти два коэффициента масштабирования и сдвига обучаются методом обратного распространения. Весьма важно иметь
такую обученную трансформацию в дополнение к части нормализации,
чтобы BatchNorm мог обучаться тождественному преобразованию при
необходимости.
BatchNorm также сохраняет экспоненциальную скользящую среднюю
(exponential moving average) среднего значения и среднеквадратичного отклонения, так как это хороший представитель-посредник для
среднего значения и среднеквадратичного отклонения всех данных
в целом и его гораздо проще вычислить. Эти скользящие средние значения сохраняются внутри состояния слоя, но отдельно от обучаемых
параметров модели.
Во время логического вывода результата нейронная сеть может работать с использованием одной точки входных данных без какого-либо
416
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
пакета, поэтому статистические характеристики пакета становятся недоступными. BatchNorm использует сохраненные значения скользящей
средней среднего значения и среднеквадратичного отклонения для
нормализации и фиксирует обученные веса для трансформаций масштабирования и сдвига.
Flax предлагает собственное руководство, описывающее, как применять
слой BatchNorm, предоставляемый этой библиотекой: https://flax-lin
en.readthedocs.io/en/latest/guides/training_techniques/batch_norm.
html.
Проблема, связанная с применением BatchNorm и других слоев
с внутренним состоянием, заключается в том, что при функцио
нальном подходе Flax не должно быть никаких побочных эффек
тов и внутренних состояний, а изменение состояния программы
представляет собой одну из разновидностей побочных эффектов.
Состояние должно передаваться как внешний параметр для устра
нения этой проблемы, поэтому классы и функции стали не имею
щими состояния (https://docs.jax.dev/en/latest/stateful-computations.
html). И снова следует отметить, что это в точности тот же паттерн,
что и применяемый для генераторов случайных чисел JAX (см. гла
ву 9), использующих внешнее состояние как ключ PRNGKey.
BatchNorm – весьма важная часть ResNet с поддержкой конволю
ционных слоев. Flax предоставляет реализации многих широко рас
пространенных слоев нейросетей в модуле flax.linen (https://flaxlinen.readthedocs.io/en/latest/api_reference/flax.linen/layers.html).
Модуль содержит конволюционные слои (https://flax-linen.readthed
ocs.io/en/latest/api_reference/flax.linen/layers.html#flax.linen.Conv)
и BatchNorm (https://flax-linen.readthedocs.io/en/latest/api_reference/
flax.linen/layers.html#flax.linen.BatchNorm).
Словарь с инициализированными переменными для BatchNorm
будет содержать в дополнение к коллекции params отдельную кол
лекцию batch_stats, включающую все текущие статистические ха
рактеристики.
Модуль Flax BatchNorm содержит параметр use_running_average
с установленным значением True, если статистические характерис
тики, хранящиеся в batch_stats, должны использоваться (во время
логического вывода результата), или False, если необходимо вычис
лять статистические характеристики по входным данным и обнов
лять коллекцию batch_stats (во время тренировки). В коде мы будем
использовать логический флаг, чтобы отличать режим тренировки
от режима логического вывода, и этот флаг будет определять пара
метр use_running_average слоев BatchNorm.
Код в листинге 11.8 определяет небольшую модель ResNet с 18 слоя
ми (ResNet18). В блокноте для этой главы вы также найдете код для
Классификация изображений с использованием ResNet
417
других, более крупных нейросетей ResNet с 34, 50, 101, 152 и 200 слоя
ми. Полностью завершенный работающий пример см. в блокноте
Colab. Здесь же выделены только самые важные части, которые ин
тересуют нас прямо сейчас. Приведенный ниже код основан на при
мере ImageNet из репозитория Flax. Примеры Flax предоставляют
великолепную возможность, для того чтобы начать эксперименти
ровать с моделями из реальной практики при изучении Flax.
Листинг 11.8
Определение нейросети ResNet с 18 слоями
from flax import linen as nn
from functools import partial
from typing import Any, Callable, Sequence, Tuple
ModuleDef = Any
class ResNetBlock(nn.Module):
"""ResNet block."""
filters: int
conv: ModuleDef
norm: ModuleDef
act: Callable
strides: Tuple[int, int] = (1, 1)
@nn.compact
def __call__(self, x,):
residual = x
y = self.conv(self.filters, (3, 3), self.strides)(x)
y = self.norm()(y)
y = self.act(y)
y = self.conv(self.filters, (3, 3))(y)
y = self.norm(scale_init=nn.initializers.zeros_init())(y)
❶
❷
❸
❹
❺
❻
if residual.shape != y.shape:
residual = self.conv(self.filters, (1, 1),
self.strides, name='conv_proj')(residual)
residual = self.norm(name='norm_proj')(residual)
return self.act(residual + y)
class ResNet(nn.Module):
"""ResNetV1."""
stage_sizes: Sequence[int]
block_cls: ModuleDef
num_classes: int
num_filters: int = 64
dtype: Any = jnp.float32
act: Callable = nn.relu
conv: ModuleDef = nn.Conv
@nn.compact
def __call__(self, x, train: bool = True):
❼
❽
❾
❿
418
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
conv = partial(self.conv, use_bias=False, dtype=self.dtype)
norm = partial(nn.BatchNorm,
use_running_average=not train,
momentum=0.9,
epsilon=1e-5,
dtype=self.dtype)
x = conv(self.num_filters, (7, 7), (2, 2),
padding=[(3, 3), (3, 3)],
name='conv_init')(x)
x = norm(name='bn_init')(x)
x = nn.relu(x)
x = nn.max_pool(x, (3, 3), strides=(2, 2), padding='SAME')
for i, block_size in enumerate(self.stage_sizes):
for j in range(block_size):
strides = (2, 2) if i > 0 and j == 0 else (1, 1)
x = self.block_cls(self.num_filters * 2 ** i,
strides=strides,
conv=conv,
norm=norm,
act=self.act)(x)
x = jnp.mean(x, axis=(1, 2))
x = nn.Dense(self.num_classes, dtype=self.dtype)(x)
x = jnp.asarray(x, self.dtype)
return x
ResNet18 = partial(ResNet, stage_sizes=[2, 2, 2, 2],
block_cls=ResNetBlock)
⓫
⓬
⓭
⓮
model = ResNet18(num_classes=NUM_LABELS)
❶
❷
❸
❹
❺
❻
❼
❽
❾
❿
⓫
⓬
⓭
⓮
Импорт для применения функции partial.
Импорт для аннотаций типов.
Класс для структурного блока ResNet.
Конволюционный слой.
Слой нормализации.
Активация.
Класс для полной ResNet.
Использование активации linen.relu.
Использование слоя linen.Conv.
Использование параметра для различения режимов тренировки и логического
вывода результата.
Использование слоя linen.BatchNorm.
Определение способа использования скользящих средних в зависимости от режима.
Использование ResNetBlocks в качестве структурных блоков.
Установка параметров для 18-слойной ResNet.
Здесь мы видим гораздо более продвинутую нейронную сеть по
сравнению с предыдущими примерами. Приведенный выше код по
419
Классификация изображений с использованием ResNet
хож на определение модели в TensorFlow или PyTorch. Такую ней
росеть или более глубокую ResNet из блокнота можно использовать
для решения реальных практических задач классификации изобра
жений.
В этом коде нет ничего особенного, за исключением новых слоев,
которые ранее не использовались, и отдельного параметра, позволя
ющего различать режимы тренировки и логического вывода резуль
тата. Состояние модели скрыто внутри слоев BatchNorm, и единствен
ной точкой, связанной с внутренним состоянием, является параметр
use_running_average, значение которого зависит от текущего режи
ма: тренировка или логический вывод результата.
Функция model.init() в основном та же, что и раньше; она про
должает возвращать все переменные модели. Теперь переменные
включают не только параметры модели, но и ее состояние (статис
тические характеристики BatchNorm). Итоговый словарь dict после
инициализации содержит ключ params для тренируемых параметров
модели, как и ранее, а также ключ batch_stats с переменными mean
и var для хранения скользящих средних каждого слоя. Все это можно
представить в наглядной форме, используя метод model.tabulate().
Листинг 11.9 Инициализация ResNet
key1, key2 = random.split(random.PRNGKey(0))
variables = model.init(key2, images)
model_state, params = flax.core.pop(variables, 'params')
print(model.tabulate(key2, images))
❶
❷
❸
>>>
>>>
ResNet Summary
>>>┌─────────┬──────────┬───────────┬────────────┬─────────────┬────────────┐
>>>│path
│ module
│ inputs
│ outputs
│ batch_stats │ params
│
>>>├─────────┼──────────┼───────────┼────────────┼─────────────┼────────────┤
>>>│
│ ResNet
│ float32[3…│ float32[3… │
│
│
>>>├─────────┼──────────┼───────────┼────────────┼─────────────┼────────────┤
>>>│conv_init│ Conv
│ float32[3…│ float32[3… │
│ kernel:
│
>>>│
│
│
│
│
│ float32[7… │
>>>│
│
│
│
│
│
│
>>>│
│
│
│
│
│ 9,408
│
>>>│
│
│
│
│
│ (37.6 KB) │
>>>├─────────┼──────────┼───────────┼────────────┼─────────────┼────────────┤
>>>│bn_init │ BatchNorm│ float32[3…│ float32[3… │ mean:
│ bias:
│
>>>│
│
│
│
│ float32[64] │ float32[6… │
>>>│
│
│
│
│ var:
│ scale:
│
>>>│
│
│
│
│ float32[64] │ float32[6… │
>>>│
│
│
│
│
│
│
>>>│
│
│
│
│ 128 (512 B) │ 128 (512
│
420
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
>>>│
│
│
│
│
│ B)
│
>>>├─────────┼──────────┼───────────┼────────────┼─────────────┼────────────┤
>>> …
❶ Инициализация модели.
❷ Разделение состояния и параметров модели.
❸ Вывод в наглядной (табличной) форме всех переменных модели (состояния и параметров).
Здесь можно видеть, что состояние и параметры модели стано
вятся явными после инициализации. В дальнейшем они будут ис
пользоваться для тренировки и логического вывода результата. Обе
переменные мы включили в объект TrainState, используемый ра
нее. Другие части кода, связанные с функцией model.apply(), опти
мизатором и метриками, не изменились.
Функция model.apply() (теперь вызываемая через train_state.
apply_fn) изменилась по сравнению с примером MLP без состояния.
Она по-прежнему продолжает принимать словарь с переменными
модели, но теперь в эти переменные включено еще и состояние.
Второй аргумент (x) – это входные данные, которые необходимо
обработать, как и ранее. Кроме того, появились еще два новых аргу
мента. Один именованный аргумент передает логический параметр
train для различения режимов тренировки и логического вывода.
Другой аргумент с именем mutable определяет, какую коллекцию
в словаре переменных модели следует считать изменяемой. В дан
ном случае изменяемым является состояние, которое необходимо
обновлять внутри цикла тренировки, поэтому обязательно требует
ся передавать ключи из части переменных, содержащей состояние
модели.
Ранее в примере MLP функция model.apply() возвращала просто
выходные данные модели, или логиты. Теперь с изменяемой частью
необходимо возвращать еще и обновленное состояние модели, по
этому если для параметра mutable установлено значение True, то
функция возвращает кортеж (output, vars), где vars – словарь из
мененных коллекций, в данном случае скользящих средних для
средних значений и отклонений данных. Кроме того, при вычис
лении градиентов необходимо помнить эти вспомогательные дан
ные вместе с состоянием модели, следовательно, трансформация
grad() (или value_and_grad()) должна знать о таких дополнитель
ных возвращаемых значениях. Мы будем использовать параметр
has_aux=True, описанный в главе 4.
Все вышеописанные изменения вносятся в функции update(),
evaluate() и в объект TrainState. Цикл тренировки не изменяется.
Короче говоря, проще посмотреть код, чем читать все эти длинные
описания.
Классификация изображений с использованием ResNet
Листинг 11.10
421
Обновленный цикл тренировки с состоянием модели
class TrainState(train_state.TrainState):
metrics: Metrics
model_state: Any
state = TrainState.create(
apply_fn=model.apply,
params=params,
model_state=model_state,
tx=optax.sgd(learning_rate=0.01, momentum=0.9),
metrics=Metrics.empty())
@jax.jit
def update(train_state, x, y):
"""A single training step"""
# Один шаг тренировки.
def loss(params):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits, new_model_state = train_state.apply_fn(
{'params': params, **train_state.model_state},
x,
mutable=list(model_state.keys()),
train=True)
loss_ce = optax.softmax_cross_entropy_with_integer_labels(
logits=logits, labels=y).mean()
return loss_ce, (logits, new_model_state)
grad_fn = jax.value_and_grad(loss, has_aux=True)
(loss_value, (logits, new_model_state)), grads = \
grad_fn(train_state.params)
train_state = train_state.apply_gradients(grads=grads,
model_state=new_model_state)
train_state = compute_metrics(train_state, loss=loss_value,
logits=logits, labels=y)
return train_state, loss_value
@jax.jit
def evaluate(train_state, x, y):
"""A single eval step"""
# Один шаг вычисления.
logits = train_state.apply_fn(
{'params': train_state.params, **train_state.model_state},
x,
mutable=False,
train=False)
❶
❷
❷
❸
❹
❺
❻
❼
❼
❼
❽
❾
❿
422
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
loss_ce = optax.softmax_cross_entropy_with_integer_labels(
logits=logits, labels=y).mean()
train_state = compute_metrics(
train_state, loss=loss_ce, logits=logits, labels=y)
return train_state
for epoch in range(num_epochs):
…
for x, y in train_data:
state, loss_value = update(state, x, y)
…
for x, y in test_data:
state = evaluate(state, x, y)
…
⓫
⓫
⓫
❶ Добавление состояния модели (нетренируемых параметров) в TrainState.
❷ Сохранение состояния модели и параметров модели по отдельности.
❸ Функция модели apply() возвращает выходные данные и обновленное со
стояние.
❹ Повторное формирование полного набора переменных модели, включающего
параметры модели и ее состояние.
❺ Пометка переменных состояния модели как изменяемых.
❻ Оповещение функции apply() о том, что мы находимся в режиме тренировки.
❼ Возврат вспомогательных параметров (состояния модели) после трансформации
❽
❾
❿
⓫
градиентов.
Обновление состояния модели внутри TrainState.
Во время вычисления состояние модели не обновляется.
Оповещение функции apply() о том, что мы НЕ находимся в режиме тренировки.
Обычный цикл тренировки с разделением частей тренировки и вычисления.
В приведенном выше примере можно видеть, что поскольку мы
начинаем с использованием TrainState, теперь стало проще вносить
изменения без постоянной модификации сигнатур функций после
добавления каждого нового функционального компонента в мо
дель. Все относящееся к состоянию тренировки размещено в этой
структуре данных, доступной внутри функций тренировки и вычис
ления цикла.
Итак, мы реализовали и натренировали остаточную нейросеть
ResNet. Остался завершающий шаг: сохранение натренированной
модели, чтобы использовать ее в дальнейшем.
11.2.2 Сохранение и загрузка модели с использованием Orbax
Для сохранения и загрузки контрольных точек Flax предлагается ис
пользовать еще одну библиотеку из экосистемы JAX – Orbax (https://
github.com/google/orbax). В Flax имеется устаревший API flax.training.checkpoints, но в настоящее время рекомендуется способ сохра
нения модели с применением Orbax.
Классификация изображений с использованием ResNet
423
Используя Orbax, можно сохранять и загружать любые конкретные
деревья pytree JAX и даже специализированные (пользовательские)
классы, производные от flax.struct.dataclass. Поэтому предостав
ляется возможность сохранять не только параметры модели, но так
же почти все сгенерированные данные, включая массивы, словари,
метаданные, конфигурации и т. д.
Orbax включает библиотеку поддержки контрольных точек, ори
ентированную на пользователей JAX, с поддержкой разнообразных
функциональных средств, требуемых различными фреймворками,
в том числе сопровождение контрольных точек, набор типов и мно
жество форматов хранения. Для установки Orbax выполните следу
ющую команду:
pip install orbax-checkpoint
Orbax также включает библиотеку сериализации для пользовате
лей JAX, позволяющую экспортировать модели JAX в формат Tensor
Flow SavedModel. Для установки этого функционального компонента
предназначена следующая команда:
pip install orbax-export
Для сохранения и восстановления параметров модели создает
ся объект контрольной точки (checkpointer) (специализированный
класс в Orbax; мы будем использовать объект контрольной точки
для pytree PyTreeCheckpointer). Для сохранения модели вызывается
метод save() с передачей в него пути и сохраняемого pytree. Для вос
становления модели предназначен метод restore(), принимающий
путь в файловой системе. По умолчанию объект контрольной точки
сохраняет каждый параметр в дереве pytree как отдельный каталог.
Существует необязательный параметр save_args, который рекомен
дуется использовать для повышения производительности, так как
он позволяет объединять массивы небольшого размера в pytree в од
ном большом файле вместо создания многочисленных маленьких
файлов.
Листинг 11.11
Сохранение и восстановление параметров модели
from flax.training import orbax_utils
import orbax.checkpoint
path = 'tmp/orbax/saved_model'
orbax_checkpointer = orbax.checkpoint.PyTreeCheckpointer()
save_args = orbax_utils.save_args_from_target(state.params)
orbax_checkpointer.save(path, state.params, save_args=save_args)
params_restored = orbax_checkpointer.restore(path)
❶
❶
❷
❸
❹
❺
❻
424
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
Импорт требуемых модулей.
Определение имени каталога для сохранения модели.
Инициализация объекта контрольной точки Orbax для деревьев pytree.
Подготовка структуры для параметра save_args, чтобы сохранить каждый параметр в одном файле.
❺ Сохранение параметров модели.
❻ Восстановление параметров модели.
❶
❷
❸
❹
Сохранение и восстановление модели выполняется просто. Также
можно использовать функции управления версиями и журналиро
вания, например сохранение контрольной точки модели после каж
дой эпохи. Далее потребуется обертка orbax.checkpoint.CheckpointManager поверх объекта контрольной точки. Для CheckpointManager
имеются параметры, управляющие интервалом сохранения конт
рольных точек, максимальным количеством сохраняемых конт
рольных точек, определяющие префикс для каталогов и т. д. Цикл
тренировки с использованием менеджера контрольных точек пока
зан в листинге 11.12.
Листинг 11.12 Использование менеджера контрольных точек
для управления созданием контрольных точек
orbax_checkpointer = orbax.checkpoint.PyTreeCheckpointer()
options = orbax.checkpoint.CheckpointManagerOptions(
max_to_keep=4,
save_interval_steps=10,
create=True)
checkpoint_manager = orbax.checkpoint.CheckpointManager(
'tmp/orbax/checkpoints', orbax_checkpointer, options)
# Внутри цикла тренировки.
for epoch in range(100):
# ... здесь выполняется тренировка.
checkpoint_manager.save(epoch, params,
save_kwargs={'save_args':save_args})
❶
❷
❸
❹
❺
❻
❼
Создание объекта контрольной точки Orbax для деревьев pytree.
Создание параметров для менеджера контрольных точек.
Нужно сохранять максимум четыре контрольные точки.
Сохранение контрольной точки на каждом десятом шаге.
Создание каталога верхнего уровня, если он пока еще не существует.
Создание менеджера контрольных точек.
Регулярно повторяющиеся вызовы менеджера контрольных точек.
❶
❷
❸
❹
❺
❻
❼
При таком подходе можно с легкостью сохранять все промежуточ
ные результаты без добавления слишком большого объема логики
в цикл тренировки. Доступны и другие объекты контрольных точек,
Использование экосистемы Hugging Face
425
в том числе асинхронные, в которых операции сохранения выпол
няются в фоновом потоке, так что можно продолжать вычисления
одновременно с сохранением параметров.
Orbax также поддерживает сохранение и загрузку деревьев pytree
с многопроцессными массивами тем же способом, что и однопро
цессные pytree. Процедура сохранения контрольных точек в много
процессном контексте использует тот же API, что и в однопроцесс
ном контексте. Но применение асинхронного объекта контрольных
точек рекомендуется для сохранения больших многопроцессных
массивов.
На текущий момент это все, что нужно знать о Flax. Мы коснулись
всего лишь малой части того, что можно делать с помощью Flax. Тем
не менее мы освоили основы, с которыми можно двигаться дальше.
Flax предоставляет превосходную документацию и набор примеров,
которые можно скопировать и начать экспериментировать с ними.
Кроме того, предлагаются подробные пошаговые инструкции для
тренировки в параллельном режиме, для инспекции и перестройки
модели, для преобразования модели PyTorch и т. п.
В предыдущих разделах мы подробно рассматривали экосисте
му JAX. Мы применили Flax для высокоуровневого моделирования
нейронной сети, Optax для оптимизации модели, CLU для структу
рирования цикла тренировки и Orbax для сохранения и загрузки мо
делей. Все это подчеркивает модульный подход, принятый членами
экосистемы. Теперь мы перейдем к более крупной экосистеме моде
лей машинного обучения – к трансформерам Hugging Face.
11.3 Использование экосистемы Hugging Face
В этом разделе мы рассмотрим экосистему Hugging Face и узнаем,
как использовать предварительно натренированные модели из Hug
ging Face Model Hub для задач обработки естественного языка и ге
нерации изображений. Hugging Face предоставляет великолепные
библиотеки трансформеров и диффузоров, а также огромное сете
вое хранилище с тысячами моделей с открытым исходным кодом.
Наиболее часто используемым фреймворком является PyTorch, хотя
на момент написания этой главы имелось в наличии более 9500 мо
делей JAX (https://huggingface.co/models?library=jax).
В таком объеме интересных моделей с открытым исходным кодом
можно особо выделить множество GPT-подобных больших языко
вых моделей (large language model – LLM), в том числе самые совре
менные семейства Llama и Gemma:
семейство моделей Llama (https://huggingface.co/meta-llama)
компании Meta;
426
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
семейство LLM с открытым исходным кодом Gemma 2 (https://
huggingface.co/collections/google/gemma-2-release-667d6600fd5
220e7b967f315) компании Google.
Кроме того, существуют более старые семейства:
GPT-J-6B (https://huggingface.co/EleutherAI/gpt-j-6b) исследо
вательской группы EleutherAI и предшествующее ему GPT-Neo
с размерностями моделей 125 млн (https://huggingface.co/Eleu
therAI/gpt-neo-125m), 1,3 млрд (https://huggingface.co/Eleuthe
rAI/gpt-neo-1.3B) и 2,7 млрд (https://huggingface.co/EleutherAI/
gpt-neo-2.7B);
языковая модель BLOOM (BigScience Large Open-science Openaccess Multilingual Language Model), содержащая до 176 млрд
параметров (https://huggingface.co/models?other=bloom&searc
h=bigscience);
модель GPT-2 (https://huggingface.co/gpt2-xl) компании OpenAI.
Компания прекратила публикацию своих моделей после GPT-2,
но даже эта модель остается хорошей отправной точкой для
тренировки пользовательской LLM.
Другие (не GPT-подобные) модели трансформеров:
около 500 моделей компании Google, в том числе самые по
следние из семейства T5 – UMT5 (https://huggingface.co/google/
umt5-xxl), поддерживающая 102 языка, и инструкционная,
точно настраиваемая FLAN-T5 с размерностью до 11,3 млрд
(https://huggingface.co/google/flan-t5-xxl), а также многие дру
гие интересные модели, например модель для реферирования
PEGASUS (https://huggingface.co/google/pegasus-large);
огромное количество общих и специализированных BERT-по
добных моделей, таких как FinBERT для анализа тональности
(эмоциональной окраски) финансовых текстов (https://hug
gingface.co/ProsusAI/fnbert) или BioBERT для майнинга био
медицинских текстов (https://huggingface.co/dmis-lab/biobertv1.1), и это далеко не полный список;
модели генерации изображений DALL·E Mini (https://hugging
face.co/dalle-mini/dalle-mini) и DALL·E Mega (https://hugging
face.co/dalle-mini/dalle-mega);
модели распознавания речи Whisper (https://huggingface.co/
openai/whisper-large-v2) и объединения изображений с встраи
ванием текста CLIP (https://huggingface.co/openai/clip-vit-largepatch14-336) компании OpenAI.
Существует таблица, содержащая модели, реализованные в биб
лиотеке Hugging Face Transformers, в которой указана поддержка раз
личных фреймворков, а именно PyTorch, TensorFlow и JAX: https://
huggingface.co/docs/transformers/main/index#supported-frameworks.
Использование экосистемы Hugging Face
427
Мы начнем с простого примера использования предварительно
натренированной GPT-подобной модели из хранилища Hugging
Face Model Hub.
Для работы с примерами в этой главе, вероятнее всего, потребует
ся установка библиотеки поддержки трансформеров:
pip install -q git+https://github.com/huggingface/transformers.git
11.3.1 Использование предварительно натренированной
модели из хранилища Hugging Face Model Hub
В хранилище Hugging Face Model Hub доступно множество разно
образных GPT-подобных моделей с различными размерностями.
Мы начнем с модели GPT-J-6B (https://huggingface.co/EleutherAI/
gpt-j-6b), представляющей собой авторегрессивную языковую мо
дель с 6 млрд параметров для английского языка, тренированную
исследовательской группой EleutherAI на наборе данных The Pile.
На момент публикации этой книги было доступно множество более
новых и мощных моделей (рекомендую ознакомиться с семействами
Llama 3 и Gemma 2), и, вероятно, в скором времени их количество
существенно увеличится, поэтому я советую внимательно следить
за этой стремительно развивающейся областью машинного обуче
ния. Для демонстрационных целей мы остановимся на GPT-J-6B, но
основные принципы останутся неизменными для более новых мо
делей. Если GPT-J-6B слишком велика для вашего компьютера, то
можно попробовать более компактную модель, например Gemma 2B
(https://huggingface.co/google/gemma-2b).
The Pile – это многообразный набор данных размером 825 Гб с от
крытым исходным кодом для языкового моделирования, состоящий
из 22 объединенных высококачественных наборов данных меньше
го размера (https://pile.eleuther.ai/). Он также сформирован иссле
довательской группой EleutherAI. Если вы намерены тренировать
собственную модель LLM, то можете воспользоваться этим набором
данных.
Существует группа моделей LLM, тренированных на наборе дан
ных The Pile. Изначально группа EleutherAI тренировала GPT-3-по
добные языковые модели GPT-Neo с вариантами 125 млн, 1,3 млрд
и 2,7 млрд параметров. Затем была создана модель GPT-J-6B. На мо
мент ее выпуска GPT-J-6B являлась самой большой в мире языковой
моделью в GPT-3 стиле с открытым доступом к ней.
ПРЕДУПРЕЖ ДЕНИЕ Модель GPT-J-6B не предназначена для
развертывания без точной настройки, диспетчерского управ
ления и/или модерации. Сама по себе она не является готовым
к применению продуктом и не может использоваться для взаи
428
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
модействий с людьми. Например, модель может генерировать
вредоносный или оскорбительный текст. Настоятельно реко
мендуется предварительно оценить риски, связанные с каждым
конкретным вариантом использования этой модели.
GPT-J-6B использовала библиотеку Mesh Transformer JAX (https://
github.com/kingoflolz/mesh-transformer-jax/), созданную на основе
Haiku. В ней применяются операторы JAX xmap/pjit для распарал
леливания трансформеров модели. Проектное решение библио
теки обеспечивает масштабирование приблизительно до 40 млрд
параметров на устройствах TPUv3, а помимо этого должны приме
няться различные стратегии распараллеливания. С 2023 г. этот про
ект активно не разрабатывается, поэтому если вы рассматриваете
перспективу тренировки собственных моделей LLM с нуля, то ре
комендуется обратить внимание на более современные и активно
разрабатываемые библиотеки. Эту тему мы более подробно обсудим
в главе 12.
Требования к памяти GPU
Для работы с описанными выше моделями потребуется GPU или TPU
с достаточным объемом памяти. Для предварительной оценки требуемого размера памяти можно умножить количество параметров модели
(например, 1,3 млрд для GPT-Neo-1.3B или 6 млрд для GPT-J-6B) на размер 32-битового числа с плавающей точкой (FP32), являющегося типовым форматом для распределенных моделей. При размере параметра
4 байта полная модель GPT-Neo-1.3B потребует 1,3 млрд × 4 = 7,2 млрд
байт, т. е. приблизительно 7,2 Гб памяти только для хранения модели.
Для GPT-J-6B необходимо 6 млрд × 4 = 24 Гб только для хранения модели. Большинство дешевых широкодоступных карт GPU не обладают
таким объемом памяти. В действительности потребуется больше памяти
для хранения промежуточных активаций. А для тренировки и точной
настройки вы должны обеспечить еще больший объем памяти для хранения градиентов.
Для более крупных моделей, таких как GPT-J-6B, возможно, потребуется
более мощный и дорогой GPU, например A100 с 40 или 80 Гб памяти
или мультиконфигурация с несколькими GPU/TPU. Если в вашем компьютере установлена видеокарта с 16 Гб памяти типа NVIDIA T4 или RTX
4080, или даже карта RTX 3090|4090 c 24 Гб памяти, то попробуйте модель меньшего размера, например GPT-Neo. Можно даже запустить эти
модели для логического вывода результатов на CPU при достаточном
объеме памяти, но это будет очень медленно.
Существуют способы снижения объема реально используемой памяти,
например преобразование модели в формат 16-битовых чисел с плавающей точкой (тип float16 или bfloat16, описанный в подразделе 3.3.2).
Использование экосистемы Hugging Face
429
В приведенном ниже примере используется вариант GPT-J-6B, преобразованный в формат FP16, чтобы сэкономить некоторое пространство
памяти. Существуют еще более интенсивные варианты, такие как квантование (дискретизация) и преобразование модели в формат INT4, но
в этой книге мы не рассматриваем подобные методики.
Загрузка модели из Hugging Face Model Hub и использование ее
для генерации текста выполняются просто.
Библиотека трансформеров Hugging Face предоставляет простой
способ загрузки различных моделей под названием auto classes
(https://huggingface.co/docs/transformers/model_doc/auto#auto-class
es). Требуемую для применения архитектуру можно более или менее
точно узнать по имени предварительно натренированной модели,
которое передается в метод from_pretrained(). Механизм auto class
es извлекает соответствующую модель автоматически.
Существуют разнообразные классы auto classes для различных за
дач и для каждого внутреннего компонента, например для PyTorch,
TensorFlow или Flax. Более подробную информацию об auto class
es можно получить здесь: https://huggingface.co/docs/transformers/
main/autoclass_tutorial.
Для GPT-подобных моделей Flax имеется общий класс, пред
назначенный для генерации текста, с именем FlaxAutoModelForCausalLM (https://huggingface.co/docs/transformers/model_doc/auto#
transformers.FlaxAutoModelForCausalLM). При передаче в него имени
(в нашем случае "EleutherAI/gpt-j-6B") он создает экземпляр типа
FlaxGPTJForCausalLM. Для модели "EleutherAI/gpt-neo-1.3B" должен
применяться тип FlaxGPTNeoForCasualLM.
Для каждой модели требуется входной текст, предварительно об
рабатывающийся заданным способом, который выполняет токени
затор: разделяет текст на токены, или лексемы (обычно это части
слов), которые, в свою очередь, обрабатываются моделью, генери
рующей продолжение входного текста, лексему за лексемой. Модели
могут использовать разнообразные стратегии токенизации и раз
личные словари токенов (лексем). Класс AutoTokenizer автоматиче
ски выполняет задачу выбора правильного токенизатора.
Теперь мы готовы к загрузке модели из Hugging Face Model Hub.
Листинг 11.13
Загрузка модели GPT-J-6B из Hugging Face Model Hub
from transformers import FlaxAutoModelForCausalLM, AutoTokenizer
import jax
❶
model_name = "EleutherAI/gpt-j-6B"
❷
model = FlaxAutoModelForCausalLM.from_pretrained(
model_name,
❸
430
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
revision="float16",
dtype=jax.numpy.float16,
from_pt=True)
tokenizer = AutoTokenizer.from_pretrained(model_name)
model.params = model.to_fp16(model.params)
model
❹
❺
❻
❼
❽
>>> <transformers.models.gptj.modeling_flax_gptj.FlaxGPTJForCausalLM
at 0x7c06953b3670>
❾
tokenizer
>>> GPT2TokenizerFast(name_or_path='EleutherAI/gpt-j-6B',
➥vocab_size=50257, model_max_length=2048, is_fast=True,
➥padding_side='right', truncation_side='right',
➥special_tokens={'bos_token': AddedToken("
➥<|endoftext|>", rstrip=False, lstrip=False,
➥single_word=False, normalized=True), 'eos_token':
➥AddedToken("<|endoftext|>", rstrip=False,
➥lstrip=False, single_word=False, normalized=True),
➥ 'unk_token': AddedToken("<|endoftext|>", rstrip=False,
➥ lstrip=False, single_word=False, normalized=True)},
➥ clean_up_tokenization_spaces=True)
❿
❶ Импорт библиотеки трансформеров Hugging Face.
❷ Определение имени используемой модели.
❸ Механизм auto class Flax для моделирования каузального (причинно-следствен-
ного) языка определяет и загружает соответствующий класс для модели GPT-Neo.
❹ Загрузка модифицированной версии FP16 используемой модели (такая версия
доступна не для каждой модели).
❺ Определение типа данных для вычисления; здесь необходимо использовать тип
FP16.
❻ Преобразование весов модели из сохраненной модели PyTorch (поскольку эта
конкретная модель не содержит отдельных весов JAX).
❼ Автоматическое определение правильного токенизатора для модели.
❽ Преобразование весов модели в FP16 (здесь в таком преобразовании нет не-
обходимости, это просто пример для других моделей, для которых недоступна
модифицированная версия FP16).
❾ Механизм auto class для FlaxAutoModelForCausalLM возвращает специализированный класс FlaxGPTNeoForCausalLM.
❿ Механизм auto class для выбора токенизатора возвращает специализированный
класс GPT2TokenizerFast.
В рассматриваемом здесь примере мы преднамеренно загрузили
версию модели с весами типа FP16 (параметр revision). Подобные
версии с меньшей точностью доступны не для всех моделей, поэтому
для некоторых моделей потребуется явное преобразование из фор
мата FP32 в формат FP16. Обратите внимание на два важных момен
та. Во-первых, параметр dtype экземпляра FlaxAutoModelForCausalLM
Использование экосистемы Hugging Face
431
определяет тип данных вычисления. Он не влияет на типы весов са
мой модели. Поэтому если версии FP16 недоступны для некоторого
варианта, то в подобном случае потребуется отдельное явное преоб
разование с применением метода to_fp16(). Также существует ме
тод to_bf16() для преобразования в тип BF16.
Следует еще отметить, что общий класс FlaxAutoModelForCausalLM
возвращает специализированный класс FlaxGPTJForCausalLM, а об
щий класс AutoTokenizer возвращает специализированный класс
токенизатора GPT2-TokenizerFast. Это не ошибка, поскольку модели
GPT-Neo и GPT-J от EleutherAI используют тот же токенизатор, что
и модель OpenAI GPT-2. Класс токенизатора также содержит важные
свойства, такие как размер языкового словаря, список специальных
токенов для пометки начала и конца предложения и т. п.
Большие языковые модели и модели каузального языка
Термин «большие языковые модели» (large language models – LLM)
имеет весьма широкий смысл и охватывает многочисленные разнообразные модели, работающие с естественным языком. В настоящее
время термин LLM наиболее часто используется для семейства GPTподобных моделей компании OpenAI, в которых применяется архитектура transformer decoder (трансформер-декодер). Существуют и другие
модели, использующие архитектуру transformer encoder (трансформеркодировщик), например BERT, или T5 – полноценный трансформер-кодировщик-декодер (transformer encoder-decoder). Джей Аламмар (Jay
Alammar) опубликовал великолепную подборку материалов о различных архитектурах трансформеров:
о полноценном трансформере-кодировщике-декодере (http://jalam
mar.github.io/illustrated-transformer/);
о трансформере-кодировщике (https://jalammar.github.io/a-visualguide-to-using-bert-for-the-first-time/);
о трансформере-декодере (http://jalammar.github.io/illustrated-gpt2/).
Кроме того, существуют модели без трансформеров, использующие другие архитектуры, например рекуррентные нейронные сети, такие как
RWKV (https://github.com/BlinkDL/RWKV-LM), или модели пространства состояний (state-space models), например Mamba (https://gonzo
ml.substack.com/p/mamba-linear-time-sequence-modeling).
Если говорить кратко, GPT-подобные большие языковые модели выполняют задачу доработки (дополнения, доведения до конечной формы)
входного текста. Вы передаете входной текст, обычно называемый текстовым запросом (prompt), и модель генерирует окончательную форму
предоставленного запроса. Такие большие языковые модели иногда называют каузальными (причинно-следственными) языковыми моделями
(causal language models – CLM).
432
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
С технической точки зрения на нижнем уровне модель выполняет задачу прогнозирования следующей лексемы (токена – token) следующим
способом. Сначала входной текст токенизируется, т. е. разделяется на
токены (лексемы) в соответствии со словарем токенизатора. Обычно
токены (лексемы) представляют собой элементы частей слов (подслов – subwords), часто встречающихся в тренировочном наборе данных. Затем последовательность текстовых лексем передается на вход
модели, а модель внутри себя формирует последовательность векторных представлений лексем (embeddings) в каждом слое. В последнем
слое CLM возвращает распределение вероятностей для следующей
лексемы, в соответствии с которым и выбирается следующая лексема.
Затем выбранная лексема добавляется в текстовый запрос, и процесс
повторяется.
Существуют многочисленные стратегии выбора следующей лексемы
(или стратегий декодирования). Можно воспользоваться жадным декодированием (greedy decoding), при котором выбирается лексема
с максимальной вероятностью, полиномиальной выборкой (multinomial
sampling) со случайным выбором лексемы из предложенного распределения или так называемым лучевым поиском, или поиском по пучку
траекторий (beam search), сохраняющим несколько предположений на
каждом временнóм шаге, с выбором в конце процедуры предположения с максимальной общей вероятностью во всей последовательности.
Другие типы декодирования описаны здесь: https://huggingface.co/
docs/transformers/generation_strategies.
Параметр temperature часто используется в методах выборки. Он
управляет степенью заострения кривой распределения вероятностей.
Когда параметр temperature равен нулю, форма кривой распределения
более заостренная. При повышении «температуры» (обычные возможные значения от 0,7 до 1,0) кривая распределения сглаживается, и вы
получаете более высокие шансы на выбор редко встречающихся лексем
(с более низкими вероятностями). В некотором смысле при «высоких
температурах» модель становится менее детерминированной и более
«креативной». Превосходное визуальное объяснение смысла парамет
ра temperature предоставлено здесь: https://blog.lukesalamone.com/
posts/what-is-temperature/.
Для получения более подробной информации рекомендуется воспользоваться руководством по LLM компании Hugging Face (https://hug
gingface.co/docs/transformers/main/llm_tutorial) и обратить внимание
на экосистему LLM с открытым исходным кодом (https://huggingface.
co/blog/os-llms).
Итак, мы готовы к использованию модели для генерации текста.
Воспользуемся текстовым запросом на генерацию SQL-запроса по
текстовому описанию того, что необходимо сделать. Будет выпол
нена токенизация текстового запроса и применена стратегия гене
Использование экосистемы Hugging Face
433
рации полиномиальной выборки, но вы можете применить любой
другой метод по своему выбору.
Листинг 11.14 Использование модели GPT-J-6B для генерации текста
prompt = """Generate SQL query from the text inside triple backticks
```select all the users from table 'users' older than 20 years```
"""
inputs = tokenizer(prompt, return_tensors="jax")
inputs
❶
❷
>>> {'input_ids': Array([[ 8645, 378, 16363, 12405, 422, 262, 2420,
...
>>>
812, 15506,
63, 198]], dtype=int32),
>>> 'attention_mask': Array([[1, 1, 1, … 1, 1, 1]], dtype=int32)}
generated_ids = model.generate(
**inputs,
do_sample=True,
num_beams=1,
max_new_tokens=300,
temperature=0.7,
pad_token_id = model.config.eos_token_id,
prng_key=jax.random.PRNGKey(4232),
no_repeat_ngram_size=2)
generated_text = tokenizer.decode(
generated_ids['sequences'].squeeze(0))
print(generated_text)
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
❸
❹
❹
❺
❻
❼
❽
❾
Generate SQL query from the text inside triple backticks
```select all the users from table 'users' older than 20 years```
```
SELECT * FROM `users` WHERE age > 20
```
```
SELECT * FROM `users` WHERE age > `20 years`
```
```
SELECT * FROM `users` WHERE `age` > 20`
```
<!--end-query-->
…
❶ Определение текстового запроса (prompt).
❷ Токенизация текстового запроса.
434
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
Генерация продолжения текстового запроса.
Использование стратегии генерации полиномиальной выборки.
Определение максимальной длины генерируемого текста.
Установка температуры модели (параметра, определяющего «креативность» модели).
❼ Предоставление ключа PRNGKey для случайной выборки.
❽ Декодирование последовательности идентификаторов лексем в текстовые лексемы.
❾ Вывод модели содержит текстовый запрос, сопровождаемый сгенерированным
текстом.
❸
❹
❺
❻
В приведенном выше примере мы разделили на лексемы тексто
вый запрос с применением соответствующего токенизатора. Основ
ным результатом этого шага стала всего лишь последовательность
чисел, содержащихся в поле input_ids. Другое поле attention_mask
помечает позиции лексем, на которые должен обратить внимание
трансформер-декодер; в нашем случае это каждая позиция в тексто
вом запросе. Будьте внимательны и не используйте токенизаторы
из других моделей, так как они будут выполнять разделение текста
по-другому с отличающимися идентификаторами лексем, и модель
интерпретирует их некорректно.
Затем мы устанавливаем значения нескольких параметров для
генерации, при этом температура и параметры, связанные со стра
тегией декодирования, являются особенно важными. Кроме того,
следует отметить, что здесь мы предоставили ключ PRNGKey для
обеспечения воспроизводимости. Далее модель генерирует последо
вательность идентификаторов лексем, которая декодируется обрат
но в текстовые лексемы с помощью того же токенизатора. Обратите
внимание: входной текстовый запрос также включен в вывод резуль
тата. Модель тоже продолжает генерацию того же результата после
метки <!--end-query-->, но эта часть вывода здесь не показана.
Полученный конкретный вывод результата в рассматриваемом
здесь примере текстового запроса содержит кое-что важное, хотя
конечный результат может оказаться не совсем точно соответству
ющим ожидаемому, и обычно его невозможно использовать в таком
виде. В данном случае мы не сможем применить сгенерированные
команды SQL напрямую по следующим причинам:
1 их
окружает отформатированный текст, поэтому они не явля
ются непосредственно выполняемыми;
2 в выводе модели содержатся различные ответы (к тому же мо
дель продолжает повторно генерировать текст);
3 некоторые команды SQL содержат ошибки, например для ко
нечной одиночной кавычки в третьем варианте вывода отсут
ствует соответствующая открывающая одиночная кавычка.
Точная настройка текстового запроса и получение именно того
результата, который требуется от модели, – это отдельная тема,
Использование экосистемы Hugging Face
435
обычно обозначаемая термином «инженерия составления тексто
вых запросов к большим языковым моделям», или просто «инжене
рия запросов» (prompt engineering), но эта тема не рассматривается
в данной книге.
Инженерия составления текстовых запросов
к большим языковым моделям
Инженерия составления текстовых запросов к большим языковым моделям, или просто инженерия запросов, – это набор методик и технологий
для структурирования текстовых запросов (входных текстов) с целью
получения требуемого результата. Существуют многочисленные разнообразные технологии, из которых самой широко известной является
«цепочка умозаключений» (chain of thought) (https://research.google/
blog/language-models-perform-reasoning-via-chain-of-thought/). Эта
область непрерывно развивается, поэтому рекомендуется внимательно
следить за ней.
Чтобы упростить начало работы с этими методиками и технологиями,
рекомендуется пройти короткий вводный курс по инженерии запросов
для разработчиков: https://www.deeplearning.ai/short-courses/chatg
pt-prompt-engineering-for-developers/.
Применения инженерии запросов не всегда достаточно; наблюдается
рост тенденции включения LLM в программные продукты, например см.
«Large Language Model Programs» (https://arxiv.org/abs/2305.05364).
Программы разделяют решение на отдельные шаги, управляют состоянием по шагам, подготавливают текстовые запросы и выполняют другие
операции оркестровки.
Вот и все об основах использования LLM. Но использование пред
варительно натренированных моделей – это всего лишь часть общей
картины. Во многих случах этого недостаточно для получения тре
буемого качества, поэтому нужно создание специализированных
решений.
11.3.2 Более подробное изучение процессов точной
настройки и предварительной тренировки
При решении задач с применением больших языковых моделей
(LLM) вы обычно начинаете с предварительно натренированных мо
делей и инженерии составления текстовых запросов (далее – прос
то «инженерия запросов»). В некоторых случаях, когда полученное
качество результата неудовлетворительное, требуется нечто боль
шее, нежели предварительно натренированные модели. Необходи
мо специализировать модели, и для этого существуют два основных
способа:
436
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
тренировка новой модели с нуля. В контексте LLM этот способ
также называют предварительной тренировкой. Такой про
цесс обычно требует огромного объема тренировочных дан
ных и длительного времени для тренировки с использованием
конфигураций с многочисленными GPU/TPU. В зависимости от
требуемых ресурсов цена такой тренировки может составлять
от десятков тысяч до десятков миллионов долларов, затрачен
ных на вычисления. Позволить себе такие расходы могут лишь
немногие пользователи/компании;
точная настройка существующей предварительно натрениро
ванной модели. При таком подходе основой служит предвари
тельно натренированная модель стороннего производителя,
а вы выполняете точную настройку (или адаптацию) модели
для своей конкретной задачи. Иногда это имеет смысл, даже
если вы сами предварительно тренируете модель, поскольку
в процессе предварительной тренировки обычно используют
ся свободно доступные общие данные огромного объема, а при
точной настройке применяется более скромный объем специа
лизированных данных. Для точной настройки существуют раз
нообразные методики – от классического трансферного обуче
ния (дообучения) с обновлениями всех или некоторых слоев
нейросети до низкоуровневой адаптации (low-rank adaptation –
LoRa), при которой формируются небольшие корректировки
(обозначаемые термином «дельта-файл» (delta file)) исходных
весов. Как правило, такой подход гораздо дешевле и практич
нее; иногда можно получить полезную точно настроенную мо
дель при затратах всего лишь нескольких сотен долларов на ис
пользование кластера GPU в облачной среде. Кроме того, можно
применить обучение с подкреплением на основе обратной свя
зи от человека (reinforcement learning from human feedback –
RLHF) для точной настройки моделей.
Генерация ответа, дополненная результатами поиска (retrievalaugmented generation – RAG), представляет собой дополнение к точ
ной настройке и может оказаться простым способом сделать пред
варительно натренированную модель более актуальной. Тут мы не
будем рассматривать методику RAG, но ознакомиться с ней можно,
например, здесь: https://www.rungalileo.io/blog/optimizing-llm-per
formance-rag-vs-finetune-vs-both.
Для тех, кто намерен дальше продвигаться в описанных выше на
правлениях, предлагается несколько примеров и руководств, опи
сывающих, как использовать библиотеку Hugging Face Transformers:
руководство по предварительной тренировке GPT-подобной
каузальной языковой модели (https://colab.research.google.
com/github/huggingface/notebooks/blob/master/examples/causal_
language_modeling_flax.ipynb) описывает тренировку авторе
Использование экосистемы Hugging Face
437
грессивной каузальной языковой модели версии, отделенной
от GPT-2, на одном из языков огромного многоязычного кор
пуса текстов OSCAR, полученного посредством классификации
языков и фильтрации корпуса текстов Common Crawl. Руковод
ство включает определение архитектуры модели, трениров
ку токенизатора, предварительную обработку набора данных
и предварительную тренировку модели на TPUv3-8;
тема, не вполне относящаяся к использованию библиотеки
Hugging Face Transformers: превосходное руководство от Google
по использованию современной общедоступной LLM Gemma
(https://huggingface.co/blog/gemma) для логического вывода ре
зультата (https://ai.google.dev/gemma/docs/jax_inference) и точ
ной настройки (https://ai.google.dev/gemma/docs/jax_fnetune)
с использованием JAX и Flax;
глава 12 этой книги содержит раздел с описанием модулей JAX
для тренировки собственных LLM;
руководство по предварительной тренировке BERT-подобной
маскированной языковой модели (https://colab.research.google.
com/github/huggingface/notebooks/blob/master/examples/
masked_language_modeling_flax.ipynb) описывает тренировку
модели на основе RoBERTa с использованием того же набора
данных OSCAR;
в руководстве по точной настройке предварительно натре
нированной модели BERT для задачи классификации текстов
с применением GLUE Benchmark (https://colab.research.google.
com/github/huggingface/notebooks/blob/master/examples/text_
classification_flax.ipynb) используется класс FlaxAutoModelFor
SequenceClassification с недавно инициализированным но
вым группированием рубрик классификации;
дополнительные примеры языковых моделей Flax (https://
github.com/huggingface/transformers/blob/main/examples/flax/
language-modeling/README.md) включают тренировку дру
гой модели на основе RoBERTa и небольшой версии GPT-2 для
норвежского языка, маскированного по интервалам языкового
моделирования на основе T5 и языкового моделирования с шу
моподавлением на основе BART;
многочисленные дополнительные примеры захвата изобра
жений (https://github.com/huggingface/transformers/tree/main/
examples/flax/image-captioning), ответов на вопросы (https://
github.com/huggingface/transformers/tree/main/examples/flax/
question-answering), резюмирования (https://github.com/hug
gingface/transformers/tree/main/examples/flax/summarization),
классификации текстов (https://github.com/huggingface/trans
formers/tree/main/examples/flax/text-classification), токенов (лек
сем) (https://github.com/huggingface/transformers/tree/main/ex
438
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
amples/flax/token-classification) и изображений (https://github.
com/huggingface/transformers/tree/main/examples/flax/vision)
с использованием Vision Transformer.
Существует множество примеров, но не имеет смысла воспроиз
водить все примеры в этой книге. Это хороший отправной пункт для
начала ваших собственных небольших экспериментов и расшире
ния их при наличии достаточного материального обеспечения.
Еще одна обширная тема: широко распространенные диффузион
ные модели преобразования текста в изображение, такие как DALL·E 2,
Stable Diffusion или Midjourney, генерирующие изображения на ос
нове текстовых описаний. Эта тема заслуживает выделения в особый
подраздел, поскольку для таких моделей существуют специализиро
ванные библиотеки диффузоров в экосистеме Hugging Face.
11.3.3 Использование библиотеки диффузоров
Diffusers – это библиотека для современных предварительно натре
нированных диффузионных моделей для генерации изображений,
звуковых фрагментов и даже трехмерных структур молекул. Биб
лиотека Diffusers предоставляет простое решение интерфейса, а так
же средства тренировки пользовательских диффузионных моделей.
Диффузионные модели
Диффузионные модели – это современная «физическая» методика генерирования изображений. Рекомендуется прочесть превосходную статью
на эту тему в журнале Quanta Magazine (https://www.quantamagazine.
org/the-physics-principle-that-inspired-modern-ai-art-20230105/).
Существуют многочисленные широко известные модели и программные продукты, такие как DALL·E 3, Stable Diffusion, Midjourney и т. д.,
и эта сфера быстро развивается.
Подобные модели обычно кодируют текстовое описание с использованием некоторого кодировщика текста, вспомогательное (необязательное) изображение с помощью некоторого кодировщика изображений,
а затем выполняют диффузионный процесс генерации изображения,
состоящий из нескольких шагов.
Ранняя история создания моделей генерации изображений представлена
в моем блоге: https://moocaholic.medium.com/openai-and-the-roadto-text-guided-image-generation-dall-e-clip-glide-dall-e-2-unclip-c6e
28f7194ea. Джей Аламмар (Jay Alammar) также предоставил превосходный
видеоматериал, описывающий внутренние механизмы диффузионных
моделей: https://jalammar.github.io/illustrated-stable-diffusion/.
Диффузионные модели применяются и для генерации видео, синтеза
звуков и т. п.
Использование экосистемы Hugging Face
439
Библиотека Diffusors предоставляет три основных компонента:
современные диффузионные конвейеры, которые могут вы
полнять логический вывод результата всего лишь в нескольких
строках кода. Существуют конвейеры для генерации изображе
ний с использованием текстовых запросов, для видоизменения
изображений, для ретуширования дефектов, для масштабиро
вания, для генерации трехмерных объектов, для преобразова
ния текста в звук (речь) и даже для генерации видео из текста,
а также для решения многих других задач;
взаимозаменяемые планировщики шума для различных скоро
стей диффузии и качества конечного результата;
предварительно натренированные модели, которые можно ис
пользовать как конструктивные элементы, а также при объеди
нении с планировщиками для создания собственных полно
стью завершенных диффузионных систем. Model Hub содержит
тысячи (на момент написания книги – более 30 300) предвари
тельно натренированных диффузионных моделей (https://hug
gingface.co/models?library=diffusers&sort=downloads).
Для установки библиотеки Diffusers с поддержкой Flax (библиотека
поддерживает Flax с версии 0.5.1) выполните следующую команду:
pip install --upgrade diffusers[flax]
Вы намерены разрабатывать что-то подобное Midjourney? Такую
разработку легко начать с использованием существующих моделей
для генерации собственных изображений и создания своего конвейе
ра генерации изображений. Но при этом следует помнить о провер
ке лицензии модели. Мы воспользуемся хорошо известной моделью
Stable Diffusion, созданной исследователями и инженерами из групп
CompVis, Stability AI, LAION и Runway (https://ru.wikipedia.org/wiki/
Stable_Diffusion). Существует несколько версий этой модели; мы
будем работать с версией runwayml/stable-diffusion-v1-5 (https://
huggingface.co/runwayml/stable-diffusion-v1-51), которая способна
генерировать изображения с разрешением 512×512. Генерация будет
выполняться на восьми TPU в параллельном режиме, но вы можете
без затруднений изменить код для работы на GPU или CPU.
В рассматриваемом здесь примере используется тип с плавающей
точкой BF16, поддерживаемый TPU. Для работы применяется вирту
альная машина Cloud TPU VM (описанная в приложении C) и соеди
нение блокнота Colab с локальным ядром. Чтобы адаптировать код
для работы на GPU, возможно, потребуется переход на тип FP16 или
1
Попытка перехода на эту страницу завершается «ошибкой 404» (страница
не найдена), и предлагается выполнить поиск в списке моделей. – Прим.
перев.
440
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
FP32 в зависимости от конкретной модели GPU и поддерживаемых
числовых форматов. Код можно выполнять и на CPU, но при этом
нужно удалить части, выполняющие сегментирование и репли
кацию.
Эта область применения нейросетей развивается чрезвычайно
быстро, и когда вы будете читать этот текст, вероятно, появятся но
вые, еще более эффективные модели генерации изображений, а биб
лиотека диффузоров существенно изменится. Поэтому рекоменду
ется как можно чаще обращаться к документации библиотеки.
Библиотека диффузоров предоставляет класс FlaxStableDiffusionPipeline (https://huggingface.co/docs/diffusers/v0.20.0/en/api/pipelines/
stable_diffusion/text2img#diffusers.FlaxStableDiffusionPipeline) для использования моделей Stable Diffusion. Как обычно, в функциональ
ных фреймворках модели не имеют состояния и параметры хра
нятся вне модели, поэтому метод инициализации конвейера воз
вращает сам конвейер и параметры модели. Для работы с классом
FlaxStableDiffusionPipeline требуется библиотека трансформеров,
поэтому необходимо установить ее в первую очередь, если она от
сутствует в вашей рабочей среде.
Общая методика: загрузка весов модели и ее инициализация,
предварительная обработка текста, репликация весов модели
и предварительно обработанного текста запроса на каждом устрой
стве, наконец, генерация нескольких изображений в параллельном
режиме.
Мы будем использовать функцию prepare_inputs() для токени
зации текстов запроса, функцию flax.jax_utils.replicate() для
репликации параметров модели на каждом устройстве и функцию
flax.training.common_utils.shard() для сегментирования дерева
массивов pytree с распределением сегментов по локальным устрой
ствам (их количество определяется функцией local_device_count()).
Реализация этой методики показана в листинге 11.15.
Листинг 11.15 Генерация изображений с использованием библиотеки
диффузоров
from jax.lib import xla_bridge
print(xla_bridge.get_backend().platform)
>>> tpu
import jax
import jax.numpy as jnp
from flax.jax_utils import replicate
from flax.training.common_utils import shard
from PIL import Image
from diffusers import FlaxStableDiffusionPipeline
❶
❶
❷
❷
❷
❷
❷
❷
441
Использование экосистемы Hugging Face
pipeline, params = FlaxStableDiffusionPipeline.from_pretrained(
"runwayml/stable-diffusion-v1-5",
revision="bf16",
dtype=jnp.bfloat16,
)
❸
❸
❹
❹
prompts = [
❺
"Three leopards holding a scarlet rose and in the centre a green rat.",
"A large glass bottle full of tiny elephants of different colors",
"HAL-9000 in the style of Picasso",
"Pink panther drinking a coffee. ",
"Golden tiger fighting vikings in a modern city, Van Gogh style",
"A blue Axolotl sitting on an emerald throne.",
"A pumpkin-shape spaceship spinning around the Moon",
"Black panther doing his homework",
]
p_params = replicate(params)
prompt_ids = pipeline.prepare_inputs(prompts)
prompt_ids = shard(prompt_ids)
seed = 42
rng = jax.random.PRNGKey(seed)
rng = jax.random.split(rng, jax.device_count())
images = pipeline(prompt_ids, p_params, rng, jit=True).images
❻
❼
❼
❽
❽
❽
❾
def image_grid(imgs, rows, cols):
❿
w,h = imgs[0].size
grid = Image.new('RGB', size=(cols*w, rows*h))
for i, img in enumerate(imgs): grid.paste(img, box=(i%cols*w, i//cols*h))
return grid
images = images.reshape((images.shape[0], ) + images.shape[-3:])
images = pipeline.numpy_to_pil(images)
image_grid(images, 2, 4)
❶
❷
❸
❹
❺
❻
❼
❽
❾
❿
⓫
⓬
⓫
⓬
⓬
Обеспечение использования TPU.
Импорт всех необходимых компонентов.
Загрузка модели Stable Diffusion 1.5.
Использование формата с плавающей точкой BF16 (вы можете заменить его на FP16 для
GPU).
Подготовка восьми текстов запросов.
Репликация параметров модели на каждом ядре TPU.
Предварительная обработка текстов запросов и сегментирование их по восьми TPU.
Подготовка восьми ключей PRNGKey.
Запуск конвейера и возврат сгенерированных изображений.
Функция для построения сетки изображений.
Удаление ненужного пакетного измерения.
Вывод изображений.
442
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
Здесь можно видеть, что код получился достаточно компакт
ным – части загрузки модели и генерации изображений занимают
всего лишь один экран. Генерация новых изображений выполняется
просто.
Первый запуск конвейера может занять некоторое время, так как
в этот момент происходит JIT-компиляция (при установленном фла
ге jit=True). Этот параметр также отвечает за выполнение pmap()версий функций генерации и оценки безопасности. Вывод результа
тов работы модели показан на рис. 11.3.
Рис. 11.3
Восемь изображений, сгенерированных моделью Stable Diffusion
Для конвейера также существуют и другие параметры, воздейст
вующие на процесс генерации изображений, поэтому настоятельно
рекомендуется внимательно изучить документацию класса FlaxStableDiffusionPipeline. Внутреннее устройство модели Stable Dif
fusion и смысл параметров, воздействующих на процесс генерации
изображений, подробно описан в одной из статей блога https://hug
gingface.co/blog/stable_diffusion.
Информацию о многочисленных доступных конвейерах из биб
лиотеки диффузоров можно получить здесь: https://huggingface.co/
docs/diffusers/api/pipelines/overview.
Библиотека диффузоров представляет собой нечто гораздо боль
шее, нежели простая коллекция предварительно натренированных
моделей. Вы можете тренировать с нуля новые диффузионные мо
дели или выполнять точную настройку существующих моделей на
собственных наборах данных. В комплект официальной докумен
тации входят подробные руководства по этому процессу, и если вы
серьезно интересуетесь этой темой, то рекомендуется начать с базо
вого руководства по тренировке (https://huggingface.co/docs/diffus
ers/tutorials/basic_training), затем перейти к основной части (https://
huggingface.co/docs/diffusers/training/overview).
Резюме
443
Существует множество интересных примеров, и библиотека
поддерживает не только изображения, но и другие типы и форма
ты данных. Особенно хотелось бы выделить Low-Rank Adaptation
of LLM (LoRA). Это методика тренировки, ускоряющая точную на
стройку больших моделей при потреблении меньшего объема па
мяти – чрезвычайно широко распространенный подход для точной
настройки моделей Stable Diffusion. На эту тему существует велико
лепный пост в блоге Педро Куэнса (Pedro Cuenca) и Саяка Пола (Sayak
Paul): https://huggingface.co/blog/lora.
Высокая эффективность использования памяти методики LoRA
позволяет выполнять точную настройку на обычных потребитель
ских GPU, таких как Tesla T4, RTX 3080 или даже RTX 2080 Ti. Такие
GPU, как T4, свободно доступны и готовы к работе в блокнотах Kaggle
и Google Colab. LoRA пока еще остается экспериментальным инстру
ментальным средством в библиотеке диффузоров, и ее API в буду
щем может измениться.
Это все, что я хотел сказать об экосистеме Hugging Face и о прак
тическом использовании высокоуровневых библиотек поддержки
нейронных сетей на текущий момент. Теперь вы обладаете прочной
основой и инструментальными средствами для более глубокого по
гружения в Flax, любую другую библиотеку высокого уровня на основе
JAX или в разнообразную экосистему Hugging Face. В следующей главе
представлен подробный обзор экосистемы JAX, включающий описа
ние инструментальных средств для тренировки собственных LLM.
Резюме
Высокоуровневые библиотеки предоставляют базисные элементы
высокого уровня, помогающие создавать нейронные сети из су
ществующих конструктивных блоков, таких как плотные, конво
люционные или LSTM-слои; многоцелевые слои самовнимания;
функции активации и т. п.
Одной из наиболее широко известных высокоуровневых библио
тек поддержки нейронных сетей является Flax компании Google.
Другие варианты – Keras и Equinox.
Основополагающий принцип («философия») Flax – предоставле
ние API, знакомого тем, кто имеет опыт работы с Keras, PyTorch
или Sonnet.
Flax API называется Linen. По своей основной сущности это функ
циональная система для определения нейронных сетей в JAX, от
личающаяся от большинства объектно ориентированных подхо
дов, применяемых в экосистемах TensorFlow и PyTorch.
В Flax абстракция Module внутри себя является сущностью с со
хранением состояния, а для внешней среды – сущностью без
444
Глава 11
Высокоуровневые библиотеки поддержки нейронных сетей
сохранения состояния. Вы пишете код своих нейросетей в объ
ектно ориентированном стиле с сохранением состояния. Но для
использования трансформаций JAX Flax создает чистые функции
из абстракций Module. Эти функции не поддерживают сохранение
состояния, они принимают и возвращают его.
Переменные модели (или параметры) инициализируются мето
дом init() созданного экземпляра Module.
Для управления прямым проходом модели с заданным набором па
раметров вызывается метод apply() созданного экземпляра Module
с передачей инициализированных переменных и входных данных.
Существует класс flax.training.train_state.TrainState для пред
ставления простого состояния тренировки.
Optax – это библиотека для компонуемых трансформаций гра
диентов. Она предоставляет обширный список предварительно
определенных современных оптимизаторов, функций потерь,
планировщиков, классов и оберток для компонуемых трансфор
маций градиентов и многих других операций.
Библиотека CLU содержит общую функциональность для написа
ния циклов тренировки машинного обучения, включая функцио
нальный интерфейс вычисления метрик.
Flax поддерживает модели с внутренним состоянием, например
слои BatchNorm. Как и параметры модели, состояние должно пере
даваться извне, чтобы обеспечить отсутствие внутреннего состоя
ния для классов и функций.
Hugging Face предоставляет превосходные библиотеки трансфор
меров и диффузоров, а также огромную коллекцию Model Hub,
содержащую тысячи моделей с открытым исходным кодом. По
количеству поддерживаемых моделей в этой коллекции JAX явля
ется третьим фреймворком.
Библиотека трансформеров Hugging Face предоставляет простой
способ загрузки различных моделей под названием auto classes.
Существует общий класс для GPT-подобных моделей Flax для ге
нерации текста: FlaxAutoModelForCausalLM.
Используя трансформеры Hugging Face, можно тренировать но
вые большие модели с нуля или выполнять точную настройку су
ществующих предварительно натренированных моделей.
Diffusers Hugging Face – это библиотека для работы с современны
ми предварительно натренированными диффузионными моде
лями для генерации изображений, аудиозаписей и даже трехмер
ных структур молекул. Библиотека Diffusers представляет собой
простое решение для логического вывода результата и инстру
ментальные средства для тренировки пользовательских диффу
зионных моделей.
LoRA – это методика тренировки, ускоряющая процедуру точной
настройки больших моделей при потреблении меньшего объема
памяти.
12
Другие члены
экосистемы JAX
Темы главы:
библиотеки утилит, библиотеки для LLM и другие
высокоуровневые библиотеки поддержки нейронных
сетей для JAX;
библиотеки для других областей машинного
обучения, в том числе для обучения с подкреплением
и эволюционных вычислений;
прочие модули JAX для информатики, физики, химии
и многого другого.
Под экосистемой JAX я подразумеваю широкий набор библиотек,
модулей и пакетов, созданных на основе JAX или спроектированных
с обеспечением возможности взаимодействия с JAX либо с другими
членами экосистемы. Экосистема JAX – это полный жизни и дина
мичный организм с сотнями модулей, предназначенных для реше
ния задач во многих сферах деятельности. Как вы уже знаете, JAX
великолепен не только для решения задач глубокого обучения, но
также для более широких областей машинного обучения, оптимиза
ции и прочих дисциплин информатики.
Наблюдается рост объема применения JAX в других областях, тре
бующих больших вычислительных мощностей: для свертки (фол
Глава 12
446
Другие члены экосистемы JAX
динга) белков, для химического моделирования и особенно в физи
ке, включая молекулярную динамику, динамику жидкостей и газов,
моделирование недеформируемого твердого тела, квантовые вы
числения, астрофизику, моделирование экосистемы океана и т. д.
Постоянно появляются новые приложения.
В этой главе мы рассмотрим экосистему JAX в более широком
смысле. Начнем с глубокого обучения, затем перейдем к более об
щей теме машинного обучения и информатики в целом, а завершим
обзор приложениями в физике и связанных с ней дисциплинах.
Глава в большей степени является кратким каталогом, нежели
исчерпывающим описанием практического применения каждого
модуля. Всю экосистему описать невозможно, и одно только пред
ставление примеров использования всех модулей потребовало бы
отдельной книги. Поиск в GitHub по образцу JAX в репозиториях
Python (https://github.com/search?q=jax+language%3APython&type=
repositories&l=Python) сразу же выдает более 2000 результатов (на
момент написания этой книги). Ситуация меняется очень быстро,
поэтому информация из текущей главы, вероятнее всего, скоро ста
нет не вполне актуальной. Используйте содержимое этой главы как
отправную точку и справочник по ссылкам на официальную доку
ментацию модулей.
12.1 Экосистема глубокого обучения
В предыдущей главе упоминались многие различные модули для
глубокого обучения. Здесь предоставлен более широкий обзор мо
дулей в экосистеме, а также отмечаются некоторые устаревшие мо
дули, которые можно встретить в разных примерах.
12.1.1 Высокоуровневые библиотеки поддержки
нейронных сетей
Начнем с высокоуровневых библиотек поддержки нейронных сетей.
Мы уже применяли на практике библиотеку Flax, но здесь предлага
ется более развернутая картина экосистемы:
Flax (https://github.com/google/flax) – высокоуровневая библио
тека поддержки нейронных сетей. Это самый первый вариант
выбора, если необходимы абстракции более высокого уровня
для нейросетей по сравнению со средствами, предоставляемы
ми только ядром JAX. Мы подробно рассматривали и использо
вали эту библиотеку в главе 11;
Equinox (https://github.com/patrick-kidger/equinox) – универ
сальная библиотека JAX с поддержкой всех необходимых
Экосистема глубокого обучения
447
средств, отсутствующих в ядре JAX, включая нейронные сети
с простым синтаксисом в стиле PyTorch, отфильтрованные API
для трансформаций, полезные подпрограммы для работы с py
tree и продвинутые функциональные средства, такие как под
держка генерации исключений (ошибок) во время выполнения.
Вы сами убедитесь, насколько полезна эта библиотека;
Keras 3.0 (https://keras.io/keras_3/) – новое поколение широко
известной библиотеки Keras, обеспечивающей возможность
выполнения рабочих процессов Keras поверх TensorFlow, JAX
и PyTorch. Кроме того, поддерживается беспроблемная интег
рация компонентов Keras (слоев, моделей, метрик) в рабочие
процессы более низкого уровня TensorFlow, JAX и PyTorch;
Ivy (https://github.com/unifyai/ivy) – транскомпилятор для ма
шинного обучения и фреймворк, в настоящее время поддержи
вающий JAX, TensorFlow, PyTorch и NumPy. Ivy унифицирует все
фреймворки машинного обучения, позволяя не только писать
код, который можно использовать в любом из этих фреймвор
ков как внутренний компонент, но также выполняет преоб
разование (или транскомпиляцию) любой функции, модели
или библиотеки, написанной в одном из вышеперечисленных
фреймворков, в код для фреймворка, который вы предпочитае
те использовать;
AXLearn (https://github.com/apple/axlearn) компании Apple –
библиотека, созданная на основе JAX и XLA для поддержки
разработки крупномасштабных моделей глубокого обучения.
В ней принят объектно ориентированный подход к инженерии
программного обеспечения.
Существуют и другие библиотеки, некоторые из которых в настоя
щее время не разрабатываются активно, тем не менее на них стоит
обратить внимание:
Haiku (https://github.com/google-deepmind/dm-haiku) – биб
лиотека поддержки нейросетей на основе JAX от компании
Google DeepMind. С июля 2023 г. компания Google DeepMind
рекомендует в новых проектах использовать Flax вместо Hai
ku. Библиотека Haiku продолжает полноценно поддерживать
ся, но в целом проект переходит в режим сопровождения, т. е.
внимание разработчиков будет сосредоточено только на ис
правлении ошибок и поддержке совместимости с новыми вер
сиями JAX. Тем не менее можно найти многочисленные при
меры от тех, кто продолжает по-прежнему пользоваться Haiku
на практике;
Trax (https://github.com/google/trax) компании Google – полно
ценная библиотека для глубокого обучения, в которой основ
ное внимание сосредоточено на ясности кода и скорости. Trax
448
Глава 12
Другие члены экосистемы JAX
включает базовые модели (такие как ResNet, LSTM, Transformer)
и алгоритмы обучения с подкреплением (такие как REINFORCE,
A2C, PPO);
Objax (https://github.com/google/objax) компании Google – ми
нималистичный объектно ориентированный фреймворк с ин
терфейсом в стиле PyTorch. Название образовано из слов Object
и JAX. Objax спроектирован и создан исследователями для ис
следователей с обеспечением простоты и легкости понимания.
Пользователи этого фреймворка должны иметь возможность
без каких-либо затруднений читать, понимать, расширять
и модифицировать его код в соответствии со своими потреб
ностями;
Stax (https://docs.jax.dev/en/latest/jax.example_libraries.stax.html) –
компонент, изначально представляющий собой эксперимен
тальную часть JAX, но в настоящее время в большей степени
становящийся библиотекой примеров. Может оказаться полез
ным, если необходимо понять, как создать небольшую, но гиб
кую библиотеку спецификаций нейронных сетей с нуля;
Elegy (https://github.com/poets-ai/elegy) – API высокого уровня
для глубокого обучения в JAX от группы poets-ai. Но в настоя
щее время не наблюдается активная деятельность по разработ
ке этого API;
Mesh Transformer JAX (https://github.com/kingoflolz/mesh-trans
former-jax) – фреймворк, который использовался для создания
языковой модели GPT-J-6B. Предоставляет распараллеленные
трансформеры модели в JAX и Haiku; предназначен для обеспе
чения масштабируемости приблизительно до 20 млрд парамет
ров на TPUv3. С 2023 г. проект активно не разрабатывается.
Отдельным подмножеством высокоуровневых библиотек под
держки нейронных сетей являются библиотеки поддержки больших
языковых моделей (LLM).
12.1.2 Большие языковые модели в JAX
С учетом невероятного прогресса в развитии больших языковых
моделей (LLM) за последние несколько лет эта тема стала особенно
актуальной. Уже создано много коммерческих моделей, таких как
GPT-4 компании OpenAI, PaLm 2 и Gemini компании Google, семей
ство Anthropic Claude, модели Cohere, Amazon Titan и т. д. В настоя
щее время начинают выходить на передний план и LLM с открытым
исходным кодом – все релизы Llama 3, Gemma 2, Mistral и прочие
с многочисленными версиями со специализированной точной на
стройкой, создаваемыми сообществом, и это только начало.
Если вы работаете в крупной компании или получили грант на
использование Cloud GPU/TPU, то можете воспользоваться этими
Экосистема глубокого обучения
449
ресурсами для тренировки собственных LLM с нуля. Ниже представ
лено описание некоторых современных библиотек, которые могут
оказаться полезными для практического применения:
EasyLM (https://github.com/young-geng/EasyLM) – разработа
на группой Berkeley AI Research как полноценное решение для
предварительной тренировки, точной настройки, выполнения
и обслуживания (сопровождения) LLM в JAX/Flax. EasyLM спро
ектирована с обеспечением легкости практического примене
ния; она скрывает все сложности распараллеливания распреде
ленной модели и распараллеливания по данным, но при этом
остаются доступными подробности ядра тренировки LLM и ло
гического вывода результата, что обеспечивает легкость специа
лизированной настройки. EasyLM способна масштабироваться
для тренировки LLM на сотнях акселераторов TPU/GPU без не
обходимости написания сложного распределенного кода тре
нировки с применением функциональности JAX pjit. Авторы
EasyLM создали версию LLaMa 1 с действительно открытым ис
ходным кодом (лицензия которой не допускает коммерческие
варианты использования) под названием OpenLLaMa (https://
github.com/openlm-research/open_llama). Этот фреймворк со
держит пример скрипта для предварительной тренировки мо
дели LLaMa с 7 млрд параметров на блоке TPU v4-512 и скрипт
для обслуживания (сопровождения) предварительно натрени
рованной модели LLaMa на компьютере с GPU или на одной
виртуальной машине TPU v3-8;
Paxml (или Pax) (https://github.com/google/paxml; текущая вер
сия 1.4.0) компании Google – фреймворк машинного обучения
на основе JAX для тренировки крупномасштабных моделей, на
столько больших, что охватывают несколько слоев или модулей
микросхемы ускорителя TPU. Главной целью Pax является эф
фективное масштабирование. Pax обеспечивает продвинутое
и полностью конфигурируемое экспериментирование и рас
параллеливание, показывая при этом лучшие в отрасли флопхарактеристики при использовании моделей. Комплект руко
водств для начала работы с Pax находится здесь: https://github.
com/google/paxml/blob/main/paxml/docs/hands-on-tutorials.md;
MaxText (https://github.com/google/maxtext) компании Google –
простая, высокопроизводительная и масштабируемая LLM,
написанная на чистом Python/JAX и ориентированная на ис
пользование Google Cloud TPU. Модель показывает высокие
флоп-характеристики (model flop utilization – MFU) – от 50 до
70 % по сравнению с 21,3 % для GPT-3 и 46,2 % для PaLM – и обес
печивает масштабирование от одного хоста до весьма крупных
кластеров, оставаясь при этом простой и «свободно оптимизи
руемой» благодаря мощи JAX и компилятора XLA. Для работы
450
Глава 12
Другие члены экосистемы JAX
с MaxText рекомендуется использовать три паттерна: локаль
ное выполнение, экспериментальный запуск в кластере или
порождение процесса в производственной среде, управляемой
Google Compute Engine (GCE). Из этого следует, что проще всего
начать с локальной разработки (Local Development), затем пе
рейти к экспериментам с кластером (Cluster Experimentation)
для некоторой узкоспециализированной доработки и, наконец,
приступить к выполнению долговременных задач в среде GCE.
Как и Pax, MaxText обеспечивает высокую производительность
и масштабируемые реализации LLM в JAX. Фреймворк Pax со
средоточен на предоставлении мощного набора параметров
конфигурации, позволяющих разработчикам изменять модель,
редактируя параметры ее конфигурации. MaxText, напротив,
представляет собой простую конкретную реализацию LLM, сти
мулирующую пользователей расширять ее, создавая ответвле
ния и напрямую редактируя исходный код;
T5X (https://github.com/google-research/t5x) компании Google –
более старая библиотека, наследующая кодовую базу транс
формера T5 (Google). Это модульный, компонуемый, удобный
для исследований фреймворк для высокопроизводительной,
конфигурируемой, самодостаточной тренировки, вычислений
и логического вывода результата языковых моделей преобра
зования многословных последовательностей с разнообразны
ми масштабами. В облачном сервисе Google Vertex AI разме
щены примеры использования T5X (https://cloud.google.com/
vertex-ai), в том числе примеры тренировки и точной настройки
моделей машинного перевода;
Levanter (https://github.com/stanford-crfm/levanter) от группы
Stanford Center for Research on Foundation Models – это фреймворк
для тренировки LLM и других базовых моделей, обязательно тре
бующих высокой разборчивости элементов текста, масштаби
руемости и воспроизводимости. Фреймворк использует допол
няющую его тензорную библиотеку Haliax (https://github.com/
stanford-crfm/haliax), которая в определенном смысле является
альтернативой внутреннему механизму xmap() в JAX. Levanter
обеспечивает масштабирование до крупных моделей и способен
выполнять процесс тренировки на разнообразном аппаратном
обеспечении, включая GPU и TPU, с гарантией масштабирования
до 20 млрд параметров и использования комплекта TPU v3-256.
Кроме того, обеспечивается высокий уровень MFU (https://crfm.
stanford.edu/2023/06/16/levanter-1_0-release.html);
Jaxformer
(https://github.com/salesforce/jaxformer) компании
Salesforce – минималистичная библиотека для тренировки LLM
на TPU в JAX с распараллеливанием по данным и распараллели
ванием модели с помощью pjit();
Экосистема глубокого обучения
451
JaxSeq (https://github.com/Sea-Snell/JAXSeq) – библиотека, соз
данная на основе библиотеки Hugging Face Transformers и под
держивающая процесс тренировки LLM в JAX. В настоящее вре
мя поддерживает модели GPT2, GPT-J, T5 и OPT;
Alpa (https://github.com/alpa-projects/alpa) – система для тре
нировки и обслуживания (сопровождения) крупномасштабных
нейронных сетей. Alpa автоматически распараллеливает по
данным пользовательский код для одного устройства по рас
пределенным кластерам, а также обеспечивает распараллели
вание операторов и конвейера. Примеры включают трениро
вочный / точно настраиваемый трансформер зрения (ViT) для
классификации изображений, точно настраиваемые языковые
модели OPT для экземпляров AWS p3.16xlarge с 8×16 Гб V100
GPU и GPT-2-подобный пример, аналогичный примеру Hugging
Face.
Вся экосистема чрезвычайно динамична и часто изменяется,
поэтому рекомендуется внимательно следить за ее развитием, по
скольку после выхода из печати этой книги могут появиться новые
интересные проекты.
12.1.3 Библиотеки утилит
Экосистема JAX содержит множество библиотек утилит, упрощаю
щих и улучшающих жизнь разработчика, использующего JAX.
Оптимизация:
библиотека Optax (https://github.com/google-deepmind/optax) –
была разработана компанией DeepMind. Мы рассматривали эту
библиотеку в главе 11. Это фреймворк для компоновки новых
оптимизаторов из многократно используемых трансформаций
градиентов, и без него практически невозможно обойтись при
тренировке собственных нейронных сетей;
JAXopt (https://github.com/google/jaxopt) – еще одна библиотека
для аппаратно ускоряемых, пакетируемых и дифференцируе
мых оптимизаторов в JAX;
Optimistix (https://github.com/patrick-kidger/optimistix) – биб
лиотека для нелинейных решателей-вычислителей в JAX
и Equnox;
Lineax (https://github.com/patrick-kidger/lineax) – библиотека
для линейных решателей-вычислителей и для линейной оцен
ки методом наименьших квадратов;
также существуют более специализированные библиотеки
оптимизации, например KFAC-JAX (https://github.com/googledeepmind/kfac-jax), созданные для оптимизации второго по
рядка нейронных сетей и для вычисления масштабируемых
аппроксимаций кривизны.
Глава 12
452
Другие члены экосистемы JAX
Улучшение качества кода:
библиотека Common Loop Utils (CLU) (https://github.com/google/
CommonLoopUtils) компании Google также рассматривалась
в главе 11. Она содержит общую функциональность для созда
ния циклов тренировки машинного обучения, чтобы сделать их
более короткими и удобными для чтения без ущерба гибкости,
необходимой при исследованиях;
Chex (https://github.com/google-deepmind/chex) компании Deep
Mind – библиотека утилит, помогающих писать надежный код
JAX. Chex предоставляет разнообразные утилиты, включая JAXсовместимые средства модульного тестирования, классы дан
ных (dataclasses), логические утверждения о свойствах типов
данных JAX, имитации и фиктивные объекты, а также среды
тестирования на нескольких устройствах;
Jax-verify (https://github.com/google-deepmind/jax_verify) – еще
одна библиотека компании DeepMind, содержащая JAX-реали
зации многих широко используемых методик верификации
нейронных сетей;
библиотека jaxtyping (https://github.com/patrick-kidger/jaxtyp
ing) – предоставляет аннотации типов и средства проверки во
время выполнения формы shape и типа dtype массивов JAX/
NumPy/PyTorch и т. д.
Инструментальные средства разработки:
Orbax (https://github.com/google/orbax) – еще одна библиотека
Google, используемая в главе 11, включает библиотеки конт
рольных точек и сериализации, ориентированные на пользо
вателей JAX. Библиотека поддерживает разнообразные функ
циональные средства, требуемые для различных фреймворков,
включая асинхронные контрольные точки, многочисленные
типы и форматы хранения, экспорт моделей JAX в формат
SavedModel TensorFlow;
инструментальное средство JAX2TF (https://www.tensorfow.org/
guide/jax2tf) компании Google – также предоставляет простой
способ преобразования модели JAX в формат SavedModel Ten
sorFlow и обеспечивает развертывание моделей с помощью
тщательно разработанной экосистемы TensorFlow. Применяя
этот инструмент, можно выполнить логический вывод резуль
тата на сервере, использующем TF Serving, на устройстве, ис
пользующем TFLite, или в веб-среде, применяя TensorFlow.js.
Также можно выполнять точную настройку, продолжив трени
ровку модели, ранее тренированной в JAX, в TensorFlow с соб
ственными имеющимися тренировочными данными и конфи
гурацией. Можно даже выполнить слияние, комбинируя части
моделей, тренированных с применением JAX, с частями моде
лей, тренированных с использованием TensorFlow;
Экосистема глубокого обучения
453
Saxml, или просто Sax (https://github.com/google/saxml), – еще
один проект Google. Это экспериментальная система, обслу
живающая модели JAX, Paxml (см. предыдущий подраздел об
LLM) и PyTorch с обеспечением логического вывода результата.
Ячейка Sax (также называемая кластером Sax) состоит из серве
ра администрирования (admin server) и группы серверов моде
лей (model servers). Сервер администрирования постоянно сле
дит за серверами моделей, распределяя по ним публикуемые
модели для обслуживания (сопровождения), и помогает клиен
там обнаруживать местонахождение серверов моделей, обслу
живающих конкретные опубликованные модели. В настоящее
время существует релиз 1.2.0 этого проекта;
TF2JAX (https://github.com/google-deepmind/tf2jax) – экспери
ментальная библиотека компании DeepMind для преобразова
ния функций/графов TensorFlow в функции JAX, позволяющая
многократно использовать и точно настраивать существующие
модели TensorFlow в кодовых базах JAX;
JAX ONNX Runtime (https://github.com/google/jaxonnxruntime) –
комплект инструментальных средств, обеспечивающий бес
проблемное выполнение моделей ONNX, использующих JAX,
как внутренних компонентов.
Некоторые другие полезные библиотеки:
JMP (https://github.com/google-deepmind/jmp) – библиотека под
держки смешанной точности для JAX. Позволяет объединять
и совместно использовать числа с плавающей точкой полной
и половинной точности для снижения требований к пропуск
ной способности памяти и для улучшения эффективности вы
числений конкретной модели;
Transformer Engine (https://github.com/NVIDIA/TransformerEn
gine) – библиотека компании NVIDIA для ускорения моделей
трансформеров на GPU NVIDIA, включая использование 8-би
товой точности чисел с плавающей точкой (FP8) на Hopper GPU
для обеспечения наилучшей производительности с использо
ванием меньшего объема памяти как при тренировке, так и при
логическом выводе результата;
JAXline (https://github.com/google-deepmind/jaxline) – распре
деленный JAX-фреймворк тренировки и вычислений, предна
значенный для создания ответвлений и охватывающий только
наиболее общие аспекты экспериментальных стандартных за
готовок (шаблонов);
Tree-math (https://github.com/google/tree-math) – упрощает реа
лизацию алгоритмов, работающих с pytree JAX, таких как ите
ративная оптимизация и методы решения уравнений;
Haliax (https://github.com/stanford-crfm/haliax) – библиотека
для создания нейронных сетей с именованными тензорами,
альтернатива механизму xmap();
Глава 12
454
Другие члены экосистемы JAX
einops (https://github.com/arogozhnikov/einops) – гибкая и мощ
ная библиотека операций с тензорами для написания надежно
го и удобного для чтения кода;
Penzai (https://github.com/google-deepmind/penzai) – исследо
вательский комплект инструментальных средств для создания,
редактирования и визуального представления нейронных се
тей. Основное внимание сосредоточено на упрощении разно
образных действий с моделями после завершения их трени
ровки, таким образом, это превосходный вариант выбора для
исследовательских работ, подразумевающих реверс-инжини
ринг, реконструирование (перестройку) модели, отладку, инс
пектирование и исследование внутренних активаций и т. п.
12.2 Модули машинного обучения
В экосистеме JAX широко представлены и другие подразделы ма
шинного обучения.
12.2.1 Обучение с подкреплением
Обучение с подкреплением (reinforcement learning – RL) за послед
нее десятилетие стало весьма популярной темой, поэтому неудиви
тельно, что существует набор библиотек, поддерживающих этот тип
обучения:
RLax (https://github.com/google-deepmind/rlax) компании Deep
Mind – предоставляет полезные конструктивные блоки для
реализации агентов обучения с подкреплением. Компоненты
RLax охватывают широкий спектр алгоритмов и принципов:
TD-обучение1, стратегии вычисления градиентов, адаптивнокритические акторы, MAP, проксимальную стратегию опти
мизации, нелинейные трансформации значений, обобщенные
функции значения состояния и группу методов исследований;
Coax (https://github.com/coax-dev/coax) – модульный пакет под
держки обучения с подкреплением на языке Python компании
Microsoft для решения задач в средах OpenAI Gym (в настоящее
время – Gymnasium) с использованием JAX;
Dopamine (https://github.com/google/dopamine) компании Goog
le – исследовательский фреймворк для быстрого прототипиро
вания алгоритмов обучения с подкреплением. Основная цель –
удовлетворение потребности в небольшой, легко осваиваемой
кодовой базе, в которой пользователи могут свободно экспери
ментировать со своими идеями;
1
TD – temporal difference – временные разности. – Прим. перев.
Модули машинного обучения
455
Acme (https://github.com/google-deepmind/acme) компании Deep
Mind – библиотека конструктивных блоков RL, цель которой –
представление простых, эффективных и удобных для чтения
агентов;
gymnax (https://github.com/RobertTLange/gymnax) – предостав
ляет JAX-ускоренные среды обучения с подкреплением;
Mctx (https://github.com/google-deepmind/mctx) компании Deep
Mind – библиотека с собственной реализацией на чистом JAX
(и с полной поддержкой JIT-компиляции) алгоритмов поиска
в дереве методом Монте-Карло (Monte Carlo tree search – MCTS),
таких как AlphaZero, MuZero и Gumbel MuZero;
Jumanji (https://github.com/instadeepai/jumanji) – разнородный
комплект масштабируемых сред обучения с подкреплением,
написанных на JAX.
12.2.2 Прочие библиотеки машинного обучения
Если необходимо использовать графовую нейронную сеть, то для
этого существует специальное инструментальное средство:
Jraph (https://github.com/google-deepmind/jraph) – библиотека
для работы с графовыми нейронными сетями компании Deep
Mind. Jraph предоставляет структуры данных для графов, набор
утилит для работы с графами и «зоопарк» расширяемых (for
kable) моделей графовых нейронных сетей.
В экосистеме JAX также существуют библиотеки для работы с эво
люционными вычислениями:
EvoJAX (https://github.com/google/evojax) – масштабируемый,
общецелевой, аппаратно ускоряемый комплект инструмен
тальных средств нейроэволюционных вычислений. Предостав
ляет нейроэволюционные алгоритмы для работы с нейрон
ными сетями, выполняющимися в параллельном режиме на
нескольких (многих) TPU/GPU;
Evosax (https://github.com/RobertTLange/evosax) – огромная биб
лиотека стратегий эволюционных вычислений в JAX.
Некоторые библиотеки специально предназначены для вероят
ностного программирования и байесовской оптимизации:
Oryx (https://github.com/jax-ml/oryx) – библиотека для вероят
ностного программирования и глубокого обучения, созданная
на основе JAX;
NumPyro (https://github.com/pyro-ppl/numpyro) – простая биб
лиотека для вероятностного программирования, предоставляю
щая внутренний компонент NumPy для Pyro, библиотеки глубо
кого вероятностного программирования. NumPyro использует
автоматическое дифференцирование и JIT-компиляцию JAX;
Глава 12
456
Другие члены экосистемы JAX
Bayex (https://github.com/alonfnt/bayex) – библиотека поддерж
ки высокопроизводительной байесовской глобальной оптими
зации, использующая гауссовы процессы. Написана полностью
на JAX;
BlackJAX (https://github.com/blackjax-devs/blackjax) – библио
тека байесовского вывода, спроектированная для обеспечения
легкости использования, скорости и модульности. Это библио
тека семплеров для JAX, а не для вероятностного программи
рования.
Кроме того, существуют средства федеративного обучения:
FedJAX (https://github.com/google/fedjax) – библиотека с откры
тым исходным кодом на основе JAX для имитаций федератив
ного обучения, в которой особый акцент сделан на легкость ис
пользования при исследованиях;
Flower (https://fower.ai/) – удобный для пользователей фрейм
ворк федеративного обучения с поддержкой JAX.
Также имеются инструментальные средства для работы с графи
кой и поддержки компьютерного зрения:
PIX (https://github.com/google-deepmind/dm_pix) – библиотека
обработки изображений в JAX. Главная цель – предоставление
функций и инструментальных средств обработки изображений
в JAX, чтобы их можно было оптимизировать и распараллели
вать с помощью jit(), vmap() и pmap();
Scenic
(https://github.com/google-research/scenic) – кодовая
база, сосредоточенная на исследованиях моделей с механиз
мом внимания для компьютерного зрения. Представляет со
бой комплект совместно используемых простых библиотек для
решения задач, часто встречающихся при тренировке крупно
масштабных моделей компьютерного зрения и в нескольких
проектах, содержащих полностью доработанные, ориентиро
ванные на конкретную задачу циклы тренировки и вычисления
с применением этих библиотек;
существует возможность дифференцируемого рендеринга с ис
пользованием JAX. Великолепный пример размещен здесь:
https://google-research.github.io/self-organising-systems/2022/
jax-raycast/;
visu3d
(https://github.com/google-research/visu3d) – уровень
абстракции между Torch/TensorFlow/JAX/NumPy и пользова
тельской программой, предоставляющий стандартные прими
тивы для трехмерной геометрии;
Big Vision (https://github.com/google-research/big_vision) – офи
циальная кодовая база, используемая для разработки Vision
Transformer, SigLIP, MLP-Mixer и т. п.
Модули JAX для других сфер деятельности
457
Другие интересные специализированные библиотеки:
JaxPruner (https://github.com/google-research/jaxpruner) – биб
лиотека с открытым исходным кодом на основе JAX с поддерж
кой средств отсечения (возможных решений) и разреживания
процесса тренировки для исследований в области машинного
обучения;
Rax (https://github.com/google/rax) – библиотека обучения ран
жирования, предоставляющая готовые к практическому при
менению реализации ранжированных потерь и метрик, кото
рые можно использовать в JAX;
OTT-JAX (https://github.com/ott-jax/ott) – библиотека для вы
числения оптимальной передачи с масштабированием и на ак
селераторах;
CoDeX (https://github.com/google/codex) – библиотека содержит
инструментальные средства сжатия данных обучения для JAX;
Foolbox (https://github.com/bethgelab/foolbox) – библиотека Py
thon, позволяющая с легкостью выполнять состязательные ата
ки на модели машинного обучения, такие как нейронные сети
глубокого обучения;
metax (https://github.com/smonsays/metax) – библиотека мета
обучения в JAX, предназначенная для исследований. Включает
разнообразные алгоритмы и архитектуры, которые можно про
извольно комбинировать и с легкостью расширять.
Последняя по списку, но не по важности библиотека для сохране
ния конфиденциальности машинного обучения:
JAX-Privacy (https://github.com/google-deepmind/jax_privacy) –
библиотека компании DeepMind содержит JAX-реализации
алгоритмов сохранения конфиденциальности машинного
обучения.
12.3 Модули JAX для других сфер деятельности
Этот раздел не является исчерпывающим обзором всех модулей JAX.
Его цель – продемонстрировать весьма широкий диапазон приме
нимости JAX и особо выделить тот факт, что область применения JAX
не ограничивается глубоким обучением и машинным обучением.
Модули, предназначенные для дифференцируемой физики:
JAX, M.D. (https://github.com/jax-md/jax-md) компании Google –
фреймворк для ускоряемой дифференцируемой молекуляр
ной динамики. Предоставляет дифференцируемые, аппаратно
ускоряемые модели молекулярной динамики, созданные на
основе JAX;
Глава 12
458
Другие члены экосистемы JAX
Brax (https://github.com/google/brax) – быстрый и полностью
дифференцируемый механизм физики, используемый для ис
следований и разработки в области робототехники, человеческо
го восприятия, материаловедения, обучения с подкреплением
и прочих прикладных областей с интенсивным имитированием.
Также существуют более специализированные модули, предна
значенные для конкретных физических дисциплин:
jax-cosmo (https://github.com/DifferentiableUniverseInitiative/jax_
cosmo) – библиотека поддержки дифференцируемой космоло
гии;
j-Wave (https://github.com/ucl-bug/jwave) – библиотека имита
ций для приложений акустики;
JAX-Fluids (https://github.com/tumaer/JAXFLUIDS) – пакет для
дифференцируемой динамики жидкостей и газов;
Veros (https://github.com/team-ocean/veros) – универсальный
имитатор поведения океана. Основная цель – предоставление
«швейцарского армейского ножа» для моделирования пове
дения океана. Это полнофункциональная модель по полным
уравнениям океана с поддержкой всего диапазона – от идеа
лизированных учебных экспериментальных моделей до реа
листичных глобальных имитаций поведений океана с высоким
разрешением;
JAXChem (https://github.com/deepchem/jaxchem) – библиотека
глубокого обучения на основе JAX для комплексного универ
сального химического моделирования;
OptimiSM (https://github.com/sandialabs/optimism) – библиотека
для формулирования и решения задач механики деформируе
мого твердого тела (теории упругости) с использованием мето
да конечных элементов;
qujax (https://github.com/CQCL/qujax) – простая, быстрая и гиб
кая библиотека Python на основе JAX для классической имита
ции квантовых схем;
PennyLane (https://github.com/PennyLaneAI/pennylane) – кроссплатформенная библиотека Python для дифференцируемого
программирования квантовых компьютеров;
NetKet (https://github.com/netket/netket) – проект с открытым
исходным кодом, предоставляющий новейшие методы для ис
следования квантовых систем многих тел (частиц) с примене
нием нейронных сетей искусственного интеллекта и методик
машинного обучения;
DeepXDE (https://github.com/lululxvi/deepxde) – библиотека для
научного машинного обучения и основанного на физических
принципах обучения, например с использованием нейрон
ных сетей, основанных на физике (physics-informed neural net
works – PINN);
Резюме
459
Diffrax (https://github.com/patrick-kidger/diffrax) – библиотека
на основе JAX, предоставляющая инструментальные средства
численного решения дифференциальных уравнений.
Здесь следует остановиться. Любой подобный список никогда не
является исчерпывающим и будет быстро устаревать. Поэтому ре
комендуется внимательно следить за экосистемой JAX, заглядывать
в репозитории GitHub и регулярно просматривать списки модулей,
один из которых размещен здесь: https://github.com/n2cholas/awe
some-jax.
Резюме
В экосистеме JAX имеется множество высокоуровневых библио
тек поддержки нейронных сетей, включая Flax, Equinox, Keras 3.0
и другие интересные варианты.
Экосистема JAX содержит огромную коллекцию библиотек, спе
циально ориентированных на тренировку и логический вывод
больших языковых моделей.
Многие библиотеки весьма полезны для различных аспектов глу
бокого обучения: организации циклов обучения, выполнения
трансформаций градиентов и оптимизации, написания исходно
го кода, создания надежных нейронных сетей и т. д.
Отдельный набор модулей помогает в развертывании и организа
ции логического вывода результата с использованием тщательно
проработанной экосистемы TensorFlow.
Многие библиотеки JAX созданы для специализированных обла
стей машинного обучения, таких как обучение с подкреплением,
компьютерное зрение, федеративное обучение, вероятностное
программирование, эволюционные вычисления и т. д.
Область применения JAX не ограничена глубоким обучением, и вы
можете обнаружить многочисленные инструментальные средства
для использования JAX в физике, химии, космологии, квантовых
вычислениях и многих других сферах деятельности.
Экосистема JAX активно развивается, поэтому прежде чем начать
писать что-либо с нуля, рекомендуется внимательно рассмотреть
огромную и динамичную экосистему JAX. Вероятнее всего, вы
найдете то, что поможет вам.
Приложение A
Установка JAX
JAX публикуется в виде двух отдельных пакетов Python:
jax – пакет с чистым кодом Python;
jaxlib – пакет, в основном написанный на C++ и содержащий
такие библиотеки, как XLA, части LLVM, используемые XLA,
инфраструктуру MLIR с привязками MHLO на Python и спе
циализированные для JAX библиотеки C++ для быстрой обра
ботки JIT и pytree.
Процесс установки JAX будет различным в зависимости от целе
вой архитектуры: CPU, GPU или TPU.
A.1
Установка JAX для CPU
JAX предназначен для высокопроизводительных вычислений и лучше
всего проявляет себя на TPU или GPU, тем не менее за счет примене
ния компилятора XLA можно достичь заметного ускорения даже на
CPU. Кроме того, может потребоваться установка варианта CPU для ло
кальной разработки. Самый простой способ установки JAX для CPU –
использование установщика пакетов pip для рабочей среды Python.
Чтобы установить версию JAX только для CPU, выполните следу
ющие команды:
pip install --upgrade pip
pip install --upgrade "jax[cpu]"
Текущий релиз jaxlib поддерживает следующие платформы и ар
хитектуры:
Linux x86_64;
Установка JAX
461
Mac x86_64;
Mac ARM;
Windows x86_64, собственный релиз или использующий WSL2 –
Windows Subsystem for Linux.
В Windows, возможно, также потребуется установка Microsoft Vi
sual Studio 2019 Redistributable, если этот пакет не установлен на ва
шем компьютере. Более подробную информацию см. в официальной
документации: https://docs.jax.dev/en/latest/installation.html#cpu.
Для других архитектур, таких как Linux aarch64, требуется уста
новка JAX из исходных кодов. Команда pip install может успешно
установить пакет jax, но библиотека jaxlib установлена не будет,
и попытка запуска JAX приведет к ошибке.
A.2
Установка JAX для GPU
Существует несколько способов установки и запуска JAX для GPU
(здесь подразумевается использование GPU компании NVIDIA):
применение CUDA и CuDNN, устанавливаемых из pip wheels
(это самый простой способ, но пакеты wheels доступны только
для Linux x86_64);
использование автоматически устанавливаемых пакетов CUDA/
CuDNN;
использование контейнера Docker.
Методы установки с помощью pip могут не работать в Windows.
На время написания этой книги экспериментальная поддержка су
ществовала только для варианта Windows WSL2 x86_64.
Для GPU AMD предоставляется экспериментальная поддержка
Linux x86_64, требующая сборки JAX из исходных кодов. Более под
робно об этом способе установки см. https://docs.jax.dev/en/latest/
developer.html#additional-notes-for-building-a-rocm-jaxlib-for-amdgpus. Также существует экспериментальная поддержка для Apple
GPU, см. https://docs.jax.dev/en/latest/installation.html#pip-installa
tion-apple-gpus. Информацию о любых изменениях и обновлениях
можно найти в официальной документации: https://docs.jax.dev/en/
latest/installation.html#nvidia-gpu.
A.2.1
Установка методом pip с поддержкой CUDA
JAX поддерживает GPU компании NVIDIA с архитектурой Maxwell
или более новой (с вычислительной мощностью 5.2 или более высо
кой). Следует отметить, что JAX больше не поддерживает GPU серии
Kepler, поскольку компания NVIDIA прекратила поддержку GPU Ke
pler в своем программном обеспечении.
Приложение А
462
Вы можете проверить вычислительную мощность своего GPU
здесь: https://developer.nvidia.com/cuda-gpus#compute. Более под
робную информацию о вычислительных мощностях GPU можно
получить здесь: https://docs.nvidia.com/cuda/cuda-c-programmingguide/index.html#compute-capabilities.
В первую очередь необходимо установить драйвер NVIDIA. Реко
мендуется устанавливать самые новые доступные версии драйверов
NVIDIA, но номер версии обязательно должен быть ≥ 525.60.13 для
CUDA 12 в Linux.
Затем выполняются следующие команды установки CUDA и JAX:
pip install --upgrade pip
# Установка CUDA 12.
# Замечание: пакеты wheels доступны только в Linux.
pip install -- upgrade "jax[cuda12]"
Вот и все – это самый простой способ.
A.2.2
Установка методом pip автоматически
устанавливаемых пакетов CUDA/CuDNN
Возможно, у вас уже установлены некоторые версии CUDA/CuDNN.
Проверьте версию CUDA с помощью следующей команды:
nvcc --version
Если CUDA/CuDNN не установлены, то выполните следующие
шаги: сначала установите драйвер NVIDIA. Рекомендуется установ
ка самого нового драйвера, доступного на сайте NVIDIA, но номер
версии драйвера обязательно должен быть ≥ 525.60.13 для CUDA 12
в Linux. Затем установите CUDA (https://developer.nvidia.com/cudadownloads) и CuDNN (https://developer.nvidia.com/CUDNN).
ПРИМЕЧАНИЕ Установленные версии CUDA и драйвера
NVIDIA должны быть достаточно новыми для поддержки ис
пользуемого GPU.
Далее устанавливается JAX. В настоящее время JAX поддерживает
один вариант пакета wheels CUDA:
созданный с использованием CUDA 12.3, CUDNN 9.0, NCCL 2.19;
совместимый с CUDA >= 12.1, CUDNN >= 9.0, <10.0, NCCL >= 2.18.
Можно воспользоваться пакетом wheel JAX с локальной установ
кой CUDA/CuDNN, если главные части (major) номера версии уста
новленных CUDA и CuDNN совпадают, а минорная (minor) часть вер
сии является настолько новой, насколько предполагает версия JAX.
Установка JAX
463
Для установки пакета wheels JAX выполните следующие команды:
pip install --upgrade pip
# Установка пакета wheel, совместимого с CUDA 12 и cuDNN 8.9
# или более нового.
# Замечание: пакеты wheels доступны только в Linux.
pip install --upgrade "jax[cuda12_local]"
Если при установке возникают ошибки или какие-либо проблемы,
то рекомендуется обратиться к документации по установке JAX.
A.2.3
Использование контейнеров Docker
Компания NVIDIA предоставляет комплект инструментальных
средств JAX Toolbox (https://github.com/NVIDIA/JAX-Toolbox) с кон
тейнерами Docker, содержащими JAX и некоторые библиотеки, такие
как T5X, Paxml и Transformer Engine.
A.3
Установка JAX для TPU
JAX предоставляет предварительно созданный пакет wheels для
Google Cloud TPU. Для установки JAX вместе с соответствующими
версиями jaxlib и libtpu необходимо выполнить следующую ко
манду в используемой облачной виртуальной машине TPU VM:
pip install jax[tpu] -f \
https://storage.googleapis.com/jax-releases/libtpu_releases.html
В приложении C описаны процедуры настройки Google Cloud
TPU, запуска экземпляра Cloud TPU и установления его соединения
с Google Colab.
ПРИМЕЧАНИЕ Среда времени выполнения Colab TPU
в Google Colab отличается от среды Google Cloud TPU. Colab
TPU в Google Colab предоставляет меньшую степень управле
ния и больше не поддерживается фреймворком JAX, начиная
с версии 0.4.
При использовании более ранних версий JAX и Colab TPU необхо
димо выполнить содержимое приведенной ниже ячейки перед им
портированием JAX:
import jax.tools.colab_tpu
jax.tools.colab_tpu.setup_tpu()
Приложение B
Использование
Google Colab
Google Colaboratory (https://colab.google/), или более кратко Colab, –
это сервис блокнота Jupyter, размещенный на специально выде
ленном хосте, не требующий специальной установки и настройки
и предоставляющий свободный доступ к вычислительным ресур
сам, в том числе к GPU и TPU. Colab особенно хорошо подходит для
решения задач машинного обучения, обработки данных, а также
для образовательных целей. Использование Colab бесплатно, хотя
существуют платные планы с увеличенным количеством вычисли
тельных элементов, более быстрыми GPU, расширенным объемом
памяти и т. п. Огромный набор специально организованных, адми
нистрируемых блокнотов доступен здесь: https://colab.google/note
books/.
Блокнот (notebook) – это список ячеек. Ячейки (cells) содержат либо
описательный текст, либо выполняемый код и его вывод (результат
работы). Вы можете выполнять содержимое ячеек в интерактивном
режиме и наблюдать соответствующие выводимые результаты (см.
рис. B.1). Основные функциональные характеристики и средства Co
lab описаны в специальном блокноте: https://colab.research.google.
com/notebooks/basic_features_overview.ipynb.
Одним из факторов, особенно важных для работы с этой книгой,
является возможность переключения среды времени выполнения
между CPU, GPU и TPU. Это можно делать с помощью пункта меню
Runtime (см. рис. B.2).
Использование Google Colab
465
Рис. B.1 Блокнот Colab представляет собой набор ячеек с кодом и текстом
Рис. B.2 Меню Runtime позволяет изменять тип среды времени выполнения
В зависимости от подписки – бесплатной или оплачиваемой – мо
жет быть предоставлен доступ к различным специальным опциям.
В моей бесплатной подписке Colab на момент написания книги до
ступными являлись опции, показанные на рис. B.3.
466
Приложение В
Рис. B.3 Доступные типы среды времени выполнения. Конкретные варианты
могут быть различными в зависимости от вашей подписки (бесплатной или
оплачиваемой)
Блокнот Colab очень удобен для изучения JAX и создания прототи
пов конкретных решений. Доступен тип среды времени выполнения
Colab TPU, который отличается от Google Cloud TPU. Colab TPU пре
доставляет меньшую степень управления и не поддерживается JAX
начиная с версии 0.4. В приложении C описана процедура настройки
Colab TPU и подключение к блокноту Colab.
Другим вариантом использования JAX в облачных интерактив
ных блокнотах с поддержкой TPU являются блокноты Kaggle (https://
www.kaggle.com/docs/notebooks). Использовать TPU можно до 20 ча
сов в неделю и до 9 часов в одном сеансе. Более подробную инфор
мацию об использовании TPU в Kaggle можно найти здесь: https://
www.kaggle.com/docs/tpu.
Приложение C
Использование
Google Cloud TPU
C.1
Настройка проекта Cloud TPU
Для работы с тензорными процессорами (TPU) Google необходимо
настроить аккаунт Google Cloud и подготовить соответствующий
комплект инструментальных средств.
ПРЕДУПРЕЖ ДЕНИЕ Использование Cloud TPU подразуме
вает некоторые расходы денежных средств, поскольку креди
ты Google Cloud Free Tier не включают оплату стоимости GPU
и TPU. Текущие расценки использования TPU варьируются
в диапазоне от 1,20 до 4,20 долл. за час работы с микросхемой
и зависят от региона. Цены на текущий день можно узнать
здесь: https://cloud.google.com/tpu/pricing. Также можно по
лучить кредиты на использование Google Cloud через раз
личные инициативные предложения для стартапов (https://
cloud.google.com/startup), для образования (https://cloud.
google.com/edu/faculty?hl=ru) и для научных исследований
(https://cloud.google.com/edu/researchers и https://sites.re
search.google/trc/about/).
Ниже описаны последовательные шаги по настройке среды ис
пользования Google Cloud TPU.
1 Настройка
аккаунта Google Cloud и создание нового проекта.
Этот процесс описан здесь: https://cloud.google.com/tpu/docs/
setup-gcp-account.
Приложение C
468
2 Установка
утилиты интерфейса командной строки (CLI) gcloud.
В консоли Google Cloud можно многое сделать вручную, и gcloud
упрощает работу. Инструкции по установке этой утилиты см.
здесь: https://cloud.google.com/sdk/docs/install.
3 Разрешение Cloud TPU API использования gcloud или консо
ли Google Cloud. Для этого необходимо выполнить следующую
команду после вывода промпта в интерфейсе командной строки:
gcloud services enable tpu.googleapis.com
4 Затем выполняется команда для создания идентификации сер
виса:
gcloud beta services identity create --service tpu.googleapis.com
После этого вы готовы к созданию облачных машин с TPU.
C.2
Запуск экземпляра Cloud TPU
и установление соединения с Google Colab
Теперь у вас установлена и настроена утилита интерфейса команд
ной строки gcloud, и вы готовы к запуску и установлению соедине
ния с TPU-машиной.
C.2.1
Установка, запуск и удаление узла Cloud TPU
Для запуска TPU-машины необходимо выбрать зону (в приведенном
ниже примере указывается зона us-central-b), тип акселератора (здесь:
TPU v2-8) и имя для нового создаваемого узла (здесь: node-jax):
$gcloud compute tpus tpu-vm create node-jax --zone \
us-central1-b --accelerator-type v2-8 --version tpu-vm-base
Доступность TPU может отличаться для разных зон, поэтому ре
комендуется попробовать указывать другие зоны и типы акселера
торов.
В любой момент можно посмотреть список работающих TPU в за
данной зоне, используя следующую команду:
$gcloud compute tpus tpu-vm list --zone us-central1-b
Для удаления узла (не забывайте выполнять эту операцию, иначе
величина оплаты может оказаться неожиданно большой) использу
ется следующая команда:
$gcloud compute tpus tpu-vm delete node-jax --zone us-central1-b
Использование Google Cloud TPU
C.2.2
469
Подготовка узла Cloud TPU
После успешного создания узла вы можете зарегистрироваться на
нем (войти), используя SSH. Для использования созданного узла
как локальной среды выполнения для конкретного блокнота Google
Colab или Jupyter необходимо создать SSH-туннель для перенаправ
ления запросов на заданный локальный порт (в приведенном ниже
примере: 8888) на удаленном компьютере. Это делается с помощью
следующей команды:
$gcloud compute tpus tpu-vm ssh --zone us-central1-b \
node-jax -- -L 8888:localhost:8888
Если вы используете Windows PowerShell, будьте особенно внима
тельны: двойной дефис обязательно должен быть заключен в оди
ночные кавычки:
$gcloud compute tpus tpu-vm ssh --zone us-central1-b \
node-jax '--' -L 8888:localhost:8888
Теперь вы находитесь на пустой машине, поэтому требуется уста
новить все необходимые инструментальные средства для работы.
Для установки JAX и начала экспериментирования с этим фреймвор
ком в консоли, не теряя времени, выполните следующую команду:
$pip install jax[tpu] -f \
https://storage.googleapis.com/jax-releases/libtpu_releases.html
При необходимости использования блокнотов Google Colab или
Jupyter нужно сделать еще кое-что. Во-первых, установить требуе
мые модули:
$pip install -U jinja2
$pip install notebook
По умолчанию Jupyter устанавливается локально и не изменяет
системную переменную PATH. Возможно, потребуется ее обновление
для включения каталога, в который установлен Jupyter (во время
установки выводятся предупреждающие сообщения об этом). Заме
ните показанный ниже путь на путь к своему домашнему каталогу:
$export PATH=$PATH:/home/grigo/.local/bin
Установите расширение для Jupyter:
$pip install jupyter_http_over_ws
Запустите сервер jupyter и разрешите установление соединений
с Google Colab:
470
Приложение C
$jupyter notebook \
--NotebookApp.allow_origin='https://colab.research.google.com' \
--port=8888 \
--NotebookApp.port_retries=0
После запуска сервера jupyter он выводит ссылку для доступа
к блокноту. В моем варианте такая ссылка выглядит следующим
образом: http://localhost:8888/?token=5dd75b993902ba1e9710471a5b
0a6c2b887bc0c35841b1c7. Значение токена в URL будет изменяться
в различных сеансах, поэтому необходимо точно скопировать этот
URL.
Этот процесс также описан здесь: https://research.google.com/co
laboratory/local-runtimes.html.
C.2.3
Установление соединения с узлом Cloud TPU
из блокнота Colab
Теперь можно установить соединение Colab с новой созданной сре
дой времени выполнения. Для этого в Google Colab перейдите к пунк
ту Reconnect -> Connect to a Local Runtime (Установление соедине
ния с локальной средой времени выполнения) (см. рис. C.1).
Рис. C.1 Установление соединения блокнота Colab с локальной средой
времени выполнения
Вставьте скопированную ссылку в поле Backend URL и щелкни
те по кнопке Connect (Соединение) (см. рис. C.2). Если соединение
с блокнотом не установлено из-за тайм-аута, то щелкните по кнопке
Reconnect (Повторное соединение).
Использование Google Cloud TPU
471
Рис. C.2 Ввод URL внутреннего компонента для локальной среды времени
выполнения
C.3
Ресурсы
Более подробно об управлении TPU:
https://cloud.google.com/tpu/docs/managing-tpus-tpu-vm.
Краткое руководство по выполнению вычислений на виртуальной
машине Cloud TPU VM:
https://cloud.google.com/tpu/docs/run-calculation-jax.
Описание архитектур Cloud TPU VM:
https://cloud.google.com/tpu/docs/system-architecture-tpu-vm#tpuarch.
Использование вырезок TPU Pod:
https://cloud.google.com/tpu/docs/jax-pods.
Проверка доступности TPU в различных зонах:
https://cloud.google.com/tpu/docs/regions-zones.
Создание экземпляра виртуальной машины (VM) глубокого обуче
ния с использованием интерфейса командной строки gcloud:
h ttps://cloud.google.com/deep-learning-vm/docs/create-vm-in
stance-gcloud.
Создание экземпляра виртуальной машины (VM) глубокого обуче
ния с использованием консоли Google Cloud:
h ttps://cloud.google.com/deep-learning-vm/docs/create-vm-in
stance-console.
Образ виртуальной машины (VM) глубокого обучения:
https://cloud.google.com/deep-learning-vm.
Цены на использование Cloud TPU:
https://cloud.google.com/tpu/pricing.
Приложение D
Экспериментальные
средства
распараллеливания
В этом приложении объединены описания двух (с половиной) экс
периментальных методик распараллеливания: xmap() (и ее полови
на shmap()) и pjit(). xmap() – более старая методика, которая была
удалена в версии JAX 0.4.31 (29 июля 2024 г.). Тем не менее она может
оставаться интересной для тех, кому необходимо разбираться в дав
но существующем коде, или для желающих узнать как можно под
робнее о развитии механизма распараллеливания в JAX.
Трансформация xmap() помогает распараллеливать функции
проще, чем pmap(), и с меньшим объемом кода, заменяя вложен
ные вызовы pmap() и vmap(), а также без изменения формы тензо
ров вручную. xmap() тоже представляет модель программирования
с именованными осями, помогающую писать код, более защищен
ный от потенциальных ошибок.
В некоторый момент активная разработка xmap() прекратилась
и предпочтение было отдано pjit() (это следующая тема нашего
обсуждения). Несмотря на статус «устаревшая и не рекомендуе
мая к применению», вполне логично описать эту методику здесь,
поскольку xmap() предоставляет в высшей степени естественный
способ обобщения pmap() и vmap(). Если вам не требуется такое
функциональное средство и вы не обязаны поддерживать давно су
ществующую кодовую базу, то можете пропустить эту часть.
Экспериментальные средства распараллеливания
473
Один из возможных вариантов замены xmap() обнаруживается
в экосистеме JAX – библиотека Haliax для создания нейронных сетей
с именованными тензорами. Другая альтернатива находится в ядре
JAX – это трансформация shmap(), в настоящее время имеющая ста
тус JEP (JAX Enhancement Proposal) и представляющая собой замену
xmap().
Иногда пользовательские функции (или нейронные сети) могут
оказываться настолько большими, что их размещение невозмож
но на одном GPU/TPU, поэтому требуется организация вычислений
в кластере. Это частый случай при тренировке и логическом выво
де результата больших языковых моделей (LLM). Современные LLM,
такие как GPT-3 и 4, более крупные версии LLaMa 2, Falcon и многие
другие, требуют применения систем с несколькими (многими) GPU.
JAX позволяет не только распределять (или разделять, или сегмен
тировать) процесс обработки данных по различным компьютерам
(распараллеливание по данным), но также разделять большой объ
ем вычислений на части, выполняемые различными компьютерами
(так называемое распараллеливание модели).
Одним из способов сегментирования большого объема вычисле
ний является трансформация pjit(), в некоторый момент ставшая
более широко применяемой, чем xmap(). Более старые LLM, такие как
GPT-J-6B, использовали xmap() (см. вариант использования здесь:
https://arankomatsuzaki.wordpress.com/2021/06/04/gpt-j/), а новые LLM
в основном предпочитают pjit() (см. вариант использования Co
here: https://cloud.google.com/blog/products/ai-machine-learning/acce
lerating-language-model-training-with-cohere-and-google-cloud-tpus).
Трансформации pjit() и jit() (это темы главы 5) были объеди
нены в единый универсальный интерфейс, поэтому настоятельно
рекомендуется использовать только jit(). Кроме того, механизм
сегментирования тензоров для распределенных массивов (тема
главы 8) предоставляет современный способ компиляции и выпол
нения функций JAX в средах со многими хостами или ядрами про
цессоров.
Эту часть можно пропустить, если не требуется поддержка давно
существующего кода, использующего pjit(), и вы не планируете на
прямую применять jit() со спецификациями сегментирования.
D.1
Использование xmap() и программирование
с именованными осями
Трансформация xmap() представляет собой интересное эксперимен
тальное функциональное средство, которое может упростить про
граммы несколькими способами (см. рис. D.1). Во-первых, xmap()
Приложение D
474
предоставляет модель программирования с именованными осями,
позволяющую перейти от индексов осей тензоров к именам осей.
Такой подход обеспечивает более надежную защиту кода от потен
циальных ошибок и упрощает его понимание и редактирование.
xmap()
Программирование
с именованными осями
pmap()+vmap()
Рис. D.1 Функциональные возможности xmap()
Во-вторых, xmap() позволяет заменить pmap() и vmap() одной
функцией, способной выполнять обе заменяемые операции, сле
довательно, избежать некоторых технических проблем, связанных
с тем, как должна быть выполнена требуемая задача, а не что долж
но быть сделано: исключаются вложенные вызовы pmap()/vmap()
и процедуры изменения формы данных перед вызовом pmap(). При
использовании xmap() пользователю не нужно беспокоиться о коли
честве доступных устройств, от него не требуется изменение фор
мы данных, добавление специального измерения для отображения
и последующее его удаление после завершения вычислений. Все
это происходит незаметно для пользователя и без его участия всего
лишь в одном вызове xmap().
Начнем с описания парадигмы программирования с именован
ными осями. Затем перейдем к части распараллеливания.
Ранее мы рассматривали примеры создания нейронных сетей
с использованием операций с тензорами, в которых схема разме
щения тензора в памяти имела важное значение, и функции обычно
были написаны с учетом того, что конкретное измерение тензора
содержит именно ту часть данных, которая ожидается. Например,
вызовы predict() и loss() в листингах 7.25 и 7.26 работали коррект
но, если входные тензоры имели ожидаемые формы, т. е. batch_size
и т. п., для образов, содержащих тензоры. Код был основан на кон
кретных позициях различных осей тензоров. Но при изменении
схемы размещения тензора в вычислениях возникали ошибки из-за
несоответствия измерений или (что еще хуже) вычисления продол
жались без видимых ошибок, однако по смыслу вычислялись совер
шенно неправильные результаты.
Подобные проблемы с измерениями тензоров обычно очень
трудно отлаживать, поэтому для глубокого обучения предлагаются
различные решения. Один из полезных способов предотвращения
ошибок такого рода – модель тензора с именами, предложенная
Экспериментальные средства распараллеливания
475
Александром Рашем (Alexander Rush) из группы Harvard NLP в ко
ротком сообщении «Tensor Considered Harmful» (https://nlp.seas.
harvard.edu/NamedTensor) и более подробно описанная в документе
«Named Tensor Notation» (https://arxiv.org/abs/2102.13196).
Основная идея заключается в явном присваивании имен измере
ниям тензора, например измерениям пакетов, характеристик, высо
ты, ширины или каналов. Если для большинства операций с тензора
ми, требующих указания измерения, передаются имена измерений,
то можно избежать необходимости тщательного отслеживания из
мерений по позиции.
Кроме того, такой подход обеспечивает дополнительную безопас
ность, выполняя автоматическую проверку того, что все API исполь
зуются корректно во время выполнения. Перегруппировка измере
ний по именам вместо позиций становится проще для написания
и понимания, следовательно, вероятность совершения ошибок сни
жается.
Существует реализация тензоров с именами PyTorch (https://docs.
pytorch.org/tutorials/intermediate/named_tensor_tutorial.html). Mesh
TensorFlow (https://github.com/tensorfow/mesh) использует имено
ванные измерения тензоров и сеток. В JAX имеется немного изме
ненная экспериментальная реализация этой идеи, описанная в ру
ководстве по xmap() (https://github.com/jax-ml/jax/blob/jax-v0.4.30/
docs/notebooks/xmap_tutorial.md).
ПРИМЕЧАНИЕ xmap() представлял собой эксперименталь
ный API и был удален в версии JAX 0.4.31 (29 июля 2024 г.).
D.1.1
Работа с именованными осями
Основная идея заключается во введении именованных осей в до
полнение к позиционным осям тензора, чтобы тензор имел оба
атрибута .dtype и .shape, описывающих его позиционные оси как
обычный массив NumPy, и дополнительный атрибут .named_shape,
описывающий его именованные оси.
Именованные оси нельзя добавить непосредственно в независи
мый отдельный тензор. В настоящее время единственным способом
создания именованных осей является использование вызова xmap().
Это будет продемонстрировано немного позже.
Атрибут .dtype хранит тип элементов (например, np.float32),
в атрибуте .shape размещен кортеж целых чисел, определяющих
форму массива (например, (3, 5) для массива 3×5), а атрибут .named_
shape представляет собой словарь (dict), отображающий имена осей
в целочисленные размеры (например, {'batch': 32, 'channel': 10}
для двух дополнительных осей массива с размерами 32 и 10 соот
ветственно).
Приложение D
476
Имена осей – это произвольно хешируемые объекты, для кото
рых обычно применяют строки, но вы можете использовать и дру
гие типы. Порядок именованных осей не имеет значения, и они не
упорядочены, т. е. {'batch': 32, 'channel': 10} и {'channel': 10,
'batch': 32} считаются равнозначными определениями.
Для поддержки обратной совместимости смысл атрибутов .ndim
и .size не изменился, и они всегда равны len(shape) и произведению
элементов .shape соответственно. Но действительный ранг массива
с непустыми именованными осями равен len(shape)+len(named_
shape), так как реальное число элементов, хранящихся в таком
массиве, равно произведению размеров всех измерений, как пози
ционных, так и именованных. Например, тензор, содержащий изо
бражения с формой (200, 200) и именованной формой {'batch': 32,
'channel': 3}, имеет действительный ранг 4.
Добавление и удаление именованных осей
В настоящее время единственным способом добавления именован
ных осей является использование xmap(), поскольку в JAX все вы
сокоуровневые операции работают в модели NumPy исключительно
с позиционными осями. Это положение можно изменить в некото
рый момент.
Программирование с именованными осями и параметр
axis_name в vmap()/pmap()
Как вы помните, функции vmap() и pmap() имеют отдельный параметр
axis_name. В некотором смысле этот параметр можно интерпретировать
как чрезвычайно ограниченную форму программирования с именованными осями, применяемую только к конкретной отображаемой оси, и вы
можете использовать эту форму лишь внутри коллективных операций.
Трансформацию xmap() можно воспринимать как обертку или
адаптер, который принимает стандартные массивы (тензоры) с по
зиционными осями, присваивает некоторым осям имена (в соответ
ствии с содержимым параметра in_axes), вызывает обертывающую
функцию и выполняет преобразование именованных осей обратно
в позиционные (используя параметр out_axes).
Отображение осей можно определить двумя способами:
используя словари, отображающие позиционные оси в имено
ванные. Например, {0: 'batch', 3: 'channel'} отображает ось
в позиции 0 в имя 'batch', а ось в позиции 3 – в имя 'channel'
(будьте особенно внимательны: эта запись отличается от записи
атрибута .named_shape, описанной выше, в которой числа соответ
ствуют размерам осей, а здесь числа обозначают позиции осей);
Экспериментальные средства распараллеливания
477
используя списки имен осей, завершающиеся специальным
объектом Python Ellipsis (... или Ellipsis; https://docs.python.
org/3/library/constants.html#Ellipsis) для отображения префикса
позиционных измерений в заданные имена. Например, ['device', 'batch', ...] для варианта, в котором отображаются
первые два позиционных измерения тензора в именованные
измерения «device» и «batch», а все прочие измерения остаются
позиционными. Здесь порядок важен.
Структура in_axes должна соответствовать сигнатуре аргументов
функции, помещенной в обертку. Такое же требование предъявляет
ся и к параметру out_axes, отвечающему за возвращаемое значение
функции.
Все позиционные оси, перечисленные в параметре in_axes, преоб
разуются в именованные. Все именованные оси, указанные в пара
метре out_axes, вставляются в позиции, заданные в этом параметре.
Эти именованные оси фактически удаляются из входных тензо
ров перед вызовом обертываемой функции; функция отображается
по ним, а затем отображаемые оси снова вставляются в заданные
позиции. Результат получается таким, как если бы применили vmap()
к каждой именованной оси (но это намного больше, чем простое
применение vmap(); это больше похоже на способ интерполяции
между стилем скрытого выполнения vmap() и pmap(), как вы сами
увидите немного ниже).
Вернемся к примеру из листинга 7.7, в котором вычислялось ска
лярное произведение двух больших массивов, организованных так,
что измерение пакетов не являлось первым. Ниже воспроизведена
схема из главы 7, наглядно представляющая все происходящее во
время обработки (см. рис. D.2).
Как вы помните, мы изменяем форму тензоров, чтобы получить
отдельное измерение с размером, равным количеству доступных
устройств, вызываем pmap() для распределения вычислений по это
му новому «техническому» измерению, внутри используем vmap()
для обработки пакета элементов на каждом устройстве, а затем уда
ляем это дополнительное измерение после завершения вычислений.
Получается сложный код: во-первых, выполняется композиция
двух функций pmap() и vmap(). Во-вторых, и это самое важное, в обо
их вызовах pmap() и vmap() мы обязательно должны отслеживать
индексы осей массива, используя параметр in_axes. Это достаточ
но сложное дело, поскольку невозможно легко и просто адаптиро
вать вашу функцию к новой схеме размещения тензора: непремен
но нужно внимательно повторно вычислять, с какими индексами
должна работать каждая функция, и при этом постоянно помнить
о том, что каждая функция видит собственные формы. При увеличе
нии вложенных вызовов все становится еще более сложным, и в та
ком коде, основанном на индексах тензоров, возрастает вероятность
Приложение D
478
возникновения ошибок. Именно эту сложность пытаются устранить
авторы «Tensor Considered Harmful».
Рис. D.2 Схема обработки данных для примера с большим транспонированным массивом
(воспроизводится рис. 7.1)
В коде из листинга D.1 требуется восемь устройств для распарал
леливания вычислений, поэтому необходимо создать виртуальную
машину Cloud TPU и использовать ее как локальную среду време
ни выполнения в Colab (см. приложение C или пример в подразде
ле 3.2.5, если требуется помощь при создании виртуальной машины)
или создать восемь виртуальных устройств CPU, как описано в гла
ве 7. Еще один вариант: если у вас есть доступ к системе с несколь
кими GPU, то можно адаптировать приведенный ниже код для такой
системы.
Листинг D.1 Воспроизведение кода из листинга 7.17
vs = random.normal(rng_key, shape=(20_000_000,3))
v1s = vs[:10_000_000,:].T
v2s = vs[10_000_000:,:].T
v1s.shape, v2s.shape
>>> ((3, 10000000), (3, 10000000))
v1sp = v1s.reshape((v1s.shape[0], 8, v1s.shape[1]//8))
❶
❶
❷
❸
479
Экспериментальные средства распараллеливания
v2sp = v2s.reshape((v2s.shape[0], 8, v2s.shape[1]//8))
v1sp.shape, v2sp.shape
❸
>>> ((3, 8, 1250000), (3, 8, 1250000))
❹
dot_parallel = jax.pmap(
jax.vmap(dot, in_axes=(1,1)),
in_axes=(1,1)
)
❺
❻
x_pmap = dot_parallel(v1sp,v2sp)
x_pmap.shape
>>> (8, 1250000)
x_pmap = x_pmap.reshape((x_pmap.shape[0]*x_pmap.shape[1]))
x_pmap.shape
>>> (10000000,)
❼
❽
❶ Теперь используются транспонированные версии исходных массивов.
❷ Первое измерение содержит компоненты вектора. Второе измерение содержит
векторы.
❸ Разделение второго измерения, содержащего векторы, на два новых измерения:
группы и векторы.
❹ Мы получили восемь групп векторов.
❺ Оповещение vmap о необходимости использования второго измерения для ото-
бражения и обработки (vmap не видит измерение групп, поэтому его вторым измерением является измерение векторов).
❻ Оповещение pmap о необходимости использования второго измерения (групп),
которое невидимо для vmap.
❼ Получение восьми групп вычисленных скалярных произведений.
❽ Удаление измерения групп.
Перепишем этот код с использованием именованных тензоров
и xmap(), не обращая внимания на то, насколько точно выполняется
часть распараллеливания (ниже в этом подразделе мы рассмотрим
это подробнее). Сосредоточимся исключительно на удобстве для
программиста написания такой обработки с применением имено
ванных осей.
Используем один вызов xmap() вместо вложенных вызовов pmap()/
vmap(). По умолчанию xmap() векторизует вычисления аналогично
vmap(), но не обеспечивает какое-либо распараллеливание, т. е. код
выполняется на одном устройстве. Поэтому в текущий момент наш
рассматриваемый здесь пример не равнозначен коду с вложенными
вызовами pmap()/vmap(). Часть, отвечающая за распараллеливание,
будет добавлена в следующем подразделе.
На рис. D.3 показано добавление и удаление именованных осей
при использовании xmap(). Не обращайте внимания на визуально
увеличившийся размер тензоров после добавления именованных
Приложение D
480
осей, это сделано только для того, чтобы вместить больше текста
в графический элемент, обозначающий тензоры. Размеры тензоров,
измеряемые в количестве элементов, остаются неизменными.
Добавление
именованных
осей
Немного
магии при вызове dot()
(в действительности
выполняются
две операции
векторизации,
подобные
vmap())
Удаление
именованных
осей
Из
м
фо ене
рм ни
ы е
Из
м
фо ене
рм ни
ы е
Векторизация по измерениям 'device' и 'batch'
Рис. D.3 Схема обработки данных для примера с большим транспонированным массивом
с использованием xmap() и именованных осей
В приведенном выше примере кода необходимо отметить, что
требуется отобразить функцию dot() по измерениям с индексом 1
(мы назвали его измерением 'device') и 2 (измерение 'batch'). Для
этого применяется форма записи {1:'device', 2:'batch'} для обо
их параметров функции внутри параметра in_axes. При выводе
результатов необходимо вставить эти измерения в позиции 0 и 1
соответственно, что отображается в содержимом параметра out_
axes=['device', 'batch', ...]. Порядок в указанном списке опреде
ляет числовую позицию начиная с нуля.
Листинг D.2
Измененный код с использованием xmap()
from jax.experimental.maps import xmap
vs = random.normal(rng_key, shape=(20_000_000,3))
v1s = vs[:10_000_000,:].T
v2s = vs[10_000_000:,:].T
❶
v1s.shape, v2s.shape
>>> ((3, 10000000), (3, 10000000))
v1sp = v1s.reshape((v1s.shape[0], 8, v1s.shape[1]//8))
v2sp = v2s.reshape((v2s.shape[0], 8, v2s.shape[1]//8))
v1sp.shape, v2sp.shape
>>> ((3, 8, 1250000), (3, 8, 1250000))
f = xmap(dot,
❷
481
Экспериментальные средства распараллеливания
in_axes=(
{1:'device', 2:'batch'},
{1:'device', 2:'batch'}
),
out_axes=['device', 'batch', ...]
)
x_xmap=f(v1sp,v2sp)
x_xmap.shape
❸
❹
❺
❻
>>> (8, 1250000)
x_xmap = x_xmap.reshape((x_xmap.shape[0]*x_xmap.shape[1]))
x_xmap.shape
>>> (10000000,)
jax.numpy.all(x_xmap == x_pmap)
>>> Array(True, dtype=bool)
❶
❷
❸
❹
❺
❻
❼
❼
❼
❼
Импорт xmap().
Использование xmap() как функции трансформации.
Сигнатура для первого аргумента dot().
Сигнатура для второго аргумента dot().
Сигнатура для возвращаемого значения dot().
Вызов трансформированной функции.
Проверка: результат должен совпадать с результатом, полученным с использованием pmap().
Что улучшилось в приведенном выше коде? Во-первых, здесь при
меняется всего лишь одна функция трансформации xmap() вместо
двух – vmap() и pmap(). Во-вторых, нет необходимости в отслежива
нии индексов по различным измерениям тензоров. Вы просто пере
даете их в параметре in_axes для преобразования в именованные
оси, позволяя обертываемой функции работать с остальными пози
ционными осями. Для возвращаемого тензора выполняется обрат
ное преобразование осей, указанных в параметре out_axes, в пози
ционные оси в позиции, заданные соответствующими значениями
этого параметра.
В приведенном выше примере мы использовали отображения,
определяемые словарями для входных аргументов, так как отобра
жаемые измерения не являлись основными (индекс 0 имеет измере
ние, хранящее отдельные компоненты вектора). Для выходных дан
ных применялись отображения, определяемые списками, поскольку
в этом случае перечислялись первые измерения, формирующие пре
фикс позиционных измерений.
Здесь один вызов xmap() равнозначен двум вложенным вызовам
vmap(). Приведенный выше код не использует распараллеливание,
Приложение D
482
поэтому необходимо уделить этому особое внимание и поговорить
о двух вызовах vmap() вместо комбинации pmap()+vmap().
Эйнштейновское обозначение суммирования Einsum
Эйнштейновское обозначение суммирования Einsum (Einstein summation) – весьма полезная функция для выражения скалярных произведений, векторных (внешних) произведений, умножений матрицы на
вектор и матрицы на матрицу. Это обобщение всех произведений со
многими измерениями.
Например, умножение матриц можно выразить, используя функцию
einsum() следующим способом:
import numpy as np
A = np.array([
[1, 1, 1],
[2, 2, 2],
[3, 3, 3]
])
B = np.array([
[1, 0, 0],
[0, 1, 0],
[0, 0, 1]
])
np.einsum("ij,jk->ik", A, B)
>>> array([[1, 1, 1],
>>>
[2, 2, 2],
>>>
[3, 3, 3]])
В основе функции einsum() заложено эйнштейновское обозначение
(форма записи) суммирования. Это чрезвычайно изящный способ выражения многих операций с тензорами. Он использует простой предметно-ориентированный язык (domain-specific language – DSL), напоминающий методику применения именованных осей. Такой DSL иногда
может компилироваться в высокопроизводительный код (именно этот
вариант рассматривался в примере функции, векторизованной вручную, в подразделе 6.1.2).
Функция einsum() принимает список входных тензоров и специальную
строку формата, такую, как показана выше: "ij,jk->ik". Часть слева от
символа -> относится к входным параметрам, часть справа определяет
вывод. Строка формата помечает измерения для всех входных тензоров
и для выводимого результата, и каждая буква соответствует конкретному измерению.
Для имен, присутствующих и во входных данных, и в выводе (так называемых свободных индексов (free indices); в приведенном здесь примере
i и k), функция einsum() фактически создает внешний цикл для соответ-
Экспериментальные средства распараллеливания
483
ствующих индексов. Для других индексов (называемых индексами суммирования (summation indices); в приведенном здесь примере j) создается
внутренний цикл с суммированием для измерений с тем же индексом.
В приведенном здесь примере строка формата "ij,jk->ik" выражает
два внешних цикла: один по первому измерению (i) первого тензора,
второй по второму измерению (k) второго тензора. Для второго измерения первого тензора и первого измерения второго тензора (оба имеют
имя j) выполняется поэлементное произведение с суммированием. Таким образом, функция einsum() для двух матриц с выбранной (описанной выше) строкой формата выполняет умножение матриц.
Другие полезные примеры: "a,a->" для скалярного (внутреннего) произведения двух векторов, "ii->i" для диагонали матрицы, "ab->ba" для
транспонирования матрицы, "Yab,Ybc->Yac" для пакетного умножения
матриц. Для векторизованного скалярного произведения, используемого в нашем примере, используется строка формата "ib,ib->b".
Функция einsum() реализована в NumPy и почти в каждом фреймворке
глубокого обучения, включая TensorFlow, PyTorch и JAX. Трансформацию
xmap() можно рассматривать как обобщенную функцию einsum(), поскольку она интерпретирует имена осей как объекты первого класса,
и имеется возможность реализации функции, работающей с такими
объектами. Функция einsum() никогда не позволяет напрямую взаи
модействовать с именованными осями. Более подробную информацию о функции einsum() можно получить здесь: https://rockt.github.
io/2018/04/30/einsum.
Правила распределения именованной оси
Рассмотрим, как именованные оси передаются в программе. Имено
ванные оси никогда не взаимодействуют с какой-либо позиционной
осью неявно, поэтому можно вызывать функцию (используя xmap()),
которая ничего не знает об именованных осях, с входными данны
ми, содержащими позиционные оси. Результат будет таким, как если
бы вы использовали vmap() для каждой отдельной именованной оси.
Именно это было сделано в предыдущем примере.
Если бинарная операция применяется к аргументам с различны
ми именованными осями, то эти оси распределяются посредством
широковещательной передачи (broadcasting) по своим именам. На
пример, если один операнд содержит именованную ось a, а в дру
гом операнде находится именованная ось b, то бинарная операция
(скажем, сложение) для таких операндов выдаст результат с обеими
осями a и b.
Предполагается, что все оси с одинаковыми именами имеют оди
наковые (или приемлемые для широковещательной передачи) фор
мы для всех аргументов в операции широковещательной передачи.
Приложение D
484
Поэтому если один операнд содержит именованную ось a, а в дру
гом операнде размещена ось с тем же именем, обе оси обязательно
должны иметь одинаковый размер или размер одной из них должен
быть равен 1. Именованная форма результата становится объедине
нием именованных форм этих входных данных.
В листинге D.3 используется функция, обрабатывающая два ар
гумента с различными именованными формами. Каждая форма со
держит именованную ось, отличающуюся от другой. Вы увидите, что
в результат включены обе именованные оси, поскольку была выпол
нена их широковещательная передача (распространение), и форма
результата стала объединением входных именованных форм.
Листинг D.3 Широковещательная передача (распространение)
в xmap()
image = random.normal(rng_key, shape=(480,640,3))
filters = random.normal(rng_key, shape=(5,3,3))
from jax.scipy.signal import convolve2d
❶
❷
def apply_filter(channel, kernel):
return convolve2d(channel, kernel, mode="same")
❸
apply_filters_to_image = xmap(apply_filter,
in_axes=(
{2:'channel'},
{0:'filter'}
),
out_axes={0:'filter', 3: 'channel'}
)
❹
❺
❻
❼
res = apply_filters_to_image(image, filters)
res.shape
>>> (5, 480, 640, 3)
❶
❷
❸
❹
❺
❻
❼
❽
❽
Генерация случайного RGB-изображения размером 640×480 пикселов.
Генерация пяти матричных фильтров размером 3×3.
Функция для применения одного фильтра к одному каналу изображения.
Генерация функции для применения многих фильтров ко многим каналам изображения.
Отображение по измерениям каналов для изображения.
Отображение по измерениям фильтров для набора фильтров.
Помещение измерения фильтров в первую позицию выводимого результата,
а измерения каналов – в последнюю позицию.
Результат содержит оба именованных измерения (фильтры ('filter’), h, w, каналы ('channel’)).
Мы применили двумерную функцию свертки к аргументам
с различными именованными осями. Двумерная функция сверт
Экспериментальные средства распараллеливания
485
ки работает с одним изображением, используя единственное ядро.
Мы добавляем именованное измерение для каналов изображения
в первый аргумент функции и другое именованное измерение для
отдельно размещаемых ядер фильтров во второй аргумент. Как
и ожидалось, результат содержит оба измерения, поскольку оба они
были распространены посредством широковещательной передачи
в другой тензор. Для выводимого значения мы выбрали конкрет
но заданные позиции этих измерений: измерение фильтров стало
первым, а измерение цветовых каналов – последним. Таким обра
зом, полученный в итоге тензор можно интерпретировать как пять
цветных изображений, сложенных в «стопку». Каждое изображение
включает три цветовых канала и результаты применения отдельно
го матричного фильтра к исходному изображению. Поскольку име
нованные оси должны быть равнозначными позиционным осям,
можно применять операции редукции к именованным осям.
ПРИМЕЧАНИЕ В интерфейсе JAX NumPy только несколько
функций поддерживают именованные оси. В настоящее вре
мя это функции jnp.sum(), jnp.max(), jnp.min().
В приведенном ниже примере (листинг D.4) имеется двумерная
матрица с именованными осями row и col. Применяется функция,
выполняющая редукцию по оси row, вычисляя сумму (с помощью
jnp.sum()) соответствующих элементов этой оси. Другая ось сохра
няется без изменений, и в результате мы получаем сумму каждого
столбца (так как строки редуцированы).
Листинг D.4
Редукция оси с использованием jnp.sum()
f = xmap(
lambda x: jnp.sum(x, axis=['row']),
in_axes=['row', 'col'],
out_axes=['col']
)
❶
❷
❸
C = jnp.array([
[1,2,3],
[4,5,6],
[7,8,9]
])
f(C)
>>> Array([12, 15, 18], dtype=int32)
❶
❷
❸
❹
❹
Функция вычисления суммы по оси row.
Объявление двух именованных осей во входных данных.
Объявление одной именованной оси в выходных данных.
Результат содержит сумму по строкам.
Приложение D
486
Также можно использовать коллективные операции. Все коллек
тивные операции, работающие внутри функции, трансформирован
ной с помощью pmap(), работают и с именованными осями. Поэтому
можно переписать функцию глобальной нормализации массива из
листинга 7.23, как показано ниже в листинге D.5.
Листинг D.5 Пример с вложенной функцией pmap(),
воспроизводящий листинг 7.23
arr = jnp.array(range(8)).reshape(2,4)
arr
❶
>>> Array([[0, 1, 2, 3],
[4, 5, 6, 7]], dtype=int32)
n = jax.pmap(
jax.pmap(
lambda x: x/jax.lax.psum(x, axis_name=('rows','cols')),
axis_name='cols'
),
axis_name='rows')
❷
jnp.sum(n(arr))
>>> Array(1., dtype=float32)
❶ Генерация небольшой матрицы.
❷ Выполнение вложенной трансформации pmap() по строкам и по столбцам.
❸ Проверка результата нормализации.
❸
Применяя xmap(), можно с легкостью переписать такую обработку
более простым способом, используя ту же функцию для обработки
одного элемента и тот же вызов коллективной функции.
Листинг D.6 Замена кода с вложенными функциями pmap()
на трансформацию xmap()
arr = jnp.array(range(8)).reshape(2,4)
arr
❶
>>> Array([[0, 1, 2, 3],
>>>
[4, 5, 6, 7]], dtype=int32)
n_xmap = xmap(
lambda x: x/jax.lax.psum(x, axis_name=('rows','cols')),
in_axes=['rows', 'cols', ...],
out_axes=['rows', 'cols', ...]
)
❷
❸
❸
jnp.sum(n_xmap(arr))
>>> Array(1., dtype=float32)
❹
Экспериментальные средства распараллеливания
❶
❷
❸
❹
487
Генерация небольшой матрицы.
Использование той же самой функции для нормализации.
Создание именованных осей в одном вызове xmap().
Проверка результата нормализации.
Код стал более простым и понятным, а результат остался тем
же. Но в текущий момент не используются параллельные устрой
ства. Это аналогично двум вызовам vmap(), выполняемым на одном
устройстве. Поэтому вполне естественно обратить внимание на понастоящему уникальное преимущество xmap() – возможность рас
параллеливания кода по сеткам аппаратных средств, сравнимых
с мощностью суперкомпьютера.
D.1.2
Распараллеливание и сетки аппаратных средств
По умолчанию xmap() векторизует вычисление тем же способом, что
и vmap(), но не организует какое-либо распараллеливание, т. е. код
выполняется на одном устройстве.
Для распараллеливания вычисления обязательно необходимо ис
пользовать оси ресурсов (resource axes). Ось ресурсов – это средство
управления тем, как xmap() выполняет вычисление.
Каждая ось, вводимая xmap(), присваивается одной или несколь
ким осям ресурсов. Источником происхождения осей ресурсов явля
ется сетка аппаратных устройств (hardware mesh), представляющая
собой n-мерный массив устройств с именованными осями.
С технической точки зрения сетка аппаратных устройств – это
объект, состоящий из двух компонентов:
n-мерный массив объектов устройств JAX – это те же объекты,
которые получаются с помощью вызовов функций jax.devices() or jax.local_devices(), но представленные типом np.array.
Будьте внимательны: это чистый массив NumPy (np.array), а не
массив JAX NumPy (jnp.array), поскольку объект устройства не
является корректным типом массива JAX;
кортеж имен осей ресурсов – длина кортежа обязательно долж
на соответствовать рангу массива устройств. То есть для трех
мерной сетки кортеж должен содержать три именованные оси
ресурсов.
Обычные TPU (по крайней мере доступные в настоящее время по
коления 2, 3 и 4) содержат четыре микросхемы (в каждой два ядра)
на плате (https://cloud.google.com/tpu/docs/system-architecture-tpuvm). TPU v2 и v3 используют топологию двумерного тора, а для TPU
v4 принята топология трехмерного тора. Поэтому можно предста
вить по умолчанию Cloud TPU как сетку микросхем 2×2 или сетку
ядер 4×2 (ядро является устройством, видимым для JAX).
Хотя физические устройства соединяются в физическую сетку
(схему), можно создать логическую сетку поверх физической. Это
Приложение D
488
позволяет абстрагировать сетку физических устройств и предостав
ляет возможность изменять форму логической сетки в соответствии
с конкретными потребностями вычислений.
JAX предоставляет специализированный контекстный менед
жер Mesh (https://docs.jax.dev/en/latest/jax.sharding.html#jax.sharding.
Mesh). Вы можете создать сетку, как показано в листинге D.7.
Листинг D.7
Создание контекстного менеджера Mesh
from jax.sharding import Mesh
import numpy as np
❶
devices = np.array(jax.devices()).reshape(4, 2)
❸
with Mesh(devices, ('x', 'y')):
...
❶ Импорт типа Mesh.
❷
❹
❺
❷ Импорт обычной библиотеки NumPy, поскольку необходим тип np.array.
❸ Создание двумерного массива устройств.
❹ Создание объекта типа Mesh с двумя осями ресурсов x и y.
❺ Теперь можно использовать вызовы xmap() с осями ресурсов.
В приведенном выше примере мы создали двумерный массив
доступных устройств и контекстный менеджер Mesh с двумя осями,
соответствующими осям сетки аппаратных устройств, x и y. После
этого можно распределить вычисление по этим двум осям.
Теперь можно отобразить логические оси, введенные xmap(), на
оси ресурсов, представленные Mesh. Например, требуется разделе
ние именованных осей rows и columns по осям ресурсов x и y. Для
такого отображения предназначен параметр axis_resources.
Параметр axis_resources представляет собой словарь, отобра
жающий оси, введенные текущим вызовом xmap(), на одну или не
сколько осей ресурсов. Каждое значение именованной оси будет
распределено по всем осям сетки, назначенной для этой именован
ной оси с помощью параметра axis_resources. Размер логической
оси непременно должен быть кратным размеру, соответствующему
измерению сетки. Мы распараллеливаем вычисление по сетке аппа
ратных устройств.
Листинг D.8 Распараллеливание вычисления xmap()
arr = jnp.array(range(10000)).reshape(100,100)
with Mesh(devices, ('x', 'y')):
n_xmap = xmap(
lambda x: x/jax.lax.psum(x, axis_name=('rows','cols')),
in_axes=['rows', 'cols', ...],
489
Экспериментальные средства распараллеливания
out_axes=['rows', 'cols', ...],
axis_resources={'rows': 'x', 'cols': 'y'}
)
❶
res = n_xmap(arr)
type(res), res.shape
>>> (jaxlib.xla_extension.ArrayImpl, (100, 100))
❶ Присваивание логических осей осям ресурсов.
❷ Результатом становится большой массив.
❷
Мы присвоили логические оси, введенные xmap(), осям ресурсов,
предоставленным имеющейся сеткой аппаратных устройств, с по
мощью словаря axis_resources.
Явное преимущество приведенного выше примера заключает
ся в том, что нам не нужно беспокоиться о количестве доступных
устройств и о соответствующем изменении формы входных данных.
Как вы помните, при использовании pmap() невозможно было пере
дать больше данных, чем количество доступных устройств, поэтому
приходилось изменять форму данных для получения размера ото
бражаемого измерения, не превышающего количество устройств
(как это было сделано в листинге 7.22). Все необходимые действия
были скрыто выполнены всего лишь в одном вызове xmap(), и это
великолепно!
При таком подходе принятое по умолчанию vmap()-подобное по
ведение xmap() становится больше похожим на поведение pmap().
Поэтому xmap() можно рассматривать как беспроблемный способ
интерполяции между стилями выполнения vmap() и pmap(). Можно
использовать xmap() как упрощенную замену pmap(), существенно
упрощающую программирование многомерных сеток аппаратных
устройств и автоматически распределяющую вычисление по не
скольким устройствам.
Трансформация xmap() выполняет два действия: разделение (или
распределение) и репликацию:
логическая ось A с размером x, отображаемая на ось ресур
сов B с размером y, распределяется по ней, т. е. разделяется на
фрагменты (размером x/y; x должен делиться на y без остатка),
и каждый фрагмент передается на собственное устройство;
логическая ось, не отображаемая на какую-либо ось ресурсов,
реплицируется на все устройства, т. е. каждое устройство полу
чает полную копию этой конкретной оси.
Например, предположим, что имеется двумерный тензор формы
(1000, 20) с первой осью rows и второй осью columns, а сетка аппарат
ных устройств для одной платы Cloud TPU с восемью ядрами имеет
форму с двумя осями x и y размером 4 и 2 соответственно.
Приложение D
490
Если выполняется некоторое вычисление без отображения осей
rows и columns на какие-либо оси ресурсов, то этот тензор полностью
реплицируется на все ядра Cloud TPU.
Если воспользоваться параметром axis_resources={'rows': 'x',
'columns': 'y'}, то первое измерение тензора ('rows' с размером
1000) разделяется на четыре фрагмента размером 1000/4 = 250,
а второе измерение ('columns' с размером 20) – на два фрагмента
размером 20/2 = 10.
Если только одно измерение отображается на ось ресурсов, то оно
разделяется на фрагменты, а второе неотображаемое измерение ре
плицируется на все устройства.
Теперь можно окончательно упростить пример из листинга D.2,
напрямую адаптировав более старый пример с использованием
pmap(). В этой адаптации мы удаляем индексы измерений тензора
и переходим на использование именованных осей. Кроме того, мы
выполняем только один вызов xmap() вместо двух вызовов vmap()
и pmap().
После получения информации о сетках аппаратных устройств
можно использовать их для распараллеливания вычислений на до
ступных устройствах. Но мы продолжаем полагаться на изменение
формы тензора вручную для соответствия количеству доступных
устройств, хотя в этом нет необходимости. При применении xmap()
можно воспользоваться автоматически распределением по выбран
ной оси, поэтому изменение формы вручную становится ненужным.
В листинге D.9 демонстрируется этот подход.
Листинг D.9
Исключение ручного изменения формы из листинга D.2
from jax.experimental.maps import xmap
rng_key = random.PRNGKey(42)
❶
vs = random.normal(rng_key, shape=(20_000_000,3))
v1s = vs[:10_000_000,:].T
v2s = vs[10_000_000:,:].T
v1s.shape, v2s.shape
>>> ((3, 10000000), (3, 10000000))
with Mesh(np.array(jax.devices()), ('device')):
f = xmap(dot,
in_axes=(
{1:'batch'},
{1:'batch'}
),
out_axes=['batch', ...],
axis_resources={'batch': 'device'}
)
❷
❸
❹
❹
❹
❺
491
Экспериментальные средства распараллеливания
x_xmap=f(v1s,v2s)
x_xmap.shape
>>> (10000000,)
❶
❷
❸
❹
❺
❻
❼
❻
❼
Импорт xmap().
Не требуется изменение формы для соответствия количеству устройств.
Использование идентификаторов сетки устройств.
Отображение по измерению пакетов.
Сегментирование измерения пакетов по сетке устройств.
Вызов трансформированной функции.
Мы получили 10 млн скалярных произведений.
Мы исключили операцию изменения формы тензора для соот
ветствия количеству доступных устройств. Мы также создали одно
мерную логическую сетку устройств. Наконец, мы распределили из
мерение batch по измерению сетки аппаратных устройств device.
Таким образом, мы автоматически сегментировали вычисление по
доступным устройствам без какого-либо изменения формы тензора
вручную.
Дополнительное преимущество заключается в том, что код ста
новится более универсальным. Теперь он не зависит от конкретного
количества устройств и может эффективно выполняться практиче
ски на любой аппаратной конфигурации с устройствами, сгруппиро
ванными в одномерную сетку.
Важный факт: любое присваивание в параметре axis_resources
никак не изменяет результаты вычисления (как минимум возни
кают проблемы с точностью значений с плавающей точкой из-за
другого порядка выполнения вычислений). Никогда не изменяется
семантика программы, меняется только способ организации вычис
ления и использования аппаратных устройств.
Благодаря такому подходу легко пробовать разнообразные спо
собы разделения на части в одной программе во многих распре
деленных сценариях, чтобы выбрать наиболее производительный
вариант. Вы можете без затруднений перемещать код между ком
пьютерами со скромными вычислительными ресурсами (например,
ноутбуком) и крупномасштабными системами (такими как TPU Pod).
Мы пропускаем более объемный пример нейронной сети с ис
пользованием xmap(), так как работа над этим экспериментальным
средством фактически была остановлена на незавершенной ста
дии. Не так-то просто без затруднений преобразовать пример SPMD
MNIST, поскольку некоторые функциональные средства отсутству
ют (правила дифференцирования для lax.pmax()), а для других от
сутствует документация (например, для lax.pdot()). Параметры
in_axes/out_axes становятся слишком сложными для древовидной
структуры (см. в репозитории книги пример, демонстрирующий эту
492
Приложение D
сложность). Механизм распараллеливания именованных осей, веро
ятно, будет проинспектирован в ближайшем будущем, и авторы JAX
могут изменить его.
Тем не менее даже в своем текущем состоянии xmap() уже спосо
бен помочь программисту упростить код и сделать его более понят
ным без изменения форм данных вручную, как это было показано
при замене вложенных вызовов pmap() и vmap().
Также существуют и другие интересные средства, заслуживающие
внимания. Во-первых, следует отметить shmap() (или shard_map();
его в шутку называют shpecialized_xmap) как современную замену
xmap(). Во время написания книги трансформация shmap() имела
статус предложения JAX Enhancement Proposal (JEP) (https://docs.jax.
dev/en/latest/jep/14273-shard-map.html).
Во-вторых, в экосистеме JAX существуют другие методики под
держки именованных (осей) тензоров. В главе 12 мы упоминали не
сколько новейших библиотек, заслуживающих внимания. На момент
написания книги одним из самых интересных вариантов является
библиотека Haliax (https://github.com/stanford-crfm/haliax).
D.2
Использование pjit()
для распараллеливания тензоров
Механизм распараллеливания, предоставляемый pmap() или xmap()
с осями ресурсов, был предназначен для варианта, в котором необ
ходимо распределить данные по различным акселераторам и вы
полнить одну функцию для всех этих распределенных фрагментов
данных. Это распараллеливание по данным (data parallelism).
При необходимости разделения (распределения) конкретной
функции вместо данных существует другой вариант. Например, это
может потребоваться для крупных нейронных сетей, которые не
способен обработать один акселератор. Современные большие ней
ронные сети, такие как GPT-3 с 175 млрд параметров или MegatronTuring NLG с 530 млрд параметров, являются примерами нейронных
сетей, с обработкой которых не справится ни один отдельный GPU.
Также может потребоваться ускорение вычислений посредством
выполнения их в параллельном режиме, если части вычислений не
зависят друг от друга.
В подобных случаях можно разделить модель на несколько частей
(например, по слоям или даже один слой может быть разделен). Это
распараллеливание модели (model parallelism). JAX поддерживает
распараллеливание модели с помощью трансформации pjit().
Как вы помните, pmap() позволяет выполнять одну программу на
нескольких устройствах. Каждое устройство получает свой сегмент
Экспериментальные средства распараллеливания
493
входных данных. Кроме того, вы можете написать программу для
обработки данных так, чтобы при необходимости обмена информа
цией между различными устройствами можно было вручную орга
низовать такой режим с помощью коллективных операций.
Используя pjit(), вы получаете возможность сегментировать
как данные, так и функцию (вариант с весами нейронной сети) по
существующей сетке аппаратных устройств. При этом сетка оста
ется той же самой, что и при использовании xmap() в предыдущем
разделе.
Необходимо точно определить, как вы хотите разделить свои вход
ные и выходные данные. Далее распределение функции по устрой
ствам происходит автоматически посредством распространения
заданных частей входных и выходных данных. Для всех промежуточ
ных тензоров pjit() автоматически определяет паттерн сегментиро
вания. Но вы также можете использовать ограничения сегментиро
вания для выбранных тензоров внутри своей программы, поскольку
это, возможно, поможет улучшить производительность функции. Нет
необходимости вручную вставлять какие-либо коллективные опера
ции в программу, pjit() сделает это при необходимости.
В итоге программа компилируется в представление XLA, как если
бы существовало только одно большое виртуальное устройство. Вы
сообщаете компилятору о необходимости сегментирования масси
вов в соответствии со стратегией сегментирования, а затем приме
няется механизм разделения XLA SPMD для генерации идентичной
программы для N устройств, которая обеспечивает обмен информа
цией между устройствами через коллективные операции. Код сохра
няет архитектуру SPMD (одна программа, множественные данные),
которая была определена при использовании pmap(). Но теперь вы
формируете программу на более высоком уровне абстракции.
Функция, возвращаемая трансформацией pjit(), сохраняет се
мантику исходной функции, но компилируется в представление
XLA, выполняемое на нескольких устройствах. В этом смысле транс
формация похожа на xmap(), которая никогда не изменяет семанти
ку программы, только способ организации вычислений и определя
ет, какие устройства используются.
Достоинство pjit() заключается в том, что эта трансформация
требует лишь нескольких изменений в исходном коде, и вы вноси
те изменения быстрее по сравнению с pmap(), где требуется больше
трудозатрат, но вам предоставляется бóльшая степень управления.
Интерфейс pjit() похож на интерфейс jit() и может работать как
декоратор для функции, требующей компиляции. В настоящее вре
мя трансформация pjit() весьма широко применяется и использу
ется гораздо чаще, чем xmap(), особенно если приходится работать
в параллельном режиме с трансформерами крупномасштабных ней
ронных сетей.
Приложение D
494
Внутренний механизм разделения на части XLA SPMD описан
в документе «GSPMD: General and Scalable Parallelization for ML Com
putation Graphs» (https://arxiv.org/abs/2105.04663), а pjit() пред
ставляет собой API, предъявляемый для использования механизма
разделения на части XLA SPMD в JAX. Подробное описание работы
механизма разделения на части XLA SPMD не относится к тематике
этой книги, поэтому всем интересующимся подробностями реко
мендуется начать с изучения вышеуказанного документа.
Для примеров в этом разделе я создал Cloud TPU с восемью ядра
ми и установил его соединение как локальной среды выполнения
с блокнотом Colab. Начинаем подробно разбираться с практическим
использованием pjit().
D.2.1
Основы работы с pjit()
Для практического использования pjit() требуется выполнение
трех условий:
создание спецификации сетки. Это та же спецификация сет
ки, которая использовалась в разделе про xmap(), т. е. логиче
ская многомерная сетка поверх физической сетки аппаратных
устройств, определяемая менеджером контекста jax.sharding.Mesh() (https://docs.jax.dev/en/latest/jax.sharding.html#jax.
sharding.Mesh). Функция будет использовать определение Mesh,
предоставляемое при вызове этой функции, а не определение
Mesh во время вызова pjit();
создание спецификации сегментирования. Здесь определяется
разделение на части входных и выходных данных с использо
ванием параметров in_shardings и out_shardings (in_axis_resources и out_axis_resources в более старых версиях) функции
pjit() (https://docs.jax.dev/en/latest/jax.experimental.pjit.html);
определение ограничений сегментирования (необязательное
требование). Чтобы определить такие ограничения для выбран
ных промежуточных тензоров внутри функции, можно предо
ставить рекомендации (подсказки), используя jax.lax.with_
sharding_constraint() (изначально jax.experimental.pjit.
with_sharding_constraint()). Такие ограничения способны
улучшить производительность. Ограничения сегментирования
кратко рассматривались в подразделе 8.3.1.
Рассмотрим несколько примеров.
Простой пример с одномерной сеткой
Сначала вернемся к простому примеру вычисления скалярных про
изведений по двум большим спискам векторов. В приведенном ниже
примере используются два длинных списка случайных векторов,
Экспериментальные средства распараллеливания
495
для которых требуется вычислить попарные скалярные произведе
ния в параллельном режиме. Мы сделали это с применением pmap()
в листингах с 7.6 по 7.9 и 7.17 и с применением pmap() в листингах D.2
и D.9. В рассматриваемом здесь примере мы создадим сетку и по
пробуем распараллелить вычисление, используя pjit(). На рис. D.4
показаны шаги этого процесса.
Определение
спецификации
сегментирования для
входных и выходных
данных
Определение
массива устройств
Функция pjit()
с заданными
спецификациями
сегментирования
Создание менеджера
контекста Mesh
с осями ресурсов
Выполнение
pjit()-трансформированной
функции внутри контекста Mesh
Рис. D.4 Процесс использования pjit()
для функции и сетки устройств
Здесь необходимо выделить несколько важных частей. Во-первых,
для части схемы, связанной с Mesh, необходимо выполнить два усло
вия для организации сетки:
простой (плоский) массив NumPy, так как менеджер контекста
Mesh принимает только массив устройств типа np.array;
имена осей ресурсов, предоставляемые в виде кортежа.
Мы используем простую одномерную сетку с одной осью. Для
этого применяется одномерный массив devices=np.array(jax.devices()), а в конструкторе Mesh() – кортеж с одним элементом для
именования оси устройств с помощью вызова Mesh(devices, ('devices',)).
Во-вторых, для части схемы, связанной с pjit(), предоставляет
ся аннотация сегментирования входных и выходных данных. Она
управляется параметрами in_shardings и out_shardings.
Значением аннотации сегментирования должно быть None, PartitionSpec или кортеж длин, равных количеству позиционных аргу
ментов распараллеливаемой функции. Значение PartitionSpec – это
кортеж, элементами которого могут быть значения None, строковые
имена осей сетки или кортежи имен осей сетки. Каждый элемент
кортежа описывает измерение сетки, по которому разделяется из
мерение входных данных.
Например, для тензора ранга 3 (тензора с тремя измерениями)
аннотация сегментирования PartitionSpec('x', 'y', None) озна
Приложение D
496
чает, что первое измерение данных сегментируется по оси x сетки,
второе измерение – по оси y, а третье измерение не сегментируется
(следовательно, реплицируется).
В рассматриваемом здесь примере сегментирование входных
данных не используется, поэтому передается значение in_shardings
=None. Это означает, что входные значения реплицируются на всех
устройствах, т. е. все устройства получают копию входных данных.
Для сегментирования выходных данных используется параметр
out_shardings=PartitionSpec('devices'), т. е. тензор, выводимый
функцией, должен быть сегментирован по оси devices.
Листинг D.10 Применение pjit() для распараллеливания вычислений
скалярных произведений
from jax import random
def dot(v1, v2):
return jnp.vdot(v1, v2)
❶
rng_key = random.PRNGKey(42)
vs = random.normal(rng_key, shape=(20_000_000,3))
v1s = vs[:10_000_000,:]
v2s = vs[10_000_000:,:]
v1s.shape, v2s.shape
❷
❷
❷
>>> ((10000000, 3), (10000000, 3))
from jax.sharding import Mesh
from jax.sharding import PartitionSpec
import numpy as np
devices = np.array(jax.devices())
devices
❸
❸
❹
❺
>>> array([TpuDevice(id=0, process_index=0, coords=(0,0,0), core_on_chip=0),
>>>
TpuDevice(id=1, process_index=0, coords=(0,0,0), core_on_chip=1),
>>>
TpuDevice(id=2, process_index=0, coords=(1,0,0), core_on_chip=0),
>>>
TpuDevice(id=3, process_index=0, coords=(1,0,0), core_on_chip=1),
>>>
TpuDevice(id=4, process_index=0, coords=(0,1,0), core_on_chip=0),
>>>
TpuDevice(id=5, process_index=0, coords=(0,1,0), core_on_chip=1),
>>>
TpuDevice(id=6, process_index=0, coords=(1,1,0), core_on_chip=0),
>>>
TpuDevice(id=7, process_index=0, coords=(1,1,0), core_on_chip=1)],
>>>
dtype=object)
f = pjit(dot,
in_shardings=None,
out_shardings=PartitionSpec('devices')
)
❻
❼
❽
497
Экспериментальные средства распараллеливания
with Mesh(devices, ('devices',)):
x_pjit=f(v1s,v2s)
❾
❿
>>> ...
>>> ValueError: One of pjit outputs is incompatible with its
➥sharding annotation NamedSharding(mesh={'devices': 8},
➥ spec=PartitionSpec('devices',)): Sharding NamedSharding(
➥mesh={'devices': 8}, spec=PartitionSpec('devices',))
➥is only valid for values of rank at least 1, but was
➥applied to a value of rank 0. For scalars the
➥PartitionSpec should be P()
# ValueError: один из элементов выходных данных pjit несовместим с
# собственной аннотацией сегментирования NamedSharding(mesh={'devices': 8},
# spec=PartitionSpec('devices',)): Сегментирование NamedSharding(
# mesh={'devices': 8}, spec=PartitionSpec('devices',))
# корректно только для значений с рангом минимум 1, но было
# применено к значению с рангом 0. Для скалярных значений
# аннотация PartitionSpec должна определять P()
❶
❷
❸
❹
❺
❻
❼
❽
❾
❿
Хорошо знакомая нам функция для вычисления скалярного произведения двух векторов.
Генерация некоторых случайных векторов.
Импорт для использования pjit() и Mesh.
Для создания объекта Mesh требуется библиотека NumPy.
Будем использовать одномерную сетку.
Вызов pjit() для функции dot().
Входные данные не разделяются (поэтому реплицируются).
Разделение выходных данных по оси devices сетки.
Создание менеджера контекста Mesh.
Вызов функции, обработанной трансформацией pjit().
Мы почти завершили работу, но получили сообщение об ошибке,
информирующее о том, что такое разделение допустимо только для
тензоров с рангом 1 или выше. Функция dot() возвращает скаляр
ное значение (тензор ранга 0), именно поэтому и возникла ошибка.
Проблема заключается в том, что функция работает с парой век
торов, выдает в результате одно число и не векторизуется (т. е. не
может работать с пакетами пар векторов). Мы знаем, что ситуация
легко исправляется: нужно просто создать векторизованную версию
функции, способную работать с массивом векторов. В листинге D.11
проблема устраняется, и мы получаем ожидаемый результат.
Листинг D.11 Применение pjit() для распараллеливания
вычислений скалярных произведений (исправленная
версия)
f = pjit(jax.vmap(dot),
in_shardings=None,
out_shardings=PartitionSpec('devices')
)
❶
Приложение D
498
with Mesh(devices, ('devices',)):
x_pjit=f(v1s,v2s)
x_pjit.shape
>>> (10000000,)
❷
❷
❸
❶ Автоматическая векторизация функции dot() с помощью трансформации vmap().
❷ Вычисление распараллеленной функции.
❸ Получение ожидаемого результата.
Здесь было внесено единственное изменение: в трансформацию
pjit() теперь передается автоматически векторизованная функция.
После этого трансформированная функция выдает корректные ре
зультаты (см. рис. D.5).
Рис. D.5
Наглядная схема обработки данных из листинга D.11
Необходимо более глубоко понять происходящее в приведенном
выше примере. Во-первых, мы подготовили функцию, способную
обрабатывать массивы векторов и формировать массивы скалярных
произведений. Это уже знакомая нам часть из главы 6, поскольку
мы воспользовались трансформацией vmap() для автоматического
создания такой функции.
Далее мы пометили входные данные как неразделяемые. Эти дан
ные реплицируются на всех устройствах, чтобы на каждом устрой
499
Экспериментальные средства распараллеливания
стве существовала полная копия входного массива. В нашем случае
это неоптимальное решение, поскольку полный входной массив не
требуется для вычисления части выводимого результата. Чтобы по
лучить часть результата, нужна только соответствующая часть вход
ных данных. Скоро мы устраним и эту неоптимальность.
Затем для выходных данных мы задали пометку, определяющую,
что результат должен сегментироваться по оси devices. Мы работа
ем с одномерной сеткой с восемью устройствами. Выходные данные
также являются одномерным массивом парных скалярных произ
ведений, поэтому здесь выходной массив разделяется на восемь ча
стей и сегментируется по всем восьми устройствам.
Фактически это означает, что каждое устройство получает два
полных массива формы (10 млн, 3), но вычисляет массив формы
(10 млн / 8), т. е. потребляет всего лишь 1/8 входных массивов.
Сделаем этот код более эффективным – разделим на части еще
и входные данные. В листинге D.12 изменяется значение параметра
in_shardings для обеспечения сегментирования входных данных.
Листинг D.12
Добавление сегментирования входных данных
f = pjit(jax.vmap(dot),
in_shardings=PartitionSpec('devices'),
out_shardings=PartitionSpec('devices')
)
❶
with Mesh(devices, ('devices',)):
x_pjit=f(v1s,v2s)
x_pjit.shape
>>> (10000000,)
❶ Добавление сегментирования входных данных.
❷ Получение того же результата.
❷
Мы добавили сегментирование входных данных по той же оси
devices. Но здесь имеется небольшой нюанс, требующий пояснения.
Функция принимает два аргумента, однако содержимое параметра
in_shardings на первый взгляд не ссылается на два значения. Что это
означает на самом деле?
Как уже было отмечено ранее, значением параметра in_shardings
должно быть None, PartitionSpec или кортеж длиной, равной коли
честву позиционных аргументов распараллеливаемой функции.
В предыдущем примере мы уже использовали значение None, теперь
же берем значение PartitionSpec.
Тем не менее параметр in_shardings=PartitionSpec('devices') не
выглядит как ссылка на два аргумента функции. В действительности
эта спецификация определяет, что первый массив сегментируется,
Приложение D
500
т. е. разделяется на восемь частей по своему первому измерению,
а второй массив реплицируется в полном объеме на всех восьми
устройствах.
Абсолютно правильная спецификация должна выглядеть так: in_sh
ardings=(PartitionSpec('devices', None), PartitionSpec('devices',
None)), чтобы оба входных аргумента функции сегментировались по
своему первому измерению (с индексом 0).
В приведенном выше примере со спецификацией in_shardings=P
artitionSpec('devices') и двумя тензорами ранга 2 (первый и вто
рой элементы входных данных для функции dot() – оба являются
массивами векторов) процедура сегментирования будет равнознач
на спецификации in_shardings=(Partition Spec('devices', None),
None).
В нашем случае это также равнозначно двум спецификациям
in_shardings=PartitionSpec('devices', None) и in_shardings=(Part
itionSpec('devices'), None), поскольку все перечисленные специ
фикации определяют сегментирование только первого аргумента
функции по первому измерению, но, строго говоря, они описывают
различные варианты. Значение None в in_shardings=PartitionSpec
('devices', None) явно ссылается на второе измерение первого
входного аргумента, а значение None в in_shardings=(PartitionSpec
('devices'), None) – на второй аргумент (см. рис. D.6).
Для сегментирования обоих массивов необходимо изменить
значение in_sharding так, чтобы оно ссылалось на оба аргумента
функции.
Листинг D.13 Добавление сегментирования входных данных для обоих
аргументов
f = pjit(jax.vmap(dot),
in_shardings=(PartitionSpec('devices'), PartitionSpec('devices')),
out_shardings=PartitionSpec('devices')
)
❶
with Mesh(devices, ('devices',)):
x_pjit=f(v1s,v2s)
x_pjit.shape
>>> (10000000,)
❶ Добавление сегментирования входных данных для обоих аргументов функции.
❷ Проверка формы.
❷
Мы окончательно сегментировали оба входных аргумента, и те
перь каждое устройство работает только с собственной частью обоих
входных массивов, следовательно, ресурсы используются эффектив
но (см. рис. D.7).
Экспериментальные средства распараллеливания
Рис. D.6
Наглядная схема обработки данных из листинга D.12
Рис. D.7 Наглядная схема обработки данных из листинга D.13
501
Приложение D
502
Теперь мы готовы перейти к более сложному примеру с двумер
ной сеткой.
Пример с двумерной сеткой
Предположим, что имеются очень большие векторы, состоящие не
из трех компонентов, как ранее, а, скажем, из 10 000 компонентов.
Хотя каждый такой вектор можно разместить на собственном от
дельном устройстве, возможно, также имеет смысл распределить
векторы по нескольким устройствам. Скалярное произведение лег
ко сегментируется, так как это всего лишь сумма соответствующих
произведений элементов вектора (см. рис. D.8).
Рис. D.8
Как сегментируется скалярное произведение
Для сегментирования по двум измерениям (по измерению самих
векторов и по измерению их компонентов) необходимо подготовить
двумерную сетку. В приведенном ниже примере сегментируются
два входных тензора ранга 2 по обоим измерениям, а результат –
только по одному измерению. Процесс может показаться достаточ
но сложным, поэтому начнем с рассмотрения его наглядной схемы,
показанной на рис. D.9.
В рассматриваемом здесь примере выполнено гораздо боль
ше операций сегментирования, чем ранее. Попробуем более вни
мательно разобраться в том, что происходит. Сначала мы создали
4000 пар векторов с 10 000 элементов в каждом. Размер векторов
намного превосходит длины векторов, использовавшихся в преды
дущих примерах. Далее мы сформировали массив устройств 4×2 для
сетки устройств и присвоили осям сетки имена x и y.
Оба входных аргумента функции dot() сегментируются по пер
вому и второму измерениям. Первое измерение (с размером 4000)
распределяется по оси x (с размером 2), и создаются сегменты с раз
мером 2000, который индексирует подмножества векторов. Второе
измерение (с размером 10 000) распределяется по оси y (с разме
ром 4), и создаются сегменты с размером 2500, индексирующим
подмножества компонентов векторов. Выходной массив распреде
ляется только по одной оси (в данном случае по оси x), так как это
тензор ранга 1.
Экспериментальные средства распараллеливания
Тензор
случайных
значений
Первый набор
векторов
Второй набор
векторов
Наборы векторов
сегментируются по сетке
устройств. Форма каждого
сегмента (2000, 2500)
Распределение каждой
пары сегментов по TPU.
Каждый TPU содержит
два сегмента с формой
(2000, 2500)
Сетка устройств
Каждый TPU вычисляет
частичное скалярное
произведение. Форма
результата (2000, 1)
Сетка устройств
Сохранение сегментирования по первому
измерению, но удаление второго
измерения посредством суммирования
всех частичных результатов по нему
Теперь нет частичных скалярных произведений,
и мы получили два сегмента для завершения
вычисления скалярных произведений. Форма
каждого сегмента (2000, )
Объединение двух сегментов в конечный
результат скалярного произведения.
Форма результата (4000, )
Рис. D.9 Сегментирование вычисления скалярного произведения
по двумерной сетке
503
Приложение D
504
Вычисление производится следующим образом: каждое устрой
ство может вычислить частичное скалярное произведение на основе
имеющихся у него сегментов. Поэтому на каждом устройстве сегмент
с 2500 элементами – векторами, состоящими из 10 000 элементов,
из первого массива поэлементно умножается на сегмент с 2500 эле
ментами – векторами, состоящими из 10 000 элементов, из второго
массива, и это выполняется для каждой из 2000 пар векторов, раз
мещенных на конкретном устройстве. Каждое частичное скалярное
произведение дает в результате одно число. Затем мы одновремен
но получаем сумму всех частичных скалярных произведений для
каждого вектора (здесь четыре таких частичных произведения). Это
делается с помощью коллективных операций, скрытых от нас.
В итоге первая строка сетки устройств содержит скалярное про
изведение для первых 2000 векторов из обоих массивов. Во второй
строке сетки устройств хранится скалярное произведение для вто
рой половины векторов из обоих массивов. Наконец, мы просто объ
единяем два оставшихся сегмента и получаем конечный результат
с 4000 скалярных произведений.
После подробного объяснения достаточно просто преобразовать
схему в исходный код.
Листинг D.14 Сегментирование по двумерной сетке
from jax.sharding import PartitionSpec as P
from jax.sharding import Mesh
import numpy as np
❶
rng_key = random.PRNGKey(42)
vs = random.normal(rng_key, shape=(8_000,10_000))
v1s = vs[:4_000,:]
v2s = vs[4_000:,:]
v1s.shape, v2s.shape
❷
❷
❷
>>> (4000, 10000), (4000, 10000))
devices = np.array(jax.devices()).reshape(2, 4)
devices
>>> array([[TpuDevice(id=0,
➥core_on_chip=0),
>>>
TpuDevice(id=1,
➥core_on_chip=1),
>>>
TpuDevice(id=2,
➥core_on_chip=0),
>>>
TpuDevice(id=3,
➥core_on_chip=1)],
>>>
[TpuDevice(id=4,
process_index=0, coords=(0,0,0),
process_index=0, coords=(0,0,0),
process_index=0, coords=(1,0,0),
process_index=0, coords=(1,0,0),
process_index=0, coords=(0,1,0),
❸
505
Экспериментальные средства распараллеливания
➥core_on_chip=0),
>>>
TpuDevice(id=5, process_index=0, coords=(0,1,0),
➥core_on_chip=1),
>>>
TpuDevice(id=6, process_index=0, coords=(1,1,0),
➥core_on_chip=0),
>>>
TpuDevice(id=7, process_index=0, coords=(1,1,0),
➥core_on_chip=1)]],
>>>
dtype=object)
def dot(v1, v2):
return jnp.vdot(v1, v2)
f = pjit(jax.vmap(dot),
in_shardings=(P('x', 'y'), P('x', 'y')),
out_shardings=P('x')
)
with Mesh(devices, ('x','y')):
x_pjit=f(v1s,v2s)
❹
❺
❻
x_pjit.shape
>>> (4000,)
❼
Эта команда импорта помогает сократить объем ручного ввода.
Генерация 8000 «широких» векторов, состоящих из 10 000 компонентов.
Подготовка двумерного массива устройств для создания сетки.
Оба входных параметра функции dot() (v1 и v2) сегментируются по обоим измерениям тензора.
❺ Выходной параметр сегментируется только по первому измерению (это единственное измерение).
❻ Создание и использование менеджера ресурсов Mesh.
❼ Проверка формы итогового результата.
❶
❷
❸
❹
Самая интересная часть находится внутри вызова pjit(), где мы
предоставляем спецификации сегментирования. Все остальное вы
полняется автоматически.
Если вы хотите подробнее изучить сгенерированный код, чтобы
увидеть автоматически добавленные коллективные операции, по
требуется заглянуть глубже уровня Jaxpr. К этой части можно полу
чить доступ только после компиляции в итоговое представление
HLO. В прилагаемом к книге блокноте вы найдете код для получения
представления HLO после компиляции и там сможете рассмотреть
операции полного редуцирования.
Многочисленные примеры применения сегментирования и рас
пределения по устройствам приведены в качественном блоге:
https://irhum.github.io/blog/pjit.
В последнем примере мы использовали pjit() для распаралле
ливания по данным и распараллеливания тензора (т. е. модели).
Списки векторов были разделены на сегменты (распараллеливание
Приложение D
506
по данным), а каждый вектор был разбит на части (распараллели
вание тензора).
В завершение темы рассмотрим пример создания многослойного
перцептрона (MLP) для классификации MNIST.
D.2.2
Пример создания многослойного перцептрона (MLP)
с использованием pjit()
Как уже было отмечено выше, pjit() может оказаться весьма полез
ным для сегментирования крупных моделей нейронных сетей. Вос
пользуемся примером создания небольшой нейросети для классифи
кации изображений MNIST, чтобы вы могли сравнить код с другими
методиками, которые мы применяли ранее. Начнем с процедуры за
грузки набора данных, которая ничем не отличается от предыдущих
версий, но все же здесь она воспроизводится для удобства.
Листинг D.15 Загрузка набора данных (воспроизводится листинг 7.25)
import tensorflow as tf
import tensorflow_datasets as tfds
data_dir = '/tmp/tfds'
data, info = tfds.load(name="mnist",
data_dir=data_dir,
as_supervised=True,
with_info=True)
data_train = data['train']
data_test = data['test']
HEIGHT = 28
WIDTH = 28
CHANNELS = 1
NUM_PIXELS = HEIGHT * WIDTH * CHANNELS
NUM_LABELS = info.features['label'].num_classes
NUM_DEVICES = jax.device_count()
BATCH_SIZE = 32
❶
def preprocess(img, label):
"""Resize and preprocess images."""
return (tf.cast(img, tf.float32)/255.0), label
train_data = tfds.as_numpy(
data_train.map(preprocess).batch(
NUM_DEVICES*BATCH_SIZE).prefetch(1)
)
test_data = tfds.as_numpy(
data_test.map(preprocess).batch(
❷
507
Экспериментальные средства распараллеливания
NUM_DEVICES*BATCH_SIZE).prefetch(1)
)
❷
len(train_data)
>>> 235
❸
❶ Новая константа для определения количества вычислительных устройств.
❷ Запрос более крупных пакетов с размером 32 * (количество устройств).
❸ Этот набор данных содержит 235 больших пакетов.
Далее используется та же структура MLP без каких-либо измене
ний. Этот код также воспроизводится здесь для удобства.
Листинг D.16
Структура MLP (воспроизводится листинг 7.26)
import jax
import jax.numpy as jnp
from jax import grad, jit, vmap, value_and_grad
from jax import random
from jax.nn import swish, logsumexp, one_hot
LAYER_SIZES = [28*28, 512, 10]
PARAM_SCALE = 0.01
❶
def init_network_params(sizes, key=random.PRNGKey(0), scale=1e-2):
"""Initialize all layers for a fully-connected neural network
with given sizes"""
# Инициализация всех слоев для полностью связанной нейронной сети
# с заданными размерами.
❷
def random_layer_params(m, n, key, scale=1e-2):
"""A helper function to randomly initialize
weights and biases of a dense layer"""
# Вспомогательная функция для случайно выбранной инициализации
# весов и отклонений плотного слоя.
w_key, b_key = random.split(key)
return (scale * random.normal(w_key, (n, m)),
scale * random.normal(b_key, (n,)))
keys = random.split(key, len(sizes))
return [random_layer_params(m, n, k, scale)
for m, n, k in zip(sizes[:-1], sizes[1:], keys)]
init_params = init_network_params(
LAYER_SIZES, random.PRNGKey(0), scale=PARAM_SCALE)
def predict(params, image):
"""Function for per-example predictions."""
# Функция для прогнозов по одному образцу.
activations = image
for w, b in params[:-1]:
outputs = jnp.dot(w, activations) + b
❸
❹
Приложение D
508
activations = swish(outputs)
final_w, final_b = params[-1]
logits = jnp.dot(final_w, activations) + final_b
return logits
batched_predict = vmap(predict, in_axes=(None, 0))
❶
❷
❸
❹
❺
Определение количества нейронов в каждом полностью связанном слое.
Функция для случайно выбранной инициализации параметров.
Подготовка начальных параметров.
Функция прямого прохода.
Генерация пакетной функции прямого прохода.
❺
Теперь необходимо отредактировать код функций потерь и об
новления. Возвращаемся к исходным функциям из главы 2, так как
нет необходимости в организации обмена информацией вручную
с использованием коллективных операций. В обновленной части
используется вызов pjit().
Листинг D.17 Функции потерь и обновления
from jax.experimental.pjit import pjit
from jax.sharding import PartitionSpec as P
from jax.sharding import Mesh
import numpy as np
❶
❶
❶
INIT_LR = 1.0
DECAY_RATE = 0.95
DECAY_STEPS = 5
NUM_EPOCHS = 20
def loss(params, images, targets):
"""Categorical cross entropy loss function."""
# Категориальная функция потерь перекрестной энтропии.
logits = batched_predict(params, images)
log_preds = logits - logsumexp(logits)
return -jnp.mean(targets*log_preds)
❷
def update(params, x, y, epoch_number):
loss_value, grads = value_and_grad(loss)(params, x, y)
lr = INIT_LR * DECAY_RATE ** (epoch_number / DECAY_STEPS)
return [(w - lr * dw, b - lr * db)
for (w, b), (dw, db) in zip(params, grads)], loss_value
❸
f_update = pjit(update,
in_shardings=(None, P('x'), P('x'), None),
out_shardings=None
)
❹
❺
❻
❶ Требуемые команды импорта.
❷ Та же функция потерь, что и в предыдущих версиях.
509
Экспериментальные средства распараллеливания
❸
❹
❺
❻
Та же функция обновления, что и в главе 2.
Вызов pjit() для функции обновления.
Сегментирование обоих параметров x и y по оси сетки x (это будет ось пакетов).
Итоговый результат не сегментируется.
Здесь самая интересная часть выделена полужирным шрифтом.
Все прочие части не изменились по сравнению с кодом из главы 2.
Сначала выполняются новые команды импорта. Они просты и по
нятны, поэтому не требуют каких-либо комментариев. Далее функ
ция update() снова упрощается. Вызовы коллективных операций
теперь не нужны, также нет необходимости в организации вручную
обмена информацией между устройствами. Код становится более
простым, и это великолепно. Сравните с кодом из листинга 7.27.
Теперь о вызове pjit(). Мы определяем, что первый элемент вход
ных данных params для функции update() должен быть реплициро
ван. Второй и третий параметры (x и y) сегментируются по первому
измерению соответствующих тензоров. Фактически это измерение
пакетов. Четвертый параметр – это просто скаляр, поэтому он ре
плицируется. В данном случае выводимый результат не сегментиру
ется, поскольку функция update() возвращает обновленный набор
весов нейронной сети, и в рассматриваемом здесь примере не тре
буется сегментирование весов.
Последний этап – организация цикла тренировки.
Листинг D.18
Полный цикл тренировки
devices = np.array(jax.devices())
devices
❶
>>> array([CpuDevice(id=0), CpuDevice(id=1), CpuDevice(id=2), CpuDevice(id=3),
>>>
CpuDevice(id=4), CpuDevice(id=5), CpuDevice(id=6), CpuDevice(id=7)],
>>>
dtype=object)
def batch_accuracy(params, images, targets):
images = jnp.reshape(images, (len(images), NUM_PIXELS))
predicted_class = jnp.argmax(batched_predict(params, images), axis=1)
return jnp.mean(predicted_class == targets)
f_batch_accuracy = pjit(batch_accuracy,
in_shardings=(None, P('x'), P('x')),
out_shardings=None
)
def accuracy(params, data):
accs = []
for images, targets in data:
accs.append(f_batch_accuracy(params, images, targets))
return jnp.mean(jnp.array(accs))
import time
❷
❸
❹
510
Приложение D
params = init_params
with Mesh(devices, ('x',)):
for epoch in range(NUM_EPOCHS):
start_time = time.time()
losses = []
for x, y in train_data:
x = jnp.reshape(x, (len(x), NUM_PIXELS))
y = one_hot(y, NUM_LABELS)
params, loss_value = f_update(params, x, y, epoch)
losses.append(jnp.sum(loss_value))
epoch_time = time.time() - start_time
❺
❻
train_acc = accuracy(params, train_data)
test_acc = accuracy(params, test_data)
print("Epoch {} in {:0.2f} sec".format(epoch, epoch_time))
print("Training set loss {}".format(jnp.mean(jnp.array(losses))))
print("Training set accuracy {}".format(train_acc))
print("Test set accuracy {}".format(test_acc))
>>>
>>>
>>>
>>>
>>>
>>>
>>>
>>>
…
Epoch 0 in 39.10 sec
Training set loss 0.41040703654289246
Training set accuracy 0.9299499988555908
Test set accuracy 0.931010365486145
Epoch 1 in 37.77 sec
Training set loss 0.37730318307876587
Training set accuracy 0.9500166773796082
Test set accuracy 0.9497803449630737
❶ Подготовка сетки устройств.
❷ Создание JIT-скомпилированной сегментированной функции для вычисления точности по
❸
❹
❺
❻
пакетам.
Мы сегментируем функцию по параметрам images и targets.
Использование JIT-скомпилированной сегментированной функции.
Установка и настройка менеджера контекста Mesh.
Вызов JIT-скомпилированной сегментированной функции для обновления параметров.
Вот и все. Если организовать вызовы pjit() как аннотации функ
ций, то код должен стать почти аналогичным коду из главы 2. Един
ственным существенным различием являлось бы использование
pjit() вместо jit() и наличие сетки устройств.
Этот код намного проще SPMD-кода из раздела 7.1, использующе
го pmap(). Но при этом получается практически тот же результат, что
и при выполнении кода, распараллеленного вручную.
Сохраняются некоторые различия, поскольку SPMD-механизм
разделения на части основан на некоторой эвристической методи
ке, поэтому, возможно, не все делает правильно. Иногда такой код
приводит к ухудшению производительности по сравнению с разде
лением вручную с использованием pmap().
Экспериментальные средства распараллеливания
511
Таким образом, имеет место следующий компромисс: либо вы
быстро пишете более простой код и используете pjit(), либо полу
чаете явные гарантии и улучшенную производительность при более
сложном коде с применением pmap(), где также может потребовать
ся набор специальных навыков, чтобы достичь желаемой эффектив
ности использования аппаратного оборудования.
Резюме
Трансформация xmap() помогает распараллеливать функции про
ще, чем pmap(), с меньшим объемом кода, заменяя вложенные
вызовы pmap() и vmap(), а также без изменения формы тензоров
вручную.
xmap() – это экспериментальное средство, применяющее модель
программирования с именованными осями, т. е. вводящее имено
ванные оси в дополнение к позиционным осям тензора.
Именованные оси никогда не взаимодействуют неявно с любыми
позиционными осями.
Оси ресурсов основаны на сетке аппаратных устройств; это
n-мерный массив устройств с именованными осями, представ
ленный менеджером контекста Mesh.
Каждая ось, введенная трансформацией xmap(), присваивается
одной или нескольким осям ресурсов.
xmap() – это беспроблемный способ интерполяции между стилями
выполнения vmap() и pmap().
Можно использовать xmap() как упрощенную замену pmap(), су
щественно упрощающую программирование многомерных сеток
аппаратных устройств и автоматически распределяющую вычис
ление по нескольким устройствам.
В настоящее время xmap() объявлен устаревшим и не рекомендуе
мым к использованию функциональным средством; его заменяет
shard_map(). xmap() был удален в версии JAX 0.4.31 (29 июля 2024 г.).
pjit() – это еще одно экспериментальное функциональное сред
ство, помогающее сегментировать как данные, так и функции
(веса для варианта нейронной сети) по существующей сетке ап
паратных устройств.
Используя pjit(), вы определяете, как нужно разделять на части
входные и выходные данные. Далее распределение функции по
устройствам происходит автоматически посредством распреде
ления сформированных частей входных и выходных данных.
Можно предоставить компилятору информацию о том, как сег
ментировать промежуточные переменные функции, применяя
функцию jax.lax.with_sharding_constraint(), во многом похо
512
Приложение D
жую на функцию jax.device_put(), но используемую внутри jitдекорированных функций.
pjit() фактически компилирует пользовательскую программу
в представление XLA так, как если бы существовало только одно
большое виртуальное устройство. pjit() использует XLA SPMD
механизм разделения на части для генерации идентичной про
граммы для N устройств с обеспечением обмена информацией
между устройствами через коллективные операции.
pjit() использует те же спецификации сетки, что и xmap().
В настоящее время pjit() и jit() (тема главы 5) объединены
в один универсальный интерфейс, поэтому настоятельно реко
мендуется использовать только jit().
Сегментирование тензоров с распределенными массивами (тема
главы 8) предоставляет современный способ компиляции и вы
полнения функций JAX в средах со многими хостами или многими
ядрами.
Предметный указатель
Символы
12_loss, стандартная функция
потерь, 408
@jit, аннотация, 72, 181, 183, 219, 220
@nn.compact, аннотация, 400
A
Acme, библиотека, 455
ahead-of-time (AOT) compilation, 211
all_gather(), коллективная
операция, 243
Alpa, библиотека, 451
AOT-компиляция, 179, 181, 211
argnums, параметр, 161, 172
array(), функция (конструктор
массива), 102
ASIC (application-specific integrated
circuit − интегральная микросхема
специального назначения), 105
Asynchronous dispatch, 111
autodiff, 39, 65, 133, 141, 149
вычисление градиента
линейная регрессия, 144
котангенс, 171
режим, 142
касательной, 166
обратный, 142, 164, 171
прямой, 142, 164, 166
сопряженное значение, 171
трассировка оценок, 165
forward mode, 142
reverse mode, 142
Autodiff Cookbook, 175
Autograd, 37
Automatic vectorization, 230
AutoTokenizer, класс, 429
Auto-vectorization, 223
AXLearn, библиотека, 447
B
Backpropagation algorithm, 171
Batch, 53
BatchNorm
слой нормализации пакетов, 415
batch_stats, коллекция, 416
params, коллекция, 416
use_running_average, параметр, 416
batch_stats, коллекция, 416
Bayex, библиотека, 456
Big Vision, библиотека, 456
BitGenerator, генератор RNG, 351
BlackJAX, библиотека, 456
block_until_ready(), метод объекта
Array, 111
Brax, библиотека, 458
broadcast_to(), функция NumPy, 241
514
Предметный указатель
BYOL (Bootstrap Your Own
Latent), методика обучения
с самоконтролем, 27
C
Categorical cross-entropy function, 67
chain(), функция, 407
Chex, библиотека, 387, 452
Classification
binary, 48
multiclass, 48
multilabel, 48
CLIP, модель объединения
изображений с встраиванием
текста, 426
clip(), функция, 96
CLM, causal language model, 431
Cloud TPU, система в Google
Cloud, 113
CLU, библиотека, 410
вычисление метрик, 410
Metric, интерфейс, 410
from_output(), функция, 411
metrics.Collection, интерфейс, 411
CNN, convolutional neural network, 92,
413
Coax, библиотека, 454
CoDeX, библиотека, 457
Colab TPU, среда времени
выполнения, 113
Collective operation (op), 242, 276
Common Loop Utils (CLU),
библиотека, 452
Composable function
transformation, 39
compute(), функция, 411
cond(), базисный элемент, 203
Constant folding, 186
conv_general_dilated, функция, 127
ConvNeXt, 414
convolve2d(), функция, 100
core_on_chip, атрибут, 116
cosine_distance, стандартная функция
потерь, 408
cost_analysis(), функция, 215
CPU, центральный процессор
(ЦП), 105
create_device_mesh(), функция, 310,
321
ctc_loss, стандартная функция
потерь, 408
D
DAG, directed acyclic graph, 148
DALL·E Mega, модель генерации
изображений, 426
DALL·E Mini, модель генерации
изображений, 426
Data parallelism, 285
DeepXDE, библиотека, 458
Dense tensor, 125
detach(), метод PyTorch, 157
device(), метод, 107
DeviceArray, тип, 306
Device mesh, 310
device_put(), метод, 185
Diffrax, библиотека, 459
Diffusers, библиотека, 438
диффузионный конвейер, 439
планировщик шума, 439
предварительно натренированная
модель, 439
Stable Diffusion, модель, 439
Directional derivative, 168
Dopamine, библиотека, 454
dot_general(), функция, 233
DSP, digital signal processing, 92
E
EasyLM, библиотека, 449
EffcientNetV2, 414
einops, библиотека, 454
Elegy, библиотека, 448
Equinox, библиотека, 446
equinox.Module, модуль библиотеки
Equinox, 387
Evaluation trace, 165
EvoJAX, библиотека, 455
Evosax, библиотека, 455
Предметный указатель
515
F
G
FedJAX, библиотека, 456
FIR-фильтр, 92
FIR, finite impulse response, 92
Fitness landscape, 66
Flax, 397
библиотека, 446
тренировка нейронной сети, 408
оптимизатор Optax, 408
управление состоянием, 415
BatchNorm params, коллекция, 416
FrozenDict, структура данных, 403
GradientTransformation,
интерфейс, 407
Linen API, 398
Module, абстракция, 399
FlaxAutoModelForCausalLM,
класс, 429, 431
FlaxAutoModelForSequenceClassificati
on, класс, 437
FlaxGPTJForCausalLM, класс, 429, 431
FlaxGPTNeoForCasualLM, класс, 429
flax.jax_utils.replicate(), функция, 440
flax.linen, модуль, 416
flax.linen.Module, класс, 400
FlaxStableDiffusionPipeline, класс, 440
flax.struct.dataclass, аннотация из
библиотеки Flax, 387
flax.struct.dataclass, класс, 423
flax.training.checkpoints, API, 422
flax.training.common_utils.shard(),
функция, 440
flax.training.train_state.TrainState,
класс, 408
flip(), функция, 90
fliplr(), функция, 344
Flower, библиотека, 456
fold_in(), функция, 358
Foolbox, библиотека, 457
from_output(), функция, 411
from_pretrained(), метод, 429
functools.partial, 189
functools.partial, аннотация, 220
functools.reduce(), функция
Python, 383
functorch, библиотека, 42
gather_from_model_output(),
функция, 411
Gaussian blur filter, 92
Gaussian noise, 90
GELU, Gaussian error linear units, 59
Gemma 2B, компактная языковая
модель, 427
Generator, генератор случайных
чисел, 351
Global device, 297
GlobalDeviceArray, тип, 306
GPT2-TokenizerFast, класс
токенизатора, 431
GPT-J-6B
авторегрессивная языковая
модель, 427
модель трансформера
естественного языка группы
EleutherAI, 27
GPU, графический процессор
(graphics processing unit), 105
grad()
трансформация, 142
функция, 420
функция трансформации, 39, 68,
149, 158, 247
Gradient descent, 65
GradientTransformation,
интерфейс, 407
gymnax, библиотека, 455
H
Haiku, библиотека, 447
Haliax, библиотека, 450, 453
has_aux=True, параметр, 420
hash(), функция Python, 359
hashlib, библиотека, 359
Hessian, 162
hessian(), функция, 163
Hessian matrix, 162
Hugging Face, 425
библиотека
диффузоров, 425
трансформеров, 425
516
Предметный указатель
auto classes, метод загрузки
моделей, 429
Hugging Face Transformers,
библиотека, 426
I
img_as_float(), функция, 90
init(), функция, 406
inline=True, параметр, 190
is_leaf, параметр, 380
Ivy, 76
библиотека, 447
J
jacfwd(), функция, 161, 168
Jacobian, 161
Jacobian matrix, 161
jacrev(), функция, 161
JAX
вычисление градиента, 149
вариант со многими
переменными, 161
производная более высокого
порядка, 158
jax.grad(), трансформация, 149
массив
изменение индекса, 119
индексирование за границами
массива, 120
неизменяемость, 117
отличия от NumPy, 117
сравнение
с NumPy, 37
с PyTorch, 42
с TensorFlow, 42
тип данных, 121
расширение, 124
bfloat16, 124
float16, 124
float64, 122
устройство
глобальное, 106
зафиксированные данные, 107
локальное, 106
незафиксированные данные, 107
экосистема, 35
модули, 36
AOT-компиляция, 211
Just After eXecution (происхождение
аббревиатуры), 28
MLIR/MHLO, 207
StableHLO, 208
XLA, 207
JAX2TF, библиотека, 452
jax.Array, тип, 306
JAXChem, библиотека, 458
jax.core.ClosedJaxpr, тип, 195
jax.core.Jaxpr, тип, 195
jax.core.Tracer, класс
трассировщика, 195
jax-cosmo, библиотека, 458
jax.custom_jvp(), функция, 175
jax.custom_vjp(), функция, 175
jax.debug.visualize_array_sharding(),
функция, 309
jax.default_device(), функция,
менеджер контекста, 108
jax_default_prng_impl, флаг
инициализации JAX, 361
jax.device_count(), функция, 106, 288
jax.device_get(), функция, 108
jax.device_put(), функция, 108, 112,
310, 319
jax.devices(), функция, 106, 107, 184,
321
jax.distributed.initialize(),
функция, 299
jax_enable_x64, переменная
конфигурации JAX, 123
jax.flatten_util.ravel_pytree(),
функция, 380, 382
JAX-Fluids, библиотека, 458
Jaxformer, библиотека, 450
jax.jit()
трансформация, 181
функция, 189, 196
параметр, 183
backend, 184
device, 184
jax.lax
модуль, 195
пакет, 233
Предметный указатель
структурированный базисный
элемент управления потоком
выполнения, 217
расширение типа, 129
jax.lax.cond(), базисный элемент, 203
jax.lax.fori_loop(), базисный
элемент, 202
jax.lax.map(), функция, работа
с pytree, 374
jax.lax.psum(), функция, 300
jax.lax.scan(), функция, 218
работа с pytree, 374
jax.lax.stop_gradient(), функция, 157
jax.lax.switch(), функция, 345, 347
jax.lax.with_sharding_constraint(),
функция, 319
jaxlib, пакет, 215
jaxlib.xla_extension.ArrayImpl, тип
массива, 264
JAXline, библиотека, 453
jax.local_device_count(), функция, 106,
262
jax.local_devices(), функция, 106, 321
jax.make_jaxpr(), функция, 193
JAX, M.D., библиотека, 457
jax.nn, библиотека, 58
jax.numpy, модуль, 100
JAX ONNX Runtime, 76
библиотека, 453
JAXopt, библиотека, 451
JAX_PLATFORMS, переменная
среды, 108
jax_platforms, флаг командной
строки, 108
Jaxpr, 190, 193
грамматическая форма, 194
промежуточное представление, 193
трассировка, 195
consts, атрибут, 195
JAX-Privacy, библиотека, 457
jax.process_index(), функция, 107, 116,
299
JaxPruner, библиотека, 457
jax.random.normal(), функция, 340,
342
jax.random.randint(), функция, 346
jax.random.split(), функция, 345
517
jax.scipy, модуль, 100
JaxSeq, библиотека, 451
jax.stages, API, 212
jax_threefry_partitionable, флаг
инициализации JAX, 362
jax.tree_util, пакет, 372
функции для работы с pytree, 376
jax.tree_util.register_pytree_node(),
функция, 390
jax.tree_util.tree_flatten(),
функция, 380
jax.tree_util.tree_leaves(),
функция, 372, 376
jax.tree_util.tree_map(), функция, 372,
376
jax.tree_util.tree_reduce(),
функция, 383
jax.tree_util.tree_structure(),
функция, 380
jax.tree_util.tree_transpose(),
функция, 385
jax.tree_util.tree_unflatten(),
функция, 380
jaxtyping, библиотека, 452
Jax-verify, библиотека, 452
jit=True, флаг, 442
JIT-компиляция, 31, 72, 179, 181
ограничения, 216
оптимизация, 189
сравнение с AOT-компиляцией, 213
статический аргумент, 186
функция
не являющаяся чистой, 216
чистая, 216
чистая функция, 190
Jaxpr, 217
Jaxpr, промежуточное
представление, 193
трассировка, 195
XLA, 204
jit(), функция трансформации, 39, 72,
179, 181, 307
оптимизация, 189
donate_argnums, аргумент, 189
keep_unused, аргумент, 190
работа с pytree, 375
JMP, библиотека, 453
518
Предметный указатель
jnp.broadcast_to(), функция, 378
jnp.convolve, функция, 127
jnp.sum(), функция, 103, 280, 283, 300
Jraph, библиотека, 455
Jumanji, библиотека, 455
jvp(), функция, 168
JVP, Jacobian-vector product, 168
j-Wave, библиотека, 458
K
keepdims=False, параметр, 317
Keras 3, 76
Keras, библиотека, 447
KFAC-JAX, библиотека, 451
kl_divergence, стандартная функция
потерь, 408
L
lax.associative_scan, 127
lax.cond, 127
lax.fori_loop, 127
lax.map, функция, 127
lax.scan, 127
базисный элемент, 218
lax.switch, 127
пример использования, 127
lax.while_loop, 127
Learning rate, 66
Levanter, библиотека, 450
Lineax, библиотека, 451
Linen, 398
Module API, 398
LLM (large language model), 27, 256,
425, 431
transformer decoder, 431
transformer encoder, 431
transformer encoder-decoder, 431
LLVM, 207
Local device, 297
local_device_count(), функция, 440
logsumexp(), функция, 68
loss(), функция потерь, 68, 292
Loss curve, 66
Loss function, 65
Lowering, 210
M
make_jaxpr(), трансформация, 195, 196
MaxText, библиотека, 449
Mctx, библиотека, 455
memory_analysis(), функция, 215
merge(), функция, 411
Mersenn Twister, 352
Mesh, менеджер контекста
сегментирования, 321
Mesh Transformer JAX,
библиотека, 428, 448
mesh_utils.create_device_mesh(),
функция, 310
metax, библиотека, 457
Metric, функциональный интерфейс
вычисления метрик, 410
metrics.Collection, интерфейс, 411
MHLO (мета-HLO), 207
MIMD (multiple instructions multiple
data – множественные инструкции,
множественные данные), 261
Minkowski distance, 187
MISD (multiple instructions single
data – множественные инструкции,
один поток данных), 261
MLIR (Multi-Level Intermediate
Representation), 207, 208
MLP-Mixer, 414
MLP, multilayer perceptron, 27, 54
model.apply(), функция, 420
model.init(), функция, 419
Model parallelism, 286
model.tabulate(), метод, 419
Moving average filter, 93
MPI, Message Passing Interface,
стандарт, 277
MPMD (multiple programs multiple
data – множественные программы,
множественные данные), 261
MSE, mean-squared error, 146
MT19937, 351
multi_transform(), функция, 407
mutable, параметр, 420
Предметный указатель
N
NamedSharding, тип объекта
сегментирования, 322
NAS, Neural Architecture Search, 59
NetKet, библиотека, 458
NumPy
генерация случайных чисел, 347
гарантия последовательной
равнозначности (sequential
equivalent guarantee), 350
начальное число, 349
состояние, 350
групповая (broadcasting)
операция, 226
массив
загрузка данных изображения, 84
обработка изображений, 83
отличия от JAX, 117
универсальная функция (ufunc), 226
einsum(), функция, 229
numpy.random, модуль, 348
numpy.random.normal(), функция, 340,
342
NumPyro, библиотека, 455
O
Objax, библиотека, 448
OpenLLaMa, фреймворк, 449
OpenXLA, 205
Optax
библиотека трансформации
градиентов, 405
библиотека утилит, 451
оптимизатор, 406
optax.apply_updates(), функция, 407
OptimiSM, библиотека, 458
Optimistix, библиотека, 451
optimizer.update(), функция, 406
Orbax
библиотека, 422, 452
объект контрольной точки
(checkpointer), 423
PyTreeCheckpointer, 423
orbax.checkpoint.CheckpointManager,
менеджер контрольных точек, 424
519
Oryx, библиотека, 455
os.urandom(), функция Python, 335
OTT-JAX, библиотека, 457
P
Pallas, расширение JAX, 110
partial(), функция, 189
PartitionSpec, объект (кортеж) для
описания сетки устройств, 322
Paxml (или Pax), библиотека, 449
PCG64DXSM, генератор случайных
чисел, 352
PCG, permuted congruential
generator, 352
pdot(), коллективная операция, 276
PennyLane, библиотека, 458
Penzai, библиотека, 454
Perceiver IO, 27
pickle, модуль, 104
Pipeline parallelism (pipelining), 328
PIX, библиотека, 456
pmap(), функция трансформации, 40,
255, 259, 263, 265, 290, 299, 377
вложенный вызов, 282
использование декоратора, 285
использование по аналогии
с vmap(), 260
распараллеливание
вычислений, 255
управление поведением, 268
большой массив, пример, 272
отображение осей тензоров, 268
in_axes, параметр, 269
out_axes, параметр, 271
axis_name, параметр, 279, 283
in_axes, параметр, 274, 377
out_axes, параметр, 291
pmax(), коллективная операция, 243
pmean(), коллективная операция, 243
pmin(), коллективная операция, 243
PositionalSharding, объект, 310
predict(), функция, 398
Predict function, 59
prepare_inputs(), функция, 440
PRNG (pseudo-random number
generator), 335
520
Предметный указатель
сохранение состояния в NumPy, 350
process_index, атрибут, 107, 116
Prompt engineering, 435
psum()
коллективная операция, 243, 292
функция, 283
PyTorch, вычисление градиента, 148
backward(), функция, 148
requires_grad=True, параметр, 148
pytree, 369
восстановление древовидной
формы, 380
иерархическая структура
данных, 369
лист, 372
преобразование в плоскую
структуру, 380
транспонирование, 384
узел, 372
специализированный,
создание, 387
None, 372
PyTreeDef, тип, 380
Q
quantumrandom, модуль Python, 336
qujax, библиотека, 458
R
RAG, retrieval-augmented
generation, 436
random_augmentation(), функция, 245,
363
random.get_state(), функция
NumPy, 350
random_noise, функция, 90
random.normal(), функция
NumPy, 348, 351, 357
random.PRNGkey(seed), функция
создания ключа, 342, 354
random.split(), функция, 356
random.SystemRandom(), генератор
случайных чисел в Python, 336
ravel_pytree(), функция, 382
Rax, библиотека, 457
ReLU, rectified linear unit, 58
replicate(axis=NUMBER), метод
сегментирования, 316
keepdims=False, параметр, 317
Residual neural network (ResNet), 413
revision, параметр, 430
RLax, библиотека, 454
RL, reinforcement learning, 454
RngBitGenerator (RNG), XLA PRNG, 361
RNG, random number generator, 335
rot90(), функция, 90
S
safetensors, пакет, 74
Saxml (Sax), библиотека, 75, 453
Scenic, библиотека, 456
scikit-image, библиотека, 85
SELU, scaled exponential linear
units, 59, 180
Sequential equivalent guarantee, 360
SGD, stochastic gradient descent, 65
ShapedArray, объекттрассировщик, 196
ShardedDeviceArray, тип, 306
sharding.replicate(0), параметр, 330
SIMD (single instruction multiple data −
одна инструкция, множественные
данные), 260
single_from_model_output(),
функция, 411
SISD (single instruction single data −
одна инструкция, один поток
данных), 260
Slicing, 89
softmax_cross_entropy, стандартная
функция потерь, 408
Sparse data, 125
split(), функция, 356
SPMD (single-program, multiple-data −
одна программа, множественные
данные), 255, 261
StableHLO, 208
static_argnames, параметр, 186
static_argnums, параметр, 186, 189
Предметный указатель
Stax, библиотека, 448
stop_gradient(), функция, 158
Supervision signal, 48
Swish, функция активации, 58
T
T5X, библиотека, 450
tabulate(), функция, 404
TensorFlow
вычисление градиента, 146
gradient(), функция, 146
Tensor parallelism (sharding), 328
Tensor sharding, 305
Test set, 48
TestU01, библиотека тестов
Crush-resistant, 352
TF2JAX, библиотека, 75, 453
The Pile, набор данных для
английского языка, 427
Threefry, PRNG в JAX, 352, 361
to_bf16(), метод, 431
to_fp16(), метод, 431
torch.func, комплект API, 42
TPU, тензорный процессор (tensor
processing unit), 105
вычисления, 112
подготовка к вычислениям, 114
TPU Pod slice, кластер, 300
Tracing, 190
Training set, 47
TrainState, объект (структура)
состояния, 420
TrainState.create(), функция, 408
Transformer Engine, библиотека, 453
Trax, библиотека, 447
tree_map(), функция, 380
Tree-math, библиотека, 453
tree_reduce(), функция, 383
tree_structure(), функция, 380
tree_transpose(), функция, 386
tree_unflatten(), функция, 380
U
UnshapedArray,
объект-трассировщик, 197
521
V
Validation set, 48
value_and_grad(), функция, 69, 155, 420
работа с pytree, 375
Veros, библиотека, 458
Vision Transformer (ViT), 27, 414
visu3d, библиотека, 456
vjp(), функция, 173
VJP, vector-Jacobian product, 173
vmap(), функция трансформации, 39,
63, 159, 230, 233, 243, 247, 251, 265,
311, 384, 386
работа с pytree, 375
управление поведением, 233
декоратор, 241
именованный аргумент, 238
коллективная операция, 242
управление осями выходного
массива, 237
управление осями массива, 233
axis_name, параметр, 242
in_axes, параметр, 234, 245
out_axes, параметр, 237
W
weak_type, свойство, 130
Wengert list, 165
Whisper, модель распознавания
речи, 426
X
XLA, 31, 38, 72, 179, 204
архитектура, 204
внешний компонент (frontend), 205
внутренний компонент
(backend), 205
HLO IR, специализированный
входной язык промежуточного
представления операций высокого
уровня, 206
xla_call, базисный элемент, 190
xla_force_host_platform_device_count,
флаг компилятора XLA, 258, 310
522
Предметный указатель
А
Автоматическая векторизация, 33,
223, 230
функции, 230
vmap(), трансформация, 230, 233
управление поведением, 233
Автоматическое
дифференцирование, 39, 133, 141
Алгоритм обратного
распространения, 171
Асинхронная диспетчеризация, 111
Б
Базисный элемент (primitive), 193
Базовый стохастический метод
градиентного спуска, 65
Большая языковая модель, 27, 256,
425, 431, 435
архитектура, 431
трансформер-декодер, 431
трансформер-кодировщик, 431
трансформер-кодировщикдекодер, 431
инженерия запросов, 435
инженерия составления текстовых
запросов, 435
каузальная
(причинно-следственная), 431
модель пространства
состояний, 431
предварительная тренировка, 436
рекуррентная нейронная сеть, 431
стратегия декодирования, 432
жадное декодирование (greedy
decoding), 432
лучевой поиск (beam search), 432
полиномиальная выборка
(multinomial sampling), 432
текстовый запрос (prompt), 431
токенизатор, 429, 432
точная настройка, 436
BERT-подобная, 426
BioBERT, 426
FinBERT, 426
FLAN-T5 (семейство T5), 426
GPT-подобная, 425, 431
BLOOM, 426
Gemma 2, семейство, 426
GPT-2, 426
GPT-J-6B, 426
GPT-Neo, 426
Llama, семейство, 425
PEGASUS (семейство T5), 426
temperature, параметр, 432
UMT5 (семейство T5), 426
В
Валидационный набор, 48
Вектор, 62
Векторизация
автоматическая, 230
вручную, 229
Венгерта список, 165
Вихрь Мерсенна, 352
Выражение в конечном виде, 136
Г
Гарантия последовательной
равнозначности, 360
Гауссов шум, 90
Генератор
псевдослучайных чисел, 335
период, 335
случайных чисел (ГСЧ), 76, 335
Генерация
ответа, дополненная результатами
поиска, 436
случайных чисел в NumPy, 347
Гессиан, 162
Гиперболический тангенс, 58
Глобальное устройство, 297
Градиент, 134
Д
Данные
древовидная структура pytree, 369
зафиксированные, 107
незафиксированные, 107
Дифференцирование, 136
Предметный указатель
автоматическое, 141
вручную, 136
символьное, 137
численное, 139
Диффузионная модель, 438
З
Задача классификации
изображений, 47
Замкнутая форма выражения, 136
И
Изображение, формат
NCHW, 87
NHWC, 87
Инженерия запросов, 435
Интерфейс (API)
высокого уровня, 126
jax.numpy, 126
низкого уровня, 126
jax.lax, 126
К
Категорийная функция потерь
перекрестной энтропии, 67
КИХ-фильтр, 92
Классификация
двоичная, 48
многозначная, 48
многоклассовая, 48
Коллективная операция, 242, 276
вложенное отображение, 282
all-gather, паттерн, 278
all-reduce, паттерн, 278
all-to-all, паттерн, 278
total exchange, 278
axis_index_groups, параметр, 280
broadcast, функция, 278
static_broadcasted_argnums, 278
gather, паттерн, 278
pmax(), 278
pmean(), 278
pmin(), 278
psum(), 278
523
reduce, паттерн, 278
scatter, паттерн, 278
Компиляция кода, 178
Компьютерное зрение, 414
Конволюционная (сверточная)
нейронная сеть, 413
Конгруэнтный генератор
с перестановками, 352
Контрольный сигнал, 48
Кривая потерь, 66
Л
Ландшафт отбора, 66
Линейный блок масштабируемой
экспоненциальной кривой, 180
Локальное устройство, 297
М
Массив, 81, 102
асинхронная диспетчеризация, 111
вычисление на TPU, 112
операция, связанная с аппаратным
устройством, 104
тип DeviceArray, 102
тип jax.Array, 102
тип numpy.ndarray, 82, 101
управление осями, 233
Array, тип JAX, 102
GlobalDeviceArray, тип, 102
jaxlib.xla_extension.ArrayImpl,
внутренний тип JAX, 101
NumPy-подобный API в JAX, 100
ShardedDeviceArray, тип, 102
Матрица, 62
Гессе, 162
Якоби, 161
Менеджер контекста, 108
Метод стохастического градиентного
спуска (SGD), 409
с импульсным параметром
возмущения, 409
Минковского расстояние, 187
Многослойный перцептрон, 54, 397
сегментирование тензоров, 325
Многоуровневый перцептрон, 27
524
Предметный указатель
Модель
загрузка (восстановление), 75
слияние (fusion), 75
точная настройка (fine-tuning), 75
развертывание, 74
сохранение, 74
Н
Направленный ациклический
граф, 148
Нейронная сеть
вес, 65
входной слой, 54
выходной слой, 54
конволюционная (сверточная), 414
остаточная нейросеть (ResNet),
вариант, 414
остаточная, 413
сверточная (конволюционная), 92
скрытый слой, 54
с прямой связью (с прямым
распространением сигнала), 54
тренировка с распараллеливанием
по данным, пример, 285
О
Обработка
изображений, 82
вырезка, 89
добавление шума, 90
массив NumPy, 83
загрузка данных, 84
сохранение тензора в файле, 98
фильтрация, 92
применение ядра фильтра
к изображению, 94
ядро, 92
цифровых сигналов, 92
Обучение с подкреплением, 454
Остаточная нейронная сеть, 413
П
Пакет, 53
Поиск минимума функции, 134
Произведение
вектора на якобиан, 173
якобиана на вектор, 168
Производная, 134
более высокого порядка, 158
наклонная, 168
по направлению, 168
Процедура градиентного спуска, 65
Р
Распараллеливание
вычислений, 254
коллективная операция, 276
конфигурация с несколькими
хостами, 296
модель программирования
мультиконтроллера
(multicontroller programming
model), 296
описание процесса по этапам, 256
pmap(), трансформация, 254, 259
конвейера, 328
модели, 286, 328
распараллеливание
конвейера, 328
распараллеливание тензора, 328
по данным, 285
восьмиканальное, 325
тренировка нейронной сети,
пример, 285
четырехканальное, 327
тензора, 328
двухканальное, 327
С
Свертка, 92
Свертывание констант, 186
Сегментирование тензора, 306
многослойный перцептрон, 325
основы, 307
сетка устройств, 310
двумерная, пример, 311
sharding.replicate(0), параметр, 330
Сетка устройств, 310
Сигмоида, 58
Предметный указатель
Скаляр, 62
Скорость обучения, 66
Снижение уровня (представления
кода), 210
Статический аргумент, 186
Схема прямого кодирования с одним
активным состоянием (one-hot
encoding), 68
Т
Тензор, 62, 81
«неровный» (ragged tensor), 126
плотный, 125
разреженный, 125
сегментирование, 305
многослойный перцептрон, 325
основы, 307
сетка устройств, 310
sharding.replicate(0),
параметр, 330
сохранение в файле, 99
форма, 86
nbytes, свойство, 88
ndim, свойство, 87
shape, свойство, 87
size, свойство, 88
Тестовый набор, 48
Токенизатор, 429
Трансформация, 39
автоматическая векторизация, 39
вычисление градиента, 39
дифференцирование, 39
компиляция кода, 39
распараллеливание кода, 40
JIT-компиляция, 39
Трансформер, 414
Трассировка, 190, 195
оценок, 165
525
Тренировочный набор, 47
Ф
Фильтр
размытия по Гауссу, 92, 94
скользящего среднего, 93
с конечной импульсной
характеристикой (КИХ-фильтр), 92
Функция
автоматическая векторизация, 230
активации, 58
декоратор, 241
именованный аргумент, 238
нулевая производная, 134
минимакс, 135
седловая точка, 135
точка максимума, 134
точка минимума, 135
потерь, 65, 67, 145
среднеквадратическая ошибка
модели, 146
прогноза, 59
прямого прохода, 59
чистая, 76
Х
Хост, 106
Ч
Чистая функция, 40
Я
Ядро фильтра, 92
применение к изображению, 94
Якобиан, 161
Книги издательства «ДМК Пресс»
можно купить оптом и в розницу на складе издательства по адресу:
Москва, ул. Электродная, д. 2, стр. 12, офис 7, тел. +7 (499) 322-19-38,
а также заказать на сайте www.dmkpress.com
с доставкой в любой регион РФ
Григорий Сапунов
Глубокое обучение с JAX
Главный редактор
Зам. главного редактора
Мовчан Д. А.
Яценков В. С.
editor@dmkpress.com
Перевод
Корректор
Верстка
Дизайн обложки
Снастин А. В.
Синяева Г. И.
Чаннова А. А.
Мовчан А. Г.
Гарнитура PT Serif. Печать цифровая.
Усл. печ. л. 42,74. Тираж 200 экз.
Веб-сайт издательства: www.dmkpress.com