Руководство по программированию для Kauldron от Google Research: конфигурации как обычные данные, компоненты, связанные строками, и тренер JAX, который можно прочитать от начала до конца
В этом руководстве мы реализуем Kauldron — библиотеку обучения на JAX от Google Research, которая позиционирует себя как оптимизированную для скорости исследований и модульности, — и воспринимаем эти два слова буквально, проверяя, что именно они дают на практике. Мы устанавливаем библиотеку, а затем посвящаем первую половину ноутбука трём механизмам, которые отличают Kauldron от набора Flax и Optax: konfig, превращающему эксперимент в дерево обычных словарей, которые можно преобразовать в JSON и обратно; kontext, соединяющему компоненты с помощью строковых путей ключей, благодаря чему функции потерь не нужно импортировать модель, оценки которой она вычисляет; и проверке форм во время выполнения, где именованные оси связываются между аргументами и сообщают, с чем они были сопоставлены, если что-то не совпадает. Затем мы создаём собственные функцию потерь и метрику в ожидаемом фреймворком формате, обучаем настоящий Trainer на синтетических данных в памяти без загрузок и ускорителя и отслеживаем внутренний слой модели, не изменяя саму модель. В завершение мы запускаем перебор из пяти вариантов, в котором каждый эксперимент отличается одной строкой конфигурации, а также позволяем процессу обучения сохранить контрольную точку и продолжить работу с места остановки.
import os
import sys
import json
import textwrap
import traceback
import subprocess
RESULTS = {}
def banner(title):
print("\n" + "=" * 78)
print(title)
print("=" * 78)
def section(name):
def wrap(fn):
def run(*a, **kw):
banner(name)
try:
out = fn(*a, **kw)
RESULTS[name] = out if isinstance(out, str) else "ok"
return out
except Exception as e:
RESULTS[name] = f"SKIPPED / FAILED -> {type(e).__name__}: {e}"
print(f"\n[!] {name} did not complete: {type(e).__name__}: {e}")
traceback.print_exc(limit=3)
return None
return run
return wrap
banner("0. Install Kauldron, and the one compatibility patch you need today")
subprocess.run([sys.executable, "-m", "pip", "install", "-q", "kauldron==1.4.2"], check=True)
import jax
from etils.enp import array_spec as _array_spec
# jax >= 0.10.1 moved `jax._src.prng`, but etils <= 1.14.0 still reaches for it whenever it
# inspects an array's dtype. Kauldron calls that code on every batch, so without this two-line
# patch a Trainer raises AttributeError before it finishes a single step. The replacement uses
# jax's own public dtype API and is a no-op on older jax.
if not hasattr(jax._src, "prng"):
_array_spec._is_jax_random_dtype = lambda dt: jax.dtypes.issubdtype(dt, jax.dtypes.prng_key)
import numpy as np
import optax
import flax
from flax import linen as nn
import kauldron
from kauldron import kd, konfig, kontext
from kauldron.typing import Float, typechecked
print(f" kauldron {kauldron.__version__} | jax {jax.__version__} | flax {flax.__version__}"
f" | optax {optax.__version__}")
print(f" devices: {jax.devices()}")
print("\n Kauldron's pitch is modularity: it is the glue, not the framework. Four pieces do the work:")
print(" konfig -> your experiment IS a Python call tree, and that tree is a plain dict")
print(" kontext -> parts are wired by string key paths, so they never import each other")
print(" ktyping -> Float['*b h w c'] checked at runtime, with named axes bound across args")
print(" kd.train -> Trainer: model + data + losses + metrics + optimizer, and nothing else")
print("\n Everything below runs on a CPU runtime with no dataset download: the data is synthetic.")
Мы устанавливаем Kauldron и применяем единственный патч совместимости, необходимый для текущего сочетания версий. В jax 0.10.1 был перемещён приватный модуль jax._src.prng, а etils вплоть до версии 1.14.0 всё ещё обращается к нему при проверке dtype массива — это путь выполнения, который Kauldron проходит на каждом батче. Без приведённой ниже двухстрочной замены, использующей собственный публичный API типов данных JAX и ничего не делающей в старых версиях, Trainer выдаёт AttributeError, не завершив ни одного шага. После этого мы импортируем четыре компонента, выполняющих основную работу: konfig для системы конфигурации, kontext для связывания, модуль typing для проверки форм во время выполнения и kd.train для самого Trainer. Всё дальнейшее выполняется на CPU, поскольку единственный набор данных в этом ноутбуке мы генерируем самостоятельно.
@section("1. A config is a call tree, and a call tree is a dict")
def config_is_a_dict():
with konfig.imports():
import optax as coptax # looks like optax, builds ConfigDict instead
cfg = coptax.adam(learning_rate=0.003)
print(f" cfg = {cfg}")
print(f" type = {type(cfg).__name__}")
print(f" __qualname__ = {cfg.__qualname__!r} <- the call, stored as data")
cfg.learning_rate = 1e-4 # configs are mutable
optimizer = konfig.resolve(cfg) # ...until you resolve them
print(f" after cfg.learning_rate = 1e-4 -> resolve() gives {type(optimizer).__name__}")
print("\n An arbitrarily complex optimizer is still just nested dicts:")
chain = coptax.chain(
coptax.clip_by_global_norm(1.0),
coptax.scale_by_adam(b2=0.99),
coptax.scale_by_learning_rate(0.003),
)
as_json = json.dumps(json.loads(chain.to_json()), indent=2)
print(textwrap.indent(as_json, " "))
rebuilt = konfig.resolve(konfig.ConfigDict(json.loads(chain.to_json())))
print(f" JSON -> ConfigDict -> resolve() -> {type(rebuilt).__name__}")
print(" optax has no idea konfig exists. No base class, no registry, no decorator.")
return f"optax.chain -> JSON -> {type(rebuilt).__name__}"
config_is_a_dict()
Начнём с konfig, поскольку на этом компоненте построена остальная библиотека. Внутри блока konfig.imports() импорт optax выглядит и работает как обычный optax, но вместо объектов создаёт конфигурацию, поэтому optax.adam(learning_rate=0.003) возвращает ConfigDict с квалифицированным именем вызова и его аргументами, а не оптимизатор. Эта конфигурация остаётся изменяемой, пока konfig.resolve не превратит её в настоящий объект, и, поскольку представляет собой всего лишь вложенные словари, сколь угодно сложный optax.chain можно сериализовать в JSON и восстановить в виде рабочего оптимизатора. Важно то, что optax пришлось сделать для поддержки этого механизма: ничего. В optax нет ни базового класса, ни реестра, ни декоратора, и то же самое относится к любой библиотеке, которую мы настраиваем таким способом.
@section("2. cfg.ref: change one number, everything downstream follows")
def config_references():
with konfig.imports():
import optax as coptax
from kauldron import kd as ckd
cfg = ckd.train.Trainer()
cfg.num_train_steps = 1000
cfg.schedules = {
"lr": coptax.warmup_cosine_decay_schedule(
init_value=0.0, peak_value=1e-3, warmup_steps=100,
decay_steps=cfg.ref.num_train_steps, # <- a reference, not the value 1000
)
}
at_1000 = konfig.resolve(cfg.schedules["lr"])
cfg.num_train_steps = 200 # one edit...
at_200 = konfig.resolve(cfg.schedules["lr"]) # ...and the schedule already knows
print(f" {'progress':>10s} {'lr @ 1000 steps':>18s} {'lr @ 200 steps':>16s}")
for frac in (0.1, 0.5, 0.9):
print(f" {frac:>9.0%} {float(at_1000(int(1000*frac))):>18.6f}"
f" {float(at_200(int(200*frac))):>16.6f}")
print("\n Without .ref the schedule would have frozen 1000 into itself, and a sweep over")
print(" num_train_steps would have silently trained on the wrong decay curve.")
return (f"lr at 90% of training: {float(at_1000(900)):.6f} (1000 steps)"
f" vs {float(at_200(180)):.6f} (200 steps)")
config_references()
Системы конфигурации обычно дают сбой, когда одно значение требуется в нескольких местах, и ответ Kauldron — cfg.ref. Мы указываем параметр decay_steps расписания warmup-cosine через cfg.ref.num_train_steps, а не через число 1000, затем меняем num_train_steps на 200 и снова разрешаем расписание. Кривая скорости обучения перестраивается, поскольку конфигурация хранит ссылку, а не копию значения. Без этой косвенной ссылки расписание сохранило бы в себе 1000, и перебор числа шагов обучения незаметно обучал бы каждый вариант по неправильной кривой затухания — именно такой ошибкой, которая выдаёт правдоподобное число и не сообщает о проблеме.
@section("3. kontext: parts are wired by string, so they never import each other")
def kontext_keys():
import dataclasses
ctx = {
"batch": {"image": np.zeros((4, 8, 8, 3)), "label": np.arange(4)},
"preds": {"logits": np.ones((4, 10)), "aux": [{"pos": np.zeros(3)}]},
}
print(" a context is just nested data; a key path reaches into it:")
for path in ["batch.image", "preds.logits", "preds.aux[0].pos"]:
print(f" {path:22s} -> {kontext.get_by_path(ctx, path).shape}")
try:
kontext.get_by_path(ctx, "batch.nope")
except KeyError as e:
print(f" {'batch.nope':22s} -> KeyError: {str(e)[:96]}...")
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class MeanGap:
preds: kontext.Key = kontext.REQUIRED
targets: kontext.Key = kontext.REQUIRED
def __call__(self, *, preds, targets):
return float(abs(np.asarray(preds).mean() - np.asarray(targets).mean()))
metric = MeanGap(preds="preds.logits", targets="batch.label")
kwargs = kontext.resolve_from_keyed_obj(ctx, metric)
print(f"\n MeanGap declared preds={metric.preds!r}, targets={metric.targets!r}")
print(f" resolved to kwargs: {{{', '.join(f'{k}: {v.shape}' for k, v in kwargs.items())}}}")
print(f" value = {metric(**kwargs)}")
print("\n MeanGap never imported the model and the model never heard of MeanGap. Point the")
print(" same metric at 'preds.aux[0].pos' and nothing but that string changes.")
return f"MeanGap(preds='preds.logits', targets='batch.label') = {metric(**kwargs)}"
kontext_keys()
kontext соединяет компоненты Kauldron, которые ничего не знают друг о друге. Контекст — это обычные вложенные данные, а путь ключа вроде batch.image или preds.aux[0].pos позволяет обратиться к ним, разрешая ключи словаря, атрибуты и индексы списков, и выдавая KeyError со списком реально доступных элементов, если путь не найден. Любой объект может объявить свои входы, аннотировав поля как kontext.Key, после чего resolve_from_keyed_obj извлечёт из контекста именно указанные пути и передаст их как именованные аргументы. Мы создаём таким способом небольшую метрику и направляем её на выходы модели: метрика не импортирует модель, модель ничего не знает о метрике, а перенаправление метрики на другой тензор требует изменить всего одну строку.
@section("4. ktyping: named axes, checked at runtime, bound across arguments")
def shape_checking():
@typechecked
def project(features: Float["*b n c"], weights: Float["c d"]) -> Float["*b n d"]:
return jax.numpy.einsum("...c,cd->...d", features, weights)
out = project(jax.numpy.zeros((2, 16, 8)), jax.numpy.zeros((8, 32)))
print(f" project(f32[2 16 8], f32[8 32]) -> {out.shape} c bound to 8, d bound to 32")
print("\n now break it: c is bound to 8 by the first argument, so 5 cannot also be c")
try:
project(jax.numpy.zeros((2, 16, 8)), jax.numpy.zeros((5, 32)))
except Exception as e:
print(textwrap.indent(str(e), " "))
print("\n 'Inferred Dims' is the part worth having: it reports what each axis name was already")
print(" bound to, so a mismatch names the axis instead of printing two anonymous shapes.")
return "mismatch named the axis: c already bound to 8, got 5"
shape_checking()
Модуль typing в Kauldron проверяет формы массивов во время выполнения с помощью именованных осей. Мы аннотируем функцию как Float[‘*b n c’] и Float[‘c d’], а декоратор связывает каждое имя оси при первом появлении и затем проверяет это соответствие во всех остальных местах сигнатуры, включая возвращаемое значение. Когда мы намеренно передаём несовместимый второй аргумент, сообщение об ошибке делает главное: вместе с фактическими формами оно выводит блок Inferred Dims, показывающий, что c уже было связано со значением 8, поэтому ошибка называет несовпавшую ось, вместо того чтобы оставлять нас один на один с двумя безымянными кортежами. В модели, где одновременно обрабатывается несколько тензоров, это разница между исправлением в одну строку и полноценной отладкой.
import dataclasses
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class LogCosh(kd.losses.Loss):
"""log(cosh(err)): quadratic near zero, linear in the tails. ~30 lines less than raw Flax."""
preds: kontext.Key = kontext.REQUIRED
targets: kontext.Key = kontext.REQUIRED
@typechecked
def get_values(self, preds: Float["*a"], targets: Float["*a"]) -> Float["*a"]:
return jax.numpy.log(jax.numpy.cosh(preds - targets))
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class WithinTol(kd.metrics.Metric):
"""Fraction of predictions landing within `tol` of the target, over every batch seen."""
preds: kontext.Key = kontext.REQUIRED
targets: kontext.Key = kontext.REQUIRED
tol: float = 0.25
@flax.struct.dataclass
class State(kd.metrics.AutoState):
# sum_field() marks a value that is ADDED when two states merge. Keeping the numerator
# and the denominator apart is what makes the pooled result exact.
n_hit: Float[""] = kd.metrics.sum_field(default=0.0)
n_total: Float[""] = kd.metrics.sum_field(default=0.0)
def compute(self) -> Float[""]:
# Return a jax scalar, like the built-in states do: the metric writer that
# `trainer.train()` logs through does not accept a bare numpy scalar.
total = jax.numpy.maximum(jax.numpy.asarray(self.n_total), 1.0)
return jax.numpy.asarray(self.n_hit) / total
@typechecked
def get_state(self, preds: Float["*a"], targets: Float["*a"]) -> "WithinTol.State":
hit = (jax.numpy.abs(preds - targets) < self.tol).astype("float32")
return self.State(n_hit=hit.sum(), n_total=jax.numpy.asarray(hit.size, "float32"))
@section("5. A custom loss and a custom metric, in the shape Kauldron expects")
def custom_loss_and_metric():
rng = np.random.default_rng(0)
p = jax.numpy.asarray(rng.normal(size=(8, 4)).astype("float32"))
t = jax.numpy.asarray(rng.normal(size=(8, 4)).astype("float32"))
loss = LogCosh(preds="preds.y", targets="batch.y")
print(f" {'LogCosh(preds, targets)':34s} {float(loss(preds=p, targets=t)):.6f}")
print(f" {'same loss, weight=0.5':34s} "
f"{float(LogCosh(preds='a', targets='b', weight=0.5)(preds=p, targets=t)):.6f} <- exactly half")
print(f" {'builtin kd.losses.L2':34s} {float(kd.losses.L2(preds='a', targets='b')(preds=p, targets=t)):.6f}")
print("\n A metric is not a number, it is a State that merges. Watch why that matters when")
print(" the last batch of an epoch is smaller than the rest:")
metric = WithinTol(preds="preds.y", targets="batch.y", tol=0.5)
big = metric.get_state(preds=p[:6], targets=t[:6])
small = metric.get_state(preds=p[6:], targets=t[6:])
merged = big.merge(small)
for label, st in [("batch of 6 rows", big), ("batch of 2 rows", small), ("big.merge(small)", merged)]:
print(f" {label:22s} {float(st.n_hit):>4.0f} / {float(st.n_total):>3.0f} = {float(st.compute()):.4f}")
naive = (float(big.compute()) + float(small.compute())) / 2
print(f" {'mean of the two rates':22s} {'':>4s} {'':>3s} = {naive:.4f} <- wrong, and quietly so")
print("\n sum_field() adds numerator and denominator separately, so the pooled value is exact")
print(" however the batches were sized. The same mechanism aggregates a metric across devices:")
print(" merge is associative, so the order the states arrive in never changes the answer.")
print("\n Subclass, annotate the keys, implement one method. The loss never sees a batch dict")
print(" and the metric never sees the model; the keys deliver exactly what was asked for.")
return (f"merged {float(merged.n_hit):.0f}/{float(merged.n_total):.0f} = {float(merged.compute()):.4f}"
f" vs {naive:.4f} from averaging the rates")
custom_loss_and_metric()
Мы создаём собственные функцию потерь и метрику ровно в том формате, который ожидает Kauldron: это замороженный dataclass с полями kontext.Key и одним методом. Функция потерь реализует get_values и возвращает массив поэлементных значений; фреймворк сам выполняет редукцию и обрабатывает аргумент weight, что мы подтверждаем проверкой: weight=0.5 ровно вдвое уменьшает результат. Метрика интереснее, поскольку метрика Kauldron возвращает не число, а объединяемое состояние State. Мы строим её на AutoState с двумя полями sum_field — числителем и знаменателем, — а затем объединяем батч из шести строк с батчем из двух строк. Итоговое значение является точным, тогда как среднее двух показателей по батчам заметно ошибочно. Именно это произошло бы с последним неполным батчем эпохи; поскольку merge ассоциативен, тот же механизм агрегирует метрику на нескольких устройствах независимо от порядка поступления результатов.
Теперь мы обучаем модель. Набор данных Kauldron — это любая вызываемая функция, возвращающая дерево массивов, поэтому kd.data.InMemoryPipeline превращает наши синтетические данные для регрессии в настоящий конвейер с пакетной обработкой и перемешиванием, причём ничего скачивать не нужно. Мы собираем Trainer из модели, этого конвейера, собственной функции потерь, собственной метрики и оптимизатора Optax, а затем напрямую запускаем его шаг обучения, чтобы вывести кривую потерь: примерно за секунду работы на CPU значение падает с 1,51 до 0,005 за триста шагов. Важны две детали. Шаг обучения не создаёт дополнительные выходы, пока мы явно не попросим об этом с помощью return_losses и return_metrics, поскольку их вычисление требует времени устройства. А столбец enc_norm считывается из interms.enc.__call__[0] — пути ключа к промежуточному выводу слоя Dense с именем enc. Поэтому для мониторинга внутренней активации достаточно одной строки конфигурации и не требуется изменять модель.
Этот шаг демонстрирует смысл всей архитектуры. Мы один раз описываем эксперимент как конфигурацию, а затем запускаем пять вариантов, каждый из которых отличается ровно одной строкой: две ширины модели и две настройки оптимизатора, включая полную замену Adam на SGD. Каждый вариант разрешается в новый Trainer и обучается в течение двухсот настоящих шагов, при этом ни один символ модели, функции потерь или цикла обучения не меняется. В ноутбуке это работает благодаря двум деталям konfig. Классы, определённые в ноутбуке, находятся в __main__, который нельзя подделать при нетерпеливом импорте, поэтому для них мы используем konfig.imports(lazy=True). Кроме того, имя, импортированное лениво без вызова, разрешается в сам объект, а не вызывается, благодаря чему функция загрузки передаётся конвейеру без изменений. В командной строке те же переопределения записываются как –cfg.model.hidden=128, поэтому перебор в Kauldron — это просто список таких строк.
Kauldron не позволяет конфигурации содержать разрешённый объект и сообщает об этом немедленно. Присваивание реального модуля Flax объекту ConfigDict сразу вызывает ошибку с предложением двух решений: обернуть импорт в konfig.imports() или использовать mock_modules в ноутбуке. Частично разрешённую конфигурацию нельзя сериализовать, сравнивать или переопределять из командной строки, поэтому система полностью запрещает такое состояние, вместо того чтобы позднее завершиться с трудноотслеживаемой ошибкой. Эта же дисциплина объясняет то, что мы видели на шаге 6: такие подчинённые объекты, как конвейеры и оценщики, по умолчанию получают seed и набор данных как ссылки на корневую конфигурацию. Эти ссылки заполняются, когда объект создаётся внутри Trainer, но не заполняются при самостоятельном создании объекта, поэтому мы явно передали конвейеру seed.
В завершение рассмотрим компоненты, превращающие скрипт обучения в полноценную задачу. Оценщик объявляется в четыре строки, поскольку наследует из корневой конфигурации модель, функции потерь и метрики и требует только собственного набора данных и расписания, определяющего частоту запуска. Checkpointer сохраняет состояние через заданный интервал шагов. Затем мы второй раз создаём тот же Trainer с тем же рабочим каталогом и просим выполнить больше шагов — индикатор прогресса начинается с 200, а не с 0: процесс нашёл контрольную точку и продолжил работу. Именно так ведёт себя система при перезапуске прерванной задачи, и это полезно один раз увидеть в ноутбуке, где весь процесс занимает секунду, а не впервые столкнуться с этим на кластере.
В сводке выводится однострочный результат, возвращённый каждым разделом, а затем указывается, что изучать дальше: два самодостаточных модуля, которые стоит прочитать отдельно, реальные конвейеры данных вместо синтетического, примеры конфигураций в репозитории и командная строка, превращающая перебор из шага 7 во флаги.
В заключение мы решили проверить два заявления Kauldron, а не просто повторять их, и оба подтвердились по простой причине. Модульность здесь — не уровень абстракции, а его отсутствие: конфигурация — это словарь, описывающий вызов Python, связывание — строка, указывающая путь через данные, а optax и наша собственная модель не потребовали ни одной строки кода, специфичного для Kauldron, чтобы участвовать в процессе. Отсюда следует и скорость исследований: мы увидели её на практике в переборе, где пять экспериментов потребовали пяти изменённых строк, и в проверке форм, которая называет несовпавшую ось вместо вывода двух форм и предоставления нам возможности самостоятельно искать причину. Компоненты, которые нам не понадобились, — наборы данных TensorFlow, Grain, шардирование и XManager — совершенно не мешали, пока Trainer работал на синтетических данных на CPU. В реальную работу мы перенесли бы не столько саму библиотеку, сколько три принципа: хранить конфигурацию как данные, связывать компоненты путями, а не импортами, и делать метрики объединяемыми состояниями — все три полезны даже в кодовой базе, которая никогда не будет использовать Kauldron.
Ознакомьтесь с ПОЛНЫМ КОДОМ здесь. Все заслуги принадлежат исследователю этого проекта. Также подпишитесь на нас в Twitter и не забудьте присоединиться к нашему ML-сообществу на Reddit с более чем 150 тыс. участников и подписаться на нашу рассылку. Постойте! Вы есть в Telegram? теперь вы также можете присоединиться к нам в Telegram.
Хотите сотрудничать с нами в продвижении вашего репозитория GitHub, страницы Hugging Face, релиза продукта, вебинара и т. д.? Свяжитесь с нами
Переведено автоматически с английского. Оригинал статьи — по ссылке ниже.