Обучение модели на графическом процессоре с использованием JAX, Optax и Fashion-MNIST.

1. Введение

Курс обучения Jax на GPU. Лабораторная работа 4: Создание простого цикла обучения на GPU.

В предыдущих практических занятиях вы проверили, что JAX видит графический процессор, изучили, как jax.jit отслеживает и компилирует функцию, и использовали профилировщик, чтобы увидеть, что на самом деле делает графический процессор. Теперь все эти элементы объединяются в наиболее распространенный рабочий процесс JAX: цикл обучения.

В этом примере вы обучаете небольшой многослойный перцептрон (MLP) на наборе данных Fashion-MNIST , реальном наборе данных для классификации изображений, содержащем 60 000 обучающих примеров и 10 000 тестовых примеров. Каждый пример представляет собой изображение предмета одежды в оттенках серого размером 28x28 пикселей. В итоге вы получите скомпилированный этап обучения Optax, достоверные показатели производительности и матрицу ошибок, которая свяжет эти показатели с реальными изображениями.

Что вы будете делать

  • Загрузите Fashion-MNIST и преобразуйте его в пакеты фиксированного размера, которые будут храниться на графическом процессоре.
  • Создайте небольшой многослойный перцептрон (MLP) в виде PyTree из массивов JAX и напишите скалярную функцию потерь.
  • Вычислите градиенты с помощью jax.grad и jax.value_and_grad , затем скомпилируйте шаг с помощью jax.jit
  • Замените написанный вручную алгоритм SGD на оптимизатор Optax AdamW и запустите короткий цикл обучения.
  • Измерьте пропускную способность в примерах/сек и сопоставьте ее с количеством токенов/сек.
  • Оцените модель, постройте матрицу ошибок и сравните вычисления для типов данных float32 и bfloat16

Что вам понадобится

  • Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
  • Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
  • Выполнение заданий 1–3 или работа в эквивалентной среде JAX GPU с поддержкой CUDA.
  • Для первой загрузки Fashion-MNIST из Pod потребуется доступ в интернет.

Примерное время выполнения: 60 минут .

Ментальная модель этапов обучения

Шаг обучения JAX — это чистая функция: она принимает массивы на вход, возвращает новые массивы на выходе и не изменяет старые параметры. В каждом из приведенных ниже фрагментов есть проверка для начинающих, которую можно применить, если что-то пойдет не так.

Кусок

Что это делает

Проверка для начинающих

params

Веса модели хранятся в виде дерева PyTree.

Та же древовидная структура, что и grads

batch

Изображения и подписи

Для предотвращения перекомпиляции формы остаются одинаковыми на каждом шаге.

loss_fn

Прямой пас плюс скалярная потеря

Для jax.grad требуется скалярная функция потерь.

jax.value_and_grad

Вычисляет функцию потерь и градиенты одновременно.

Градиенты соответствуют формам параметров.

optimizer.update

Преобразует градиенты в обновления.

Адам хранит состояние оптимизатора

optax.apply_updates

Формирует следующие параметры

Параметры неизменяемы, поэтому возвращаем новое дерево.

Эти фрагменты выполняются в фиксированном четырехэтапном цикле: от прямого шага до обновления градиентов и повтора.

2. Прежде чем начать

Выберите свой проект

В консоли Google Cloud выберите или создайте проект с включенной функцией выставления счетов.

Открытая облачная оболочка

Чтобы запустить сеанс Cloud Shell , нажмите кнопку «Активировать Cloud Shell» (значок терминала в правом верхнем углу консоли), а затем укажите в качестве исполнителя свой проект:

gcloud config set project <YOUR_PROJECT_ID>

Подготовка среды для работы с графическим процессором.

Выполните следующие команды в Cloud Shell . В Codelab 1 подробно объясняется каждая команда, включая требования к квоте GPU и зоне.

gcloud services enable \
  container.googleapis.com \
  compute.googleapis.com \
  iam.googleapis.com \
  cloudresourcemanager.googleapis.com \
  logging.googleapis.com \
  monitoring.googleapis.com

git clone https://github.com/Google-Cloud-AI/partner-ai-nvidia.git
cd partner-ai-nvidia/05-workshops/jax-on-gpu/terraform

cp terraform.tfvars.example terraform.tfvars

Отредактируйте файл terraform.tfvars и установите project_id = " " Затем выполните настройку кластера и разверните project_id = " " :

terraform init
terraform apply
$(terraform output -raw get_credentials_command)

cd ..
kubectl apply -f deploy/jupyter.yaml

terraform apply занимает около 10 минут. После завершения дождитесь завершения работы Pod и LoadBalancer, затем считайте одноразовый токен JupyterLab из лога Pod:

kubectl get pod jax-jupyter -w         # wait for Running, then Ctrl+C
kubectl get svc jax-jupyter-svc -w     # wait for EXTERNAL-IP, then Ctrl+C
kubectl logs jax-jupyter | grep -o 'token=[a-z0-9]*' | head -1

Откройте http:// :8884 , вставьте токен и создайте новый блокнот Python 3 в /workspace . Каждый блок кода из этого практического занятия помещается в ячейку этого блокнота.

Установите все необходимое для этого практического занятия.

!pip install --quiet optax matplotlib

Оба пакета обычно поставляются в контейнере NVIDIA JAX, поэтому эта команда pip, как правило, ничего не делает.

Настройте и проверьте графический процессор.

Импортируйте JAX, Optax и несколько вспомогательных программ. Эта ячейка также проверяет, что в качестве бэкенда по умолчанию используется графический процессор (GPU).

import gzip
import hashlib
import html
import math
import pathlib
import struct
import time
import urllib.request
from functools import partial

from IPython.display import HTML, display
import matplotlib.pyplot as plt
import numpy as np

import jax
import jax.numpy as jnp

try:
    import optax
except ModuleNotFoundError as exc:
    raise ModuleNotFoundError(
        "This lesson requires Optax. Install it with: pip install optax"
    ) from exc


devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
device = gpu_devices[0] if gpu_devices else None

print(f"JAX version:     {jax.__version__}")
print(f"Optax version:   {getattr(optax, '__version__', 'unknown')}")
print(f"Default backend: {jax.default_backend()}")
print(f"Devices:         {devices}")

assert gpu_devices, f"This lesson assumes a GPU backend. Available devices: {devices}"
print(f"Using GPU:       {device}")


CLASS_NAMES = np.array([
    "T-shirt/top",
    "Trouser",
    "Pullover",
    "Dress",
    "Coat",
    "Sandal",
    "Shirt",
    "Sneaker",
    "Bag",
    "Ankle boot",
])


def block_tree(tree):
    """Wait until a PyTree of JAX arrays is ready on device."""
    return jax.block_until_ready(tree)


def tree_l2_norm(tree):
    """L2 norm of all leaves in a PyTree treated as one long vector."""
    leaves = jax.tree_util.tree_leaves(tree)
    return jnp.sqrt(sum(jnp.sum(jnp.square(x)) for x in leaves))


def count_params(params):
    """Total number of scalar values across all leaves of a parameter PyTree."""
    return sum(x.size for x in jax.tree_util.tree_leaves(params))


def show_table(headers, rows, title=None, aligns=None):
    """Render rows as an HTML table; `aligns` is an optional per-column list of "left"/"right"/"center"."""
    aligns = aligns or ["left"] * len(headers)
    parts = ["<div style='font-family: system-ui; max-width: 980px;'>"]
    if title:
        parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    parts.append("<table style='border-collapse: collapse; width: 100%; font-size: 13px;'>")
    parts.append("<thead><tr>")
    for h, a in zip(headers, aligns):
        parts.append(
            f"<th style='text-align:{a}; border-bottom:1px solid #d0d7de; padding:6px;'>"
            f"{html.escape(str(h))}</th>"
        )
    parts.append("</tr></thead><tbody>")
    for row in rows:
        parts.append("<tr>")
        for cell, a in zip(row, aligns):
            parts.append(
                f"<td style='text-align:{a}; border-bottom:1px solid #eef1f4; padding:6px;'>"
                f"{html.escape(str(cell))}</td>"
            )
        parts.append("</tr>")
    parts.append("</tbody></table></div>")
    display(HTML("".join(parts)))


def show_bars(rows, title, unit="", lower_is_better=False):
    """Render (label, value) pairs as a horizontal bar chart in HTML. Scales bars to the largest value."""
    max_value = max(float(value) for _, value in rows) or 1.0
    color = "#1a7f37" if not lower_is_better else "#0969da"
    parts = ["<div style='font-family: Arial, sans-serif; max-width: 760px;'>"]
    parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    for label, value in rows:
        width = max(3, 100 * float(value) / max_value)
        parts.append(
            "<div style='display:grid; grid-template-columns: 190px 1fr 130px; gap: 8px; "
            "align-items:center; margin: 6px 0;'>"
            f"<div style='font-size:13px;'>{html.escape(str(label))}</div>"
            "<div style='background:#f6f8fa; border-radius:6px; overflow:hidden; height:22px;'>"
            f"<div style='height:22px; width:{width:.1f}%; background:{color};'></div></div>"
            f"<div style='font-size:13px; font-variant-numeric: tabular-nums;'>{float(value):,.1f} {html.escape(unit)}</div>"
            "</div>"
        )
    parts.append(
        f"<div style='font-size:12px; color:#57606a;'>"
        f"{'Lower' if lower_is_better else 'Higher'} is better.</div></div>"
    )
    display(HTML("".join(parts)))

Вы должны увидеть версию JAX, версию Optax, gpu в качестве бэкенда по умолчанию и список устройств CUDA, за которым следует графический процессор, который будет использоваться в остальной части практического задания. Вспомогательные функции show_table и show_bars отображают таблицы результатов и гистограммы, которые вы увидите на последующих шагах.

3. Загрузка и проверка Fashion-MNIST

Прежде чем приступать к обучению, вам необходимы данные на графическом процессоре в таком формате, который не будет меняться от шага к шагу. На этом этапе загружается набор данных Fashion-MNIST, анализируется он и преобразуется в пакеты фиксированного размера, которые хранятся на устройстве.

Скачать набор данных

Fashion-MNIST использует тот же формат файлов IDX, что и MNIST. Приведенные ниже вспомогательные функции загружают сжатые файлы, проверяют их контрольные суммы и преобразуют изображения и метки в массивы NumPy.

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

DATA_DIR = pathlib.Path.home() / ".cache" / "jax-course" / "fashion-mnist"
DATA_DIR.mkdir(parents=True, exist_ok=True)

FILES = {
    "train-images-idx3-ubyte.gz": "8d4fb7e6c68d591d4c3dfef9ec88bf0d",
    "train-labels-idx1-ubyte.gz": "25c81989df183df01b3e8a0aad5dffbe",
    "t10k-images-idx3-ubyte.gz": "bef4ecab320f06d8554ea6380940ec79",
    "t10k-labels-idx1-ubyte.gz": "bb300cfdad3c16e7a12a480ee83cd310",
}

PRIMARY_BASE_URL = "https://github.com/zalandoresearch/fashion-mnist/raw/master/data/fashion"

def md5sum(path):
    """Stream `path` in 1 MB chunks and return its MD5 hex digest."""
    digest = hashlib.md5()
    with open(path, "rb") as f:
        for chunk in iter(lambda: f.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def download_if_needed(filename, expected_md5):
    """Download `filename` if missing or its MD5 doesn't match."""
    path = DATA_DIR / filename
    if path.exists() and md5sum(path) == expected_md5:
        print(f"Using cached {filename}")
        return path

    urls = [f"{PRIMARY_BASE_URL}/{filename}"]
    last_error = None
    for url in urls:
        try:
            print(f"Downloading {filename} from {url}")
            urllib.request.urlretrieve(url, path)
            actual_md5 = md5sum(path)
            if actual_md5 != expected_md5:
                raise ValueError(f"MD5 mismatch: expected {expected_md5}, got {actual_md5}")
            return path
        except Exception as exc:
            last_error = exc
            if path.exists():
                path.unlink()
            print(f"  failed: {exc}")

    raise RuntimeError(f"Could not download {filename}") from last_error


def read_idx_images(path):
    """Parse the Fashion-MNIST IDX-3 image file at `path` and return a (N, rows, cols) uint8 array."""
    with gzip.open(path, "rb") as f:
        magic, num_images, rows, cols = struct.unpack(">IIII", f.read(16))
        assert magic == 2051, f"Unexpected image magic number {magic} in {path}"
        data = np.frombuffer(f.read(), dtype=np.uint8)
    return data.reshape(num_images, rows, cols)


def read_idx_labels(path):
    """Parse the IDX-1 label file at `path` and return a 1-D uint8 array of class indices."""
    with gzip.open(path, "rb") as f:
        magic, num_labels = struct.unpack(">II", f.read(8))
        assert magic == 2049, f"Unexpected label magic number {magic} in {path}"
        data = np.frombuffer(f.read(), dtype=np.uint8)
    return data.reshape(num_labels)


paths = {name: download_if_needed(name, checksum) for name, checksum in FILES.items()}

train_images = read_idx_images(paths["train-images-idx3-ubyte.gz"])
train_labels = read_idx_labels(paths["train-labels-idx1-ubyte.gz"])
test_images = read_idx_images(paths["t10k-images-idx3-ubyte.gz"])
test_labels = read_idx_labels(paths["t10k-labels-idx1-ubyte.gz"])

show_table(
    ["Split", "Images", "Image shape", "Labels"],
    [
        ("train", f"{len(train_images):,}", train_images.shape[1:], f"{len(train_labels):,}"),
        ("test", f"{len(test_images):,}", test_images.shape[1:], f"{len(test_labels):,}"),
    ],
    title="Fashion-MNIST loaded from IDX files",
)

В таблице должно быть указано 60 000 обучающих изображений и 10 000 тестовых изображений, каждое с формой (28, 28) и таким же количеством меток, как и изображения в обоих случаях.

Рассмотрим несколько примеров.

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

fig, axes = plt.subplots(2, 5, figsize=(10, 4))
for label, ax in enumerate(axes.flat):
    idx = np.flatnonzero(train_labels == label)[0]
    ax.imshow(train_images[idx], cmap="gray")
    ax.set_title(CLASS_NAMES[label], fontsize=10)
    ax.axis("off")
fig.suptitle("One Fashion-MNIST example per class")
fig.tight_layout()
plt.show()

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

Подготовьте пакеты данных для графического процессора фиксированного размера.

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

Ячейка нормализует пиксели до диапазона [0, 1] , преобразует каждое изображение размером 28x28 в вектор из 784 значений, перемешивает обучающий набор данных один раз и преобразует данные в пакеты фиксированного размера. Пакеты перемещаются на графический процессор один раз с помощью jax.device_put , поэтому цикл обучения индексирует только те массивы, которые уже находятся на устройстве.

TRAIN_EXAMPLES = 60_000
TEST_EXAMPLES = 10_000
BATCH_SIZE = 512
INPUT_DIM = 28 * 28
NUM_CLASSES = 10


def prepare_images(images):
    """Cast uint8 images to float32 in [0, 1] and flatten each one into a 1-D feature vector."""
    images = images.astype(np.float32) / 255.0
    return images.reshape(images.shape[0], -1)


rng = np.random.default_rng(0)
train_perm = rng.permutation(len(train_images))[:TRAIN_EXAMPLES]

x_train = prepare_images(train_images[train_perm])
y_train = train_labels[train_perm].astype(np.int32)
x_test = prepare_images(test_images[:TEST_EXAMPLES])
y_test = test_labels[:TEST_EXAMPLES].astype(np.int32)


def make_fixed_batches(x, y, batch_size):
    """Trim trailing examples that don't fill a batch, reshape, and move to device."""
    usable = (len(x) // batch_size) * batch_size
    x = x[:usable].reshape(usable // batch_size, batch_size, x.shape[-1])
    y = y[:usable].reshape(usable // batch_size, batch_size)
    return jax.device_put(jnp.asarray(x), device), jax.device_put(jnp.asarray(y), device)


x_train_batches, y_train_batches = make_fixed_batches(x_train, y_train, BATCH_SIZE)
x_test_batches, y_test_batches = make_fixed_batches(x_test, y_test, BATCH_SIZE)
first_batch = (x_train_batches[0], y_train_batches[0])

show_table(
    ["Array", "Shape", "Dtype", "Devices"],
    [
        ("x_train_batches", x_train_batches.shape, x_train_batches.dtype, x_train_batches.devices()),
        ("y_train_batches", y_train_batches.shape, y_train_batches.dtype, y_train_batches.devices()),
        ("x_test_batches", x_test_batches.shape, x_test_batches.dtype, x_test_batches.devices()),
        ("y_test_batches", y_test_batches.shape, y_test_batches.dtype, y_test_batches.devices()),
    ],
    title="Fixed-size batches on GPU",
)

В таблице каждый массив должен иметь float32 или int32 , каждая фигура должна заканчиваться на 784 для изображений и 512 для меток, а в столбце Devices должно отображаться устройство CUDA, выбранное вами в ячейке настроек.

4. Определите модель и выполните один шаг градиентного спуска.

При работе с пакетами данных на устройстве вам нужны две вещи: модель, которая преобразует пакет данных в логиты, и скалярная функция потерь, которую можно дифференцировать. На этом этапе создаются обе функции, а затем вручную выполняется один шаг градиентного спуска, чтобы вы могли точно увидеть, что возвращает jax.grad .

Определите модель и функцию потерь.

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

Прямой проход преобразует веса и активации в compute_dtype , затем преобразует логиты обратно в float32 перед обработкой функции потерь. В данный момент тип данных compute_dtype — float32 , но позже вы измените его на bfloat16 и измерите, какие изменения произойдут.

HIDDEN1 = 256
HIDDEN2 = 128
LEARNING_RATE = 3e-3


def init_mlp_params(key, input_dim=INPUT_DIM, hidden1=HIDDEN1, hidden2=HIDDEN2, num_classes=NUM_CLASSES):
    """Initialize a 3-layer MLP with He-style weight scaling and zero biases."""
    k1, k2, k3 = jax.random.split(key, 3)
    return {
        "w1": jax.random.normal(k1, (input_dim, hidden1), dtype=jnp.float32) * math.sqrt(2.0 / input_dim),
        "b1": jnp.zeros((hidden1,), dtype=jnp.float32),
        "w2": jax.random.normal(k2, (hidden1, hidden2), dtype=jnp.float32) * math.sqrt(2.0 / hidden1),
        "b2": jnp.zeros((hidden2,), dtype=jnp.float32),
        "w3": jax.random.normal(k3, (hidden2, num_classes), dtype=jnp.float32) * math.sqrt(2.0 / hidden2),
        "b3": jnp.zeros((num_classes,), dtype=jnp.float32),
    }


def mlp(params, x, compute_dtype=jnp.float32):
    """Forward pass: cast inputs/params to `compute_dtype`, two GELU hidden layers, then cast logits back to float32."""
    x = x.astype(compute_dtype)
    w1 = params["w1"].astype(compute_dtype)
    b1 = params["b1"].astype(compute_dtype)
    w2 = params["w2"].astype(compute_dtype)
    b2 = params["b2"].astype(compute_dtype)
    w3 = params["w3"].astype(compute_dtype)
    b3 = params["b3"].astype(compute_dtype)

    x = jax.nn.gelu(x @ w1 + b1)
    x = jax.nn.gelu(x @ w2 + b2)
    logits = x @ w3 + b3
    return logits.astype(jnp.float32)


def cross_entropy_loss(params, batch, compute_dtype=jnp.float32):
    """Scalar softmax cross-entropy loss."""
    x, y = batch
    logits = mlp(params, x, compute_dtype=compute_dtype)
    return optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()


def loss_with_metrics(params, batch, compute_dtype=jnp.float32):
    """Same loss, but also returns batch accuracy in an aux dict."""
    x, y = batch
    logits = mlp(params, x, compute_dtype=compute_dtype)
    loss = optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()
    accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == y)
    return loss, {"accuracy": accuracy}


params = init_mlp_params(jax.random.key(1))
params = jax.device_put(params, device)

rows = []
for name, value in params.items():
    rows.append((name, value.shape, value.dtype, value.devices()))
show_table(["Parameter", "Shape", "Dtype", "Devices"], rows, title=f"MLP parameters: {count_params(params):,} trainable values")

В таблице перечислены шесть листов — три матрицы весов и три вектора смещения — все float32 и все выполняются на графическом процессоре. В заголовке указано общее количество обучаемых значений.

Вычислите один шаг градиента вручную.

jax.grad и jax.value_and_grad по умолчанию различаются по первому аргументу. В данном случае первым аргументом является params , поэтому градиент имеет ту же древовидную структуру, что и словарь параметров.

Используйте jax.grad когда вам нужны только градиенты. Используйте jax.value_and_grad на этапе обучения, когда вам нужны функция потерь и градиенты из одного и того же прямого и обратного проходов.

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

def sgd_step(params, batch):
    """One un-jitted SGD update and returns (new_params, loss, grads)."""
    loss, grads = jax.value_and_grad(cross_entropy_loss)(params, batch)
    new_params = jax.tree.map(lambda p, g: p - LEARNING_RATE * g, params, grads)
    return new_params, loss, grads


grads_only = jax.grad(cross_entropy_loss)(params, first_batch)
loss_value, grads = jax.value_and_grad(cross_entropy_loss)(params, first_batch)
loss_value, grads, grads_only = block_tree((loss_value, grads, grads_only))

rows = []
for name in params:
    rows.append((name, params[name].shape, grads[name].shape, grads[name].dtype))
show_table(["Leaf", "Param shape", "Grad shape", "Grad dtype"], rows, title="Gradient tree matches the parameter tree")

grad_difference = tree_l2_norm(jax.tree.map(lambda a, b: a - b, grads, grads_only))

print(f"loss before update: {float(loss_value):.4f}")
print(f"gradient L2 norm:   {float(tree_l2_norm(grads)):.4f}")
print(f"grad vs value_and_grad difference: {float(grad_difference):.6f}")

params_after_one, loss_after_one, _ = sgd_step(params, first_batch)
params_after_one, loss_after_one = block_tree((params_after_one, loss_after_one))
print(f"loss used for one SGD update: {float(loss_after_one):.4f}")

В выходных данных необходимо проверить три момента. Каждая строка имеет Grad shape , идентичную Param shape , grad vs value_and_grad difference равна нулю или очень близка к нему, поскольку оба преобразования вычисляют одну и ту же производную, а последняя выведенная функция потерь — это функция потерь, вычисленная до обновления на том же пакете данных.

sgd_step возвращает новое дерево параметров, а не изменяет params , что на практике и означает "чистая функция".

5. Скомпилируйте этап обучения с помощью jax.jit.

Шаг SGD, написанный вручную, корректен, но это не тот способ, которым следует запускать цикл обучения на графическом процессоре. Без jit Python постоянно отправляет множество мелких операций. С jit JAX отслеживает весь шаг один раз, а XLA компилирует его в исполняемый файл для данной формы пакета и типа данных.

Первый вызов JIT-компилятора включает компиляцию. Приведенные ниже данные показывают время выполнения: сначала происходит предварительная прогревка, затем измеряется время выполнения с использованием кэша.

@jax.jit
def sgd_step_jit(params, batch):
    loss, grads = jax.value_and_grad(cross_entropy_loss)(params, batch)
    new_params = jax.tree.map(lambda p, g: p - LEARNING_RATE * g, params, grads)
    return new_params, loss


def batch_at(step):
    """Pick batch `step % num_batches`."""
    i = step % x_train_batches.shape[0]
    return x_train_batches[i], y_train_batches[i]


def time_loop(step_fn, params, steps):
    """Run `step_fn` for `steps` iterations and time it."""
    start = time.perf_counter()
    loss = None
    for step in range(steps):
        params, loss = step_fn(params, batch_at(step))
    params, loss = block_tree((params, loss))
    elapsed = time.perf_counter() - start
    return params, loss, elapsed


params_warm, loss_warm = sgd_step_jit(params, first_batch)
block_tree((params_warm, loss_warm))

EAGER_STEPS = 20
JIT_STEPS = 100

# sgd_step returns (params, loss, grads).
_, eager_loss, eager_elapsed = time_loop(lambda p, b: sgd_step(p, b)[:2], params, EAGER_STEPS)
_, jit_loss, jit_elapsed = time_loop(sgd_step_jit, params, JIT_STEPS)

eager_rate = EAGER_STEPS * BATCH_SIZE / eager_elapsed
jit_rate = JIT_STEPS * BATCH_SIZE / jit_elapsed

show_table(
    ["Mode", "Steps", "Final loss", "Elapsed seconds", "Examples/sec"],
    [
        ("Python dispatch", EAGER_STEPS, f"{float(eager_loss):.4f}", f"{eager_elapsed:.3f}", f"{eager_rate:,.0f}"),
        ("jitted step", JIT_STEPS, f"{float(jit_loss):.4f}", f"{jit_elapsed:.3f}", f"{jit_rate:,.0f}"),
    ],
    title="Cached training-step throughput",
    aligns=["left", "right", "right", "right", "right"],
)
show_bars([("Python dispatch", eager_rate), ("jitted step", jit_rate)], "Examples per second", "examples/s")

Вы должны увидеть две строки и диаграмму с двумя столбцами, при этом шаг компиляции с использованием JIT-редактора достигает более высокой производительности (примеров в секунду), чем диспетчеризация Python. Точное соотношение зависит от вашего графического процессора и формы пакета, но порядок имеет значение: скомпилированный шаг выполняет ту же работу с гораздо меньшими накладными расходами на каждую операцию.

6. Обучите программу с помощью оптимизатора Optax.

Для демонстрации градиентов достаточно чистого SGD, но в реальных циклах обучения необходимы импульс, затухание весов и расписание. На этом этапе заменяется Optax, запускается реальный цикл обучения и результат преобразуется в показатель пропускной способности.

Настройте оптимизатор Optax.

Optax , a gradient processing and optimization library for JAX, gives you composable optimizers such as SGD with momentum, Adam, AdamW, gradient clipping, and learning-rate schedules. An Optax optimizer has two important methods: optimizer.init(params) creates the optimizer state, such as Adam's momentum buffers, and optimizer.update(grads, opt_state, params) converts gradients into updates and returns the next optimizer state. Then optax.apply_updates(params, updates) returns the next parameter tree.

optimizer = optax.adamw(learning_rate=LEARNING_RATE, weight_decay=1e-4)
opt_state = optimizer.init(params)


@jax.jit
def train_step(params, opt_state, batch):
    (loss, metrics), grads = jax.value_and_grad(loss_with_metrics, has_aux=True)(params, batch)
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)
    metrics = {
        "loss": loss,
        "accuracy": metrics["accuracy"],
        "grad_norm": optax.global_norm(grads),
    }
    return params, opt_state, metrics


params_opt = init_mlp_params(jax.random.key(2))
params_opt = jax.device_put(params_opt, device)
opt_state = optimizer.init(params_opt)

params_opt, opt_state, metrics = train_step(params_opt, opt_state, first_batch)
params_opt, opt_state, metrics = block_tree((params_opt, opt_state, metrics))

show_table(
    ["Metric", "Value"],
    [
        ("loss", f"{float(metrics['loss']):.4f}"),
        ("accuracy", f"{100 * float(metrics['accuracy']):.1f}%"),
        ("gradient L2 norm", f"{float(metrics['grad_norm']):.4f}"),
    ],
    title="One compiled Optax training step",
    aligns=["left", "right"],
)

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

Запустите короткий цикл обучения.

Один шаг работает, поэтому его можно повторить. Цикл ниже вызывает тот же скомпилированный train_step для пакетов одинаковой формы. Это основной шаблон обучения JAX.

Обратите внимание на схему логирования. Метрики остаются на устройстве на протяжении всего шага и преобразуются в числа с плавающей запятой Python только каждые несколько шагов. Вывод значения функции потерь в Python на каждой итерации удобен, но он также синхронизирует хост с графическим процессором на каждой итерации.

def train_many_steps(params, opt_state, steps=400, log_every=25):
    """Run a sized training loop, log a metric snapshot every `log_every` steps, and return history plus final state."""
    history = []
    start = time.perf_counter()
    metrics = None

    for step in range(steps):
        params, opt_state, metrics = train_step(params, opt_state, batch_at(step))

        if step % log_every == 0 or step == steps - 1:
            metrics = block_tree(metrics)
            history.append(
                {
                    "step": step,
                    "loss": float(metrics["loss"]),
                    "accuracy": float(metrics["accuracy"]),
                    "grad_norm": float(metrics["grad_norm"]),
                }
            )

    params, opt_state, metrics = block_tree((params, opt_state, metrics))
    elapsed = time.perf_counter() - start
    return params, opt_state, history, elapsed, metrics


params_train = init_mlp_params(jax.random.key(3))
params_train = jax.device_put(params_train, device)
# Fresh optimizer state, paired with this fresh `params_train`.
opt_state = optimizer.init(params_train)

params_train, opt_state, _ = train_step(params_train, opt_state, first_batch)
block_tree((params_train, opt_state))

TRAIN_STEPS = 400
params_train, opt_state, history, elapsed, final_metrics = train_many_steps(
    params_train, opt_state, steps=TRAIN_STEPS, log_every=25
)
examples_per_sec = TRAIN_STEPS * BATCH_SIZE / elapsed

show_table(
    ["Step", "Loss", "Accuracy", "Grad norm"],
    [(h["step"], f"{h['loss']:.4f}", f"{100*h['accuracy']:.1f}%", f"{h['grad_norm']:.3f}") for h in history],
    title=f"Training metrics, {examples_per_sec:,.0f} examples/sec",
    aligns=["right", "right", "right", "right"],
)

steps = [h["step"] for h in history]
losses = [h["loss"] for h in history]
accuracies = [h["accuracy"] for h in history]

fig, ax1 = plt.subplots(figsize=(8, 4))
ax1.plot(steps, losses, marker="o", color="#0969da", label="loss")
ax1.set_xlabel("step")
ax1.set_ylabel("loss", color="#0969da")
ax1.tick_params(axis="y", labelcolor="#0969da")
ax1.grid(True, alpha=0.25)

ax2 = ax1.twinx()
ax2.plot(steps, accuracies, marker="s", color="#1a7f37", label="accuracy")
ax2.set_ylabel("batch accuracy", color="#1a7f37")
ax2.tick_params(axis="y", labelcolor="#1a7f37")
ax2.set_ylim(0.0, 1.0)

fig.suptitle("Fashion-MNIST training curve")
fig.tight_layout()
plt.show()

Вы должны увидеть таблицу метрик, в которой для каждого зарегистрированного шага отображается одна строка, а также двухкоординатный график, на котором кривая потерь снижается, а кривая пакетной точности возрастает на протяжении 400 шагов.

7. Оцените обученную модель.

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

Проведите оценку на отложенных тестовых данных.

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

Функция оценки использует jax.vmap для тестовых пакетов фиксированного размера. Это позволяет сделать код компактным и применять одну и ту же логику прогнозирования ко всем пакетам.

@jax.jit
def evaluate_batches(params, x_batches, y_batches):
    """Vmap `loss_with_metrics` over every (x, y) batch and return the mean loss and accuracy."""
    def eval_one_batch(x, y):
        loss, metrics = loss_with_metrics(params, (x, y))
        return loss, metrics["accuracy"]

    losses, accuracies = jax.vmap(eval_one_batch)(x_batches, y_batches)
    return {"loss": jnp.mean(losses), "accuracy": jnp.mean(accuracies)}


test_metrics = evaluate_batches(params_train, x_test_batches, y_test_batches)
test_metrics = block_tree(test_metrics)

show_table(
    ["Split", "Loss", "Accuracy", "Examples evaluated"],
    [
        ("train batch", f"{float(final_metrics['loss']):.4f}", f"{100 * float(final_metrics['accuracy']):.1f}%", BATCH_SIZE),
        ("test", f"{float(test_metrics['loss']):.4f}", f"{100 * float(test_metrics['accuracy']):.1f}%", int(np.prod(y_test_batches.shape))),
    ],
    title="Evaluation after the short training run",
    aligns=["left", "right", "right", "right"],
)

Значения в тестовой строке должны находиться в том же диапазоне, что и в итоговой обучающей выборке, а не быть значительно хуже. В столбце Examples evaluated отображается усеченный тестовый набор, а не все 10 000, поскольку из ранее созданных вами выборок были удалены оставшиеся 272 примера.

Визуализация прогнозов

Теперь перейдём к прогнозам. Зелёные заголовки обозначают правильные прогнозы. Красные заголовки — ошибочные.

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

@jax.jit
def predict(params, x, compute_dtype=jnp.float32):
    """Return the predicted class index (argmax of the logits) for each row of `x`."""
    logits = mlp(params, x, compute_dtype=compute_dtype)
    return jnp.argmax(logits, axis=-1)


rng = np.random.default_rng(7)
sample_count = 25
sample_indices = rng.choice(len(test_images), size=sample_count, replace=False)
sample_pixels = test_images[sample_indices]
sample_x = prepare_images(sample_pixels)
sample_y = test_labels[sample_indices].astype(np.int32)

sample_x_device = jax.device_put(jnp.asarray(sample_x), device)
sample_pred = np.asarray(block_tree(predict(params_train, sample_x_device)))

fig, axes = plt.subplots(5, 5, figsize=(10, 10))
for ax, image, true_label, pred_label in zip(axes.flat, sample_pixels, sample_y, sample_pred):
    correct = int(true_label) == int(pred_label)
    ax.imshow(image, cmap="gray")
    ax.set_title(
        f"pred: {CLASS_NAMES[pred_label]}\ntrue: {CLASS_NAMES[true_label]}",
        fontsize=9,
        color="#1a7f37" if correct else "#d1242f",
    )
    ax.axis("off")
fig.suptitle("Sample Fashion-MNIST predictions")
fig.tight_layout()
plt.show()

Вы должны увидеть сетку пять на пять, в которой преобладают зеленые заголовки, а также несколько красных. Обратите внимание на красные заголовки: большинство из них должны представлять собой ошибки в обозначении типов одежды, которые выглядят одинаково на изображении в оттенках серого размером 28x28.

Матрица ошибок

Используемые нами изображения представляют собой небольшую выборку. Матрица ошибок показывает, какие классы модель путает во всем тестовом наборе. Строки содержат истинные метки, а столбцы — предсказанные метки, причем диагональ означает, что модель, как правило, права.

@jax.jit
def predict_batches(params, x_batches):
    """Run `predict` over every batch with vmap; output shape is (num_batches, batch_size)."""
    return jax.vmap(lambda x: predict(params, x))(x_batches)


test_pred = np.asarray(block_tree(predict_batches(params_train, x_test_batches))).reshape(-1)
test_true = np.asarray(y_test_batches).reshape(-1)

confusion = np.zeros((NUM_CLASSES, NUM_CLASSES), dtype=np.int32)
np.add.at(confusion, (test_true, test_pred), 1)
confusion_percent = confusion / confusion.sum(axis=1, keepdims=True)

fig, ax = plt.subplots(figsize=(8, 7))
im = ax.imshow(confusion_percent, cmap="Blues", vmin=0.0, vmax=1.0)
ax.set_xticks(np.arange(NUM_CLASSES), CLASS_NAMES, rotation=45, ha="right")
ax.set_yticks(np.arange(NUM_CLASSES), CLASS_NAMES)
ax.set_xlabel("Predicted label")
ax.set_ylabel("True label")
ax.set_title("Fashion-MNIST confusion matrix")
fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04, label="fraction of true class")

for i in range(NUM_CLASSES):
    for j in range(NUM_CLASSES):
        value = confusion_percent[i, j]
        if value >= 0.08 or i == j:
            ax.text(j, i, f"{100 * value:.0f}%", ha="center", va="center", fontsize=8, color="white" if value > 0.45 else "black")

fig.tight_layout()
plt.show()

Вы должны увидеть темную диагональ, причем клетки, расположенные вне диагонали, будут преимущественно светлыми, а самые темные клетки, расположенные вне диагонали, будут сгруппированы между рубашками, футболками/топами, пуловерами и пальто.

8. Сравните вычисления float32 и bfloat16.

float32 используется по умолчанию для начальных циклов обучения. Тип bfloat16 использует меньше битов, поэтому может уменьшить объем используемой памяти и, возможно, использовать более быстрые аппаратные пути на поддерживаемых графических процессорах NVIDIA. Он не всегда быстрее для каждой модели, особенно для небольших, поэтому в целом: измените тип данных, прогрейте модель, проведите измерения.

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

@partial(jax.jit, static_argnames=("compute_dtype",))
def train_step_mixed(params, opt_state, batch, compute_dtype=jnp.float32):
    """Same as `train_step`, but `compute_dtype` is a static argument."""
    (loss, metrics), grads = jax.value_and_grad(loss_with_metrics, has_aux=True)(
        params, batch, compute_dtype=compute_dtype
    )
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)
    metrics = {
        "loss": loss,
        "accuracy": metrics["accuracy"],
        "grad_norm": optax.global_norm(grads),
    }
    return params, opt_state, metrics


def time_mixed_precision(compute_dtype, steps=100):
    """Fresh init + one warmup compile for this dtype, then time `steps` steps and report examples/sec."""
    params_mp = init_mlp_params(jax.random.key(10))
    params_mp = jax.device_put(params_mp, device)
    opt_state_mp = optimizer.init(params_mp)

    params_mp, opt_state_mp, metrics = train_step_mixed(
        params_mp, opt_state_mp, first_batch, compute_dtype=compute_dtype
    )
    block_tree((params_mp, opt_state_mp, metrics))

    start = time.perf_counter()
    for step in range(steps):
        params_mp, opt_state_mp, metrics = train_step_mixed(
            params_mp, opt_state_mp, batch_at(step), compute_dtype=compute_dtype
        )
    params_mp, opt_state_mp, metrics = block_tree((params_mp, opt_state_mp, metrics))
    elapsed = time.perf_counter() - start
    return {
        "dtype": str(jnp.dtype(compute_dtype)),
        "loss": float(metrics["loss"]),
        "accuracy": float(metrics["accuracy"]),
        "elapsed": elapsed,
        "examples_per_sec": steps * BATCH_SIZE / elapsed,
    }


MIXED_PRECISION_STEPS = 100
mp_results = [
    time_mixed_precision(jnp.float32, steps=MIXED_PRECISION_STEPS),
    time_mixed_precision(jnp.bfloat16, steps=MIXED_PRECISION_STEPS),
]

show_table(
    ["Compute dtype", "Final loss", "Accuracy", "Elapsed seconds", "Examples/sec"],
    [
        (
            r["dtype"],
            f"{r['loss']:.4f}",
            f"{100 * r['accuracy']:.1f}%",
            f"{r['elapsed']:.3f}",
            f"{r['examples_per_sec']:,.0f}",
        )
        for r in mp_results
    ],
    title="Mixed-precision timing after warmup",
    aligns=["left", "right", "right", "right", "right"],
)
show_bars([(r["dtype"], r["examples_per_sec"]) for r in mp_results], "Mixed-precision examples/sec", "examples/s")

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

Проверьте, что осталось на видеокарте.

Цикл обучения вернул новые PyTrees-объекты параметров и состояний оптимизатора. На графическом процессоре они по-прежнему представляют собой массивы JAX. Метрики преобразуются в значения Python только после их явного логирования.

В реальных конвейерах обработки входных данных пакеты часто начинаются на хосте. Это нормально, но передавайте весь пакет целиком за раз и избегайте преобразований, таких как np.asarray(loss) или float(loss) внутри активной части цикла.

def devices_in_tree(tree):
    """Set of devices that any JAX-array leaf in `tree` currently lives on."""
    devices = set()
    for leaf in jax.tree_util.tree_leaves(tree):
        if hasattr(leaf, "devices"):
            devices.update(leaf.devices())
    return devices


show_table(
    ["Object", "Where its arrays live"],
    [
        ("trained params", devices_in_tree(params_train)),
        ("optimizer state", devices_in_tree(opt_state)),
        ("training batches", x_train_batches.devices()),
        ("test batches", x_test_batches.devices()),
    ],
    title="Device placement check",
)

print(f"final training loss = {float(final_metrics['loss']):.4f}")

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

9. Уборка

Удалите рабочую нагрузку Jupyter, включая балансировщик нагрузки и постоянный том:

kubectl delete -f deploy/jupyter.yaml

Уничтожьте кластер, пул узлов, VPC и учетную запись службы:

cd terraform
terraform destroy

При появлении запроса введите yes , затем подтвердите, что ничего не осталось:

gcloud container clusters list
gcloud compute instances list

Оба поля должны быть пустыми для этого проекта. Если вы создали проект только для этой серии, вы можете удалить весь проект из консоли Cloud .

10. Поздравляем!

Вы перешли от отдельных концепций JAX к полному циклу обучения на графическом процессоре на реальном наборе данных и обучили многослойный перцептрон на наборе данных Fashion-MNIST от начала до конца.

Что вы узнали

  • Как хранить параметры модели в виде массива JAX в PyTree и обрабатывать пакеты данных фиксированного размера на графическом процессоре?
  • Как написать скалярную функцию потерь для целевой функции, которую вы хотите оптимизировать?
  • Когда использовать jax.grad (только для градиентов), а когда jax.value_and_grad (функция потерь и градиенты за один проход)?
  • Как встроить обновление Optax AdamW в один шаг обучения, скомпилированный с помощью jax.jit
  • Как честно измерить пропускную способность: сначала прогрейте систему, используйте block_until_ready() перед остановкой тактового генератора, а чтение данных хостом шлюза должно осуществляться с соблюдением условий логирования.
  • Как количество примеров в секунду соотносится с количеством токенов в секунду, и почему показатель «токенов в секунду» здесь является прогнозом, а не точным измерением.
  • Как визуализировать прогнозы и матрицу ошибок, чтобы метрики соответствовали реальным примерам.
  • Почему bfloat16 — полезный вариант для измерения, а не гарантированное ускорение, и почему параметры остаются в float32

Следующие шаги

  • Практический семинар 5: Ускорение механизма внимания на GPU с помощью cuDNN и TransformerEngine, где вы примените тот же принцип компиляции к ядрам механизма внимания, которые доминируют в обучении трансформеров.
  • Измените BATCH_SIZE и запустите заново. Большие пакеты часто улучшают использование графического процессора до тех пор, пока объем памяти не достигнет предела, а создание каждой новой фигуры требует одной перекомпиляции.
  • Измените значения HIDDEN1 и HIDDEN2 . Увеличение количества матричных умножений обычно приводит к большей нагрузке на графический процессор, поэтому понаблюдайте за тем, что происходит с количеством примеров в секунду.
  • Измените log_every в train_many_steps . Большее количество логов означает более быструю синхронизацию хоста, и это должно отражаться в показателе пропускной способности.

Справочная документация