Загрузка 0
ПОДЕЛИТЬСЯ

Мой блог

Листай вниз

Практическое руководство по Kauldron от Google Research: конфигурации как данные и JAX-тренер

Практическое руководство по Kauldron от Google Research: конфигурации как данные и JAX-тренер

Я подготовил подробный разбор Google Research Kauldron — современной библиотеки обучения на базе JAX, созданной для максимальной модульности и скорости исследований. В этом материале я проведу вас через ключевые механизмы фреймворка: от систем кон конфигурации konfig и связывания компонентов kontext до написания собственных лоссов, метрик и запуска экспериментов на CPU без лишних скачиваний.

Чтобы библиотека работала стабильно с актуальными версиями окружения, я начал с установки и применения важного патча совместимости. В релизе jax 0.10.1 внутренний модуль jax._src.prng был перемещен, однако версии etils вплоть до 1.14.0 продолжают обращаться к нему при проверке типов массивов. Поскольку Kauldron задействует этот путь на каждом батче, стандартный Trainer без небольшой двухстрочной правки выдаст ошибку AttributeError. Я использовал встроенный публичный API типов JAX, после чего успешно импортировал четыре базовых элемента фреймворка: konfig, kontext, модуль проверки размерностей и сам тренировочный модуль.

Система конфигураций konfig и динамические ссылки

Первый фундаментальный блок библиотеки — это konfig. Когда мы импортируем сторонние библиотеки вроде optax внутри специального блока konfig.imports(), на выходе получаются не готовые объекты, а структуры данных, которые выглядят и дополняются как оригинал, но формируют конфигурацию. Например, вызов optax.adam(learning_rate=0.003) создает объект ConfigDict, содержащий квалифицированное имя вызова и его аргументы. До момента вызова konfig.resolve() эти настройки остаются обычными вложенными словарями, которые можно легко сериализовать в JSON и восстановить обратно.

Реклама

Использование ссылок cfg.ref для зависимых параметров

Одной из главных проблем систем конфигураций является дублирование значений, которые должны синхронно меняться в разных частях кода. В Kauldron эта задача решена с помощью концепции cfg.ref. Если указать расписанию скорости обучения decay_steps ссылку на общее число шагов cfg.ref.num_train_steps вместо жестко зашитого числа, то при изменении num_train_steps с 1000 до 200 график затухания перестроится автоматически. Без такого косвенного обращения параметры зафиксировались бы на старых значениях, что привело бы к незаметному обучению по неверной кривой.

Связывание компонентов через строковые пути в kontext

Модуль kontext отвечает за объединение независимых частей архитектуры. Контекст в данном случае представляет собой обычные вложенные данные, а строковые пути вроде batch.image, preds.logits или preds.aux[0].pos позволяют извлекать нужные элементы независимо от того, являются они ключами словаря, атрибутами или индексами списка. Если путь отсутствует, фреймворк выбрасывает ошибку KeyError с подробным перечнем доступных ключей.

Реклама

Создание независимых метрик и потерь

Благодаря конктексту любые компоненты могут объявлять свои зависимости через аннотации kontext.Key. Я создал пользовательскую метрику MeanGap, которая получает на вход прогнозы и цели исключительно через строковые ключи. Сама метрика никогда не импортирует модель, а модель ничего не знает о метрике — их связывает исключительно текстовое описание в конфигурации. Перенаправление на другой тензор требует изменения всего одной строки.

Проверка размерностей в рантайме с помощью ktyping

Для контроля тензоров во время выполнения в Kauldron предусмотрен модуль ktyping, использующий именованные оси. Декоратор @typechecked проверяет сигнатуры функций, например, для операций эйнштейнова суммирования с размерностями Float["*b n c"] и Float["c d"]. Именованная ось фиксируется при первом появлении и жестко проверяется во всех остальных частях функции.

Диагностика несоответствий размеров

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

Пользовательские функции потерь и агрегируемые метрики

Чтобы запрограммировать собственные потери и метрики, достаточно создать замороженный датакласс (frozen dataclass) с полями kontext.Key. Я реализовал функцию потерь LogCosh, возвращающую массив для каждого элемента, где фреймворк самостоятельно управляет весами и редукцией. Проверка показала, что параметр weight=0.5 ровно вдвое уменьшает итоговое значение потерь.

Реклама

Состояние метрик и корректное объединение батчей

Метрики в Kauldron устроены сложнее: они возвращают не просто число, а специальный объект State, поддерживающий слияние (merge). Наследуя состояние от AutoState с полями sum_field() для числителя и знаменателя, мы можем корректно агрегировать результаты на последнем, неполном батче эпохи. Обычное усреднение долей дало бы погрешность, в то время как покомпонентное суммирование через ассоциативный метод merge гарантирует точный результат независимо от размера батчей и распределения по устройствам.

Сборка тренера, анализ внутренних слоев и запуск тренировки

Для проверки пайплайна я использовал kd.data.InMemoryPipeline, который превращает синтетические данные регрессии в полноценный генератор с батчами и перемешиванием без необходимости что-либо скачивать из интернета. Модель MLP была собрана вместе с оптимизатором Optax и кастомными метриками в единый объект kd.train.Trainer.

Мониторинг скрытых активаций без изменения модели

Запуск тренировки на CPU продемонстрировал быстрое падение функции потерь за 300 шагов. При этом особую ценность представляет мониторинг внутренних слоев: с помощью строкового пути interms.enc.__call__[0] я подключил метрику нормы к промежуточному выходу плотного слоя enc. Это позволило отслеживать внутренние активации сети без единого изменения в исходном коде архитектуры MLP.

Масштабирование экспериментов через циклы конфигураций

Главное достоинство такого проектирования раскрывается при создании сеток экспериментов (sweep). Базовый конфиг можно запускать в цикле с изменением ровно одного параметра за раз: варьируя ширину скрытых слоев модели (например, 4 или 128 нейронов), меняя скорость обучения или полностью заменяя оптимизатор Adam на SGD. При этом структура модели и цикл обучения остаются абсолютно нетронутыми.

Особенности ленивого импорта для интерактивной среды

Поскольку классы вроде MLP определены внутри блокнота (в пространстве имён __main__), стандартный импорт для них не подходит. Я применил konfig.imports(lazy=True), который передает ссылки на объекты как есть, не пытаясь вызвать их преждевременно. В командной строке те же самые переопределения задаются лаконичными флагами вроде --cfg.model.hidden=128.

Защитные механизмы консистентности конфигураций

Библиотека жестко пресекает попытки сохранить разрешенные объекты внутри конфигураций. Попытка напрямую присвоить готовый модуль Flax в ConfigDict вызывает немедленную ошибку с рекомендациями использовать правильный импорт или обертки. Это сделано намеренно: полуразрешенную конфигурацию невозможно корректно сериализовать, сравнить или переопределить через консоль, поэтому система исключает подобные состояния на этапе инициализации.

Оценка результатов, чекпоинты и возобновление работы

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

Автоматическое восстановление прерванного обучения

Повторный запуск тренера с тем же рабочим каталогом позволил автоматически обнаружить существующие чекпоинты и продолжить обучение с достигнутого шага (например, с 200 до 300), а не с нуля. Точно так же система ведет себя при восстановлении прерванных задач на кластерах. Изучив Kauldron, я убедился, что заложенные в него принципы — хранение конфигураций в виде чистых данных, связывание компонентов по путям и мерджинг состояний метрик — полезны даже за пределами экосистемы JAX.

01.