1. Введение

В этом практическом задании вы сохраняете тот же этап обучения JAX и распределяете работу между двумя графическими процессорами на вашем узле.
Существует несколько способов распределения обучения между устройствами: параллельная обработка данных , параллельная обработка тензоров или моделей, а также конвейерная параллельная обработка. В этом практическом занятии рассматривается параллельная обработка данных — простейший вариант. Каждый графический процессор получает свой собственный фрагмент пакета данных, запускает одну и ту же модель и один и тот же шаг обучения, а также вносит градиенты в общее обновление.
Важно то, что код этапа обучения практически не меняется. Меняется только способ размещения массивов на устройствах, а JAX обрабатывает распределенное выполнение.
Что вы будете делать
- Создайте
Mesh— логическую сетку из графических процессоров с именованными осями. - Обучающие пакеты шардирования и репликация параметров с использованием
NamedShardingиPartitionSpec - Проверьте полученное размещение с помощью
jax.debug.visualize_array_sharding - Запустите тот же этап обучения
jax.jitна сегментированных массивах и позвольте JAX распараллелить его. - Перепишите вычисление градиента с использованием
shard_mapдля явного управления каждым сегментом. - Проведите замеры производительности на одном графическом процессоре по сравнению с многопроцессорной системой и измените глобальный размер пакета данных.
Что вам понадобится
- Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
- Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
- В среде, где JAX видит два или более графических процессора . Данный практический пример останавливается на первой ячейке, если виден только один из них.
- Выполнение заданий с 1 по 5 или использование эквивалентной среды JAX GPU. В частности, задание 4 включает настройку кэша Fashion-MNIST и цикла обучения
optaxиспользуемого здесь.
Примерное время выполнения: 60 минут .
Как работает параллельное обучение на основе данных
Параллельная обработка данных — это простейшая стратегия для работы с несколькими графическими процессорами. Она разделяет пакет данных и дублирует модель. Вот как это происходит:
- Необходимо скопировать параметры модели таким образом, чтобы каждый графический процессор хранил полную копию весов.
- Разделите пакет данных по размеру пакета, и каждый графический процессор получит свой собственный фрагмент.
- Прямое и обратное преобразование данных выполняется на каждом графическом процессоре независимо, при этом каждый вычисляет градиенты на своем локальном участке.
- Чтобы усреднить градиенты по всем графическим процессорам, каждая копия получала одинаковое обновление.
- Обновление параметров на каждом графическом процессоре должно быть одинаковым, с использованием одних и тех же градиентов, что означает одинаковые новые веса.
Когда каждый графический процессор обладает достаточными локальными вычислительными ресурсами, параллельная обработка данных позволяет обрабатывать примеры per_device_batch * num_gpus с лишь незначительным увеличением времени выполнения по сравнению с обработкой одного графического процессора per_device_batch .
В этом и заключается преимущество. Недостаток в том, что каждый графический процессор должен хранить полную копию модели, поэтому параллельная обработка данных не помогает, когда сама модель слишком велика для одного устройства. Кроме того, она требует синхронизации градиентов между устройствами на каждом шаге, что может стать узким местом для очень больших моделей, небольших пакетов данных или более медленных межсоединений.
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 занимает около 12 минут. После завершения дождитесь завершения работы 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:// , вставьте токен и создайте новый блокнот Python 3 в /workspace . Каждый блок кода из этого практического занятия помещается в ячейку этого блокнота.
Установите все необходимое для этого практического занятия.
!pip install --quiet optax matplotlib
Настройте и проверьте графический процессор.
Этот код импортирует JAX вместе с тремя необходимыми примитивами шардинга: Mesh , PartitionSpec и NamedSharding , все из jax.sharding , и выводит то, что видит JAX. Строка assert len(gpu_devices) >= 2 является определяющим фактором для этого практического задания: всё, что следует за ней, предполагает наличие более одного устройства, поэтому, если виден только один графический процессор, код останавливается здесь, чтобы избежать ошибок на последующих этапах.
import os
os.environ["LD_LIBRARY_PATH"] = "/usr/local/nvidia/lib64:" + os.environ.get("LD_LIBRARY_PATH", "")
import gzip
import gc
import hashlib
import shutil
import subprocess
import html
import math
import pathlib
import struct
import time
import urllib.request
import warnings
from functools import partial
from IPython.display import HTML, display
import matplotlib.pyplot as plt
import numpy as np
warnings.filterwarnings("ignore", category=DeprecationWarning)
warnings.filterwarnings("ignore", message=".*ml_dtypes.*")
warnings.filterwarnings("ignore", message=".*JAX_PLATFORMS.*")
import jax
import jax.numpy as jnp
import optax
from jax.sharding import Mesh, PartitionSpec as P, NamedSharding
devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
NUM_DEVICES = len(gpu_devices)
print(f"JAX version: {jax.__version__}")
print(f"Default backend: {jax.default_backend()}")
print(f"GPU devices: {gpu_devices}")
print(f"GPU count: {NUM_DEVICES}")
assert len(gpu_devices) >= 2, (
f"This lab needs at least 2 GPUs. Found {len(gpu_devices)}. "
f"Available devices: {devices}"
)
def block_tree(tree):
"""Wait until a PyTree of JAX arrays is ready on device."""
return jax.block_until_ready(tree)
def drop_device_refs(*names, clear_compilation_cache=False):
"""Drop global references that may hold device buffers, then run cleanup."""
for name in names:
globals().pop(name, None)
gc.collect()
if clear_compilation_cache and hasattr(jax, "clear_caches"):
jax.clear_caches()
def show_table(headers, rows, title=None, aligns=None):
"""Render rows as an HTML table."""
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."""
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, gpu в качестве бэкенда по умолчанию и список из двух устройств CUDA с количеством графических процессоров, равным 2.
3. Загрузите набор данных Fashion-MNIST и определите ресурсоемкий многослойный перцептрон (MLP).
На этом этапе используется тот же набор данных Fashion-MNIST, что и в предыдущей лабораторной работе, но с вычислительно сложной моделью, чтобы эффект использования нескольких графических процессоров был более заметен. Загрузка набора данных и подготовка пакета на стороне хоста происходят один раз, до каких-либо измерений времени. В дальнейшем в этом практическом задании будет измеряться только этап обучения на скомпилированном графическом процессоре.
Загрузите и подготовьте данные.
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",
}
PRIMARY_BASE_URL = "https://github.com/zalandoresearch/fashion-mnist/raw/master/data/fashion"
def md5sum(path):
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):
path = DATA_DIR / filename
if path.exists() and md5sum(path) == expected_md5:
return path
for base in [PRIMARY_BASE_URL]:
try:
print(f"Downloading {filename}")
urllib.request.urlretrieve(f"{base}/{filename}", path)
if md5sum(path) != expected_md5:
raise ValueError("MD5 mismatch")
return path
except Exception:
if path.exists():
path.unlink()
raise RuntimeError(f"Could not download {filename}")
def read_idx_images(path):
with gzip.open(path, "rb") as f:
_, n, rows, cols = struct.unpack(">IIII", f.read(16))
return np.frombuffer(f.read(), dtype=np.uint8).reshape(n, rows, cols)
def read_idx_labels(path):
with gzip.open(path, "rb") as f:
_, n = struct.unpack(">II", f.read(8))
return np.frombuffer(f.read(), dtype=np.uint8).reshape(n)
paths = {name: download_if_needed(name, cs) for name, cs 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"])
# Shuffle once on the host
perm = np.random.default_rng(0).permutation(len(train_images))
x_train_all = (train_images[perm].astype(np.float32) / 255.0).reshape(len(train_images), -1)
y_train_all = train_labels[perm].astype(np.int32)
drop_device_refs("train_images", "train_labels", "perm")
Определите модель совместного использования блоков.
Модель многократно применяет один и тот же скрытый блок . Повторное использование одного и того же блока увеличивает объем локальных вычислений, выполняемых каждым графическим процессором за шаг, без увеличения количества значений градиента, которые необходимо синхронизировать между графическими процессорами.
В основе конструкции лежит это разделение. Стоимость синхронизации градиента зависит от количества параметров, а вычислительные затраты — от объема арифметических операций, выполняемых для каждого примера. BLOCK_REPEATS повышает второй порог, не затрагивая первый.
INPUT_DIM = 28 * 28
WIDTH = 1024
NUM_CLASSES = 10
BLOCK_REPEATS = 128
BLOCK_MIX = 0.10
LEARNING_RATE = 3e-4
PER_DEVICE_BATCH = 1024
GLOBAL_BATCH = PER_DEVICE_BATCH * NUM_DEVICES
NUM_TRAIN_BATCHES = 8
BENCHMARK_WARMUP = 4
BENCHMARK_STEPS = 15
BENCHMARK_REPEATS = 3
def init_params(seed=0):
rng = np.random.default_rng(seed)
def normal(shape, scale):
return rng.standard_normal(shape).astype(np.float32) * scale
return {
"w_in": normal((INPUT_DIM, WIDTH), math.sqrt(2.0 / INPUT_DIM)),
"b_in": np.zeros((WIDTH,), dtype=np.float32),
"w_block": normal((WIDTH, WIDTH), math.sqrt(2.0 / WIDTH)),
"b_block": np.zeros((WIDTH,), dtype=np.float32),
"w_out": normal((WIDTH, NUM_CLASSES), math.sqrt(2.0 / WIDTH)),
"b_out": np.zeros((NUM_CLASSES,), dtype=np.float32),
}
def make_fashion_batches(batch_size, num_batches=NUM_TRAIN_BATCHES):
needed = batch_size * num_batches
if needed > len(x_train_all):
raise ValueError(
f"Need {needed:,} examples, but Fashion-MNIST has {len(x_train_all):,}."
)
x = x_train_all[:needed].reshape(num_batches, batch_size, INPUT_DIM)
y = y_train_all[:needed].reshape(num_batches, batch_size)
return x, y
def model(params, x):
h = jax.nn.gelu(x @ params["w_in"] + params["b_in"])
def block(h, _):
z = jax.nn.gelu(h @ params["w_block"] + params["b_block"])
h = (1.0 - BLOCK_MIX) * h + BLOCK_MIX * z
return h, None
h, _ = jax.lax.scan(block, h, xs=None, length=BLOCK_REPEATS)
return h @ params["w_out"] + params["b_out"]
def loss_with_metrics(params, batch):
x, y = batch
logits = model(params, x)
loss = optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()
accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == y)
return loss, {"accuracy": accuracy}
optimizer = optax.adamw(learning_rate=LEARNING_RATE, weight_decay=1e-4)
param_template = init_params(seed=1)
PARAM_COUNT = sum(x.size for x in param_template.values())
GRADIENT_MB = PARAM_COUNT * np.dtype(np.float32).itemsize / 1e6
drop_device_refs("param_template")
show_table(
["", "Value"],
[
("Dataset", f"Fashion-MNIST train ({len(x_train_all):,} examples)"),
("Input shape", "28 x 28 grayscale, flattened to 784"),
("Model", f"shared-block MLP, width={WIDTH}, repeats={BLOCK_REPEATS}"),
("Parameters", f"{PARAM_COUNT:,}"),
("Gradient size", f"{GRADIENT_MB:.1f} MB per step"),
("Per-GPU batch", PER_DEVICE_BATCH),
("Global batch on all GPUs", GLOBAL_BATCH),
("Benchmark", f"median of {BENCHMARK_REPEATS} x {BENCHMARK_STEPS} steps"),
],
title="Fashion-MNIST compute-heavy workload",
)
Вы должны увидеть сводную таблицу, описывающую рабочую нагрузку, включая количество параметров, размер градиентов, которые будут синхронизироваться на каждом шаге, а также размеры пакетов для каждого графического процессора и общий размер пакета. На узле с двумя графическими процессорами общий размер пакета вдвое превышает размер пакета для каждого графического процессора.
4. Измерьте базовый показатель производительности с использованием одной видеокарты.
Перед добавлением второй видеокарты вам потребуется число для сравнения. Базовый вариант выполняется на одной видеокарте с размером пакета PER_DEVICE_BATCH — это точное значение объема работы, которое каждая видеокарта получит в многопроцессорном режиме.
Для начала закрепите данные и параметры на одном устройстве:
single_device = gpu_devices[0]
x_batches_1gpu, y_batches_1gpu = make_fashion_batches(PER_DEVICE_BATCH)
x_batches_1gpu = jax.device_put(x_batches_1gpu, single_device)
y_batches_1gpu = jax.device_put(y_batches_1gpu, single_device)
params_1gpu = jax.device_put(init_params(seed=1), single_device)
opt_state_1gpu = optimizer.init(params_1gpu)
Определите этап обучения и эталонный показатель.
train_step — это обычный шаг для одной видеокарты, включающий обработку значений и градиентов, обновление optax и новые параметры. При этом отсутствуют устройства и шардинг, которые являются частью данного практического задания.
benchmark_training сначала выполняет разминку, поэтому компиляция не учитывается, затем измеряет время выполнения трех повторений по пятнадцать шагов и сообщает медианное значение. block_tree обеспечивает справедливое измерение времени при асинхронной диспетчеризации JAX.
@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)
return params, opt_state, {"loss": loss, "accuracy": metrics["accuracy"]}
def benchmark_training(
step_fn,
params,
opt_state,
x_batches,
y_batches,
warmup=BENCHMARK_WARMUP,
steps=BENCHMARK_STEPS,
repeats=BENCHMARK_REPEATS,
):
"""Warm up, then report median steady-state throughput."""
num_batches = x_batches.shape[0]
for i in range(warmup):
batch = (x_batches[i % num_batches], y_batches[i % num_batches])
params, opt_state, _ = step_fn(params, opt_state, batch)
block_tree((params, opt_state))
batch_size = x_batches.shape[1]
timings = []
metrics = None
for repeat in range(repeats):
start = time.perf_counter()
for i in range(steps):
batch_index = (repeat * steps + i) % num_batches
batch = (x_batches[batch_index], y_batches[batch_index])
params, opt_state, metrics = step_fn(params, opt_state, batch)
params, opt_state, metrics = block_tree((params, opt_state, metrics))
timings.append(time.perf_counter() - start)
elapsed = float(np.median(timings))
return {
"examples_per_sec": steps * batch_size / elapsed,
"ms_per_step": 1000 * elapsed / steps,
"final_loss": float(metrics["loss"]),
"final_accuracy": float(metrics["accuracy"]),
}
Проведите базовое тестирование.
Последние две строки кода ниже вызывают функцию drop_device_refs , которая повторяется на протяжении всего этого практического занятия: каждый запуск выделяет параметры, состояние оптимизатора и пакеты на графическом процессоре, и эти буферы остаются активными до тех пор, пока на них ссылается глобальная переменная Python. Удаление имен и запуск gc.collect() освобождает память устройства до того, как следующий запуск выделит свою собственную, поэтому последующий шаг не завершится ошибкой нехватки памяти, вызванной уже завершенным запуском.
result_1gpu = benchmark_training(
train_step,
params_1gpu,
opt_state_1gpu,
x_batches_1gpu,
y_batches_1gpu,
)
show_table(
["Metric", "Value"],
[
("GPUs used", "1"),
("Batch per step", PER_DEVICE_BATCH),
("Throughput", f"{result_1gpu['examples_per_sec']:,.0f} examples/sec"),
("Step time", f"{result_1gpu['ms_per_step']:.2f} ms"),
("Final loss", f"{result_1gpu['final_loss']:.4f}"),
("Final accuracy", f"{100 * result_1gpu['final_accuracy']:.1f}%"),
],
title="Single-GPU baseline",
)
# Keep scalar timing results, but free device buffers from the single-GPU run.
drop_device_refs(
"params_1gpu",
"opt_state_1gpu",
"x_batches_1gpu",
"y_batches_1gpu",
)
Выполнение кода занимает некоторое время, поскольку он компилирует шаг, выполняет предварительную подготовку, а затем запускает 45 шагов с заданным временем. В результате вы должны получить таблицу базовых показателей для одной видеокарты, в которой будет указана пропускная способность в примерах в секунду, время выполнения шага в миллисекундах, а также значения потерь и точности с последнего шага. result_1gpu сохраняется после очистки, поскольку он содержит обычные числа с плавающей запятой Python, а не массивы устройств.
5. Создайте сетку устройств.
Mesh отображает физические графические процессоры на логическую сетку с именованными осями. Для параллельной обработки данных создается одномерная сетка, в которой все графические процессоры расположены вдоль одной 'data' .
Названия осей на выбор
В этом практическом задании ось сетки называется 'data' потому что она разделяет пакет данных на части для параллельной обработки. В более крупных моделях для описания других видов параллелизма можно использовать такие названия, как 'model' , 'tensor' , 'pipeline' или 'fsdp' . Для двумерной сетки может использоваться ('data', 'model') , где одна ось разделяет пакет данных, а другая — веса или активации модели.
Эти имена не имеют особого значения для JAX. Они обретают смысл только благодаря PartitionSpec и коллективным переменным, которые на них ссылаются.
На практике эти три примитива используются вместе: PartitionSpec описывает структуру, NamedSharding прикрепляет эту структуру к сети устройств, а jax.device_put перемещает массив в эту структуру.
mesh = Mesh(np.array(gpu_devices), ("data",))
show_table(
["", "Value"],
[
("Mesh shape", str(mesh.shape)),
("Axis names", str(mesh.axis_names)),
("Devices", ", ".join(str(d) for d in mesh.devices.flat)),
],
title="Device mesh",
)
В таблице должна быть одна ось с именем data , размер которой равен количеству ваших графических процессоров, и должны быть указаны оба устройства CUDA.
Проверьте топологию графического процессора.
Параллельное обучение данных уменьшает градиенты на каждом шаге, поэтому путь между двумя графическими процессорами проходит непосредственно по критическому пути. Пути NVLink , которые nvidia-smi обозначает как NV* , гораздо лучше подходят для этой цели, чем пути PHB , проходящие через мост хоста и PCIe.
if shutil.which("nvidia-smi"):
topo = subprocess.run(
["nvidia-smi", "topo", "-m"],
check=False,
text=True,
capture_output=True,
)
print(topo.stdout or topo.stderr)
else:
print("nvidia-smi is not available in this environment.")
Вы должны увидеть матрицу с одной строкой и одним столбцом для каждого графического процессора. На g2-standard-24 два L4-блока соединены через PCIe, поэтому ожидайте, что код GPU0 -to- GPU1 сообщит путь класса PHB . Именно это выдает данный тип машины, это не ошибка конфигурации, а ограничение масштабируемости параллельного обучения данных, и это объясняет любой результат, полученный на этапе сравнения.
6. Разделите данные на сегменты и продублируйте параметры.
В параллельном обучении данных есть ровно два варианта размещения:
- Данные разделяются по размерности пакета, и каждый графический процессор получает свой собственный фрагмент пакета.
- Параметры реплицируются, при этом каждый графический процессор хранит полную копию, поэтому прямой проход выполняется идентично.
PartitionSpec('data', None) разделяет первое измерение по оси сетки 'data' и дублирует второе измерение. PartitionSpec() без аргументов просто дублирует все.
В пакетных массивах используется ведущее измерение, поскольку make_fashion_batches возвращает все обучающие пакеты, сгруппированные вместе. Именно поэтому используется P(None, "data", None) , чтобы оставить измерение 0 целым, распределить примеры в измерении 1 по графическим процессорам и продублировать признаки.
batch_data_sharding = NamedSharding(mesh, P("data", None))
batch_label_sharding = NamedSharding(mesh, P("data"))
all_data_sharding = NamedSharding(mesh, P(None, "data", None))
all_label_sharding = NamedSharding(mesh, P(None, "data"))
replicated = NamedSharding(mesh, P())
x_batches_multi, y_batches_multi = make_fashion_batches(GLOBAL_BATCH)
x_batches_multi = jax.device_put(x_batches_multi, all_data_sharding)
y_batches_multi = jax.device_put(y_batches_multi, all_label_sharding)
params_multi = jax.device_put(init_params(seed=1), replicated)
opt_state_multi = optimizer.init(params_multi)
print(
f"Global batch: {GLOBAL_BATCH} examples "
f"({PER_DEVICE_BATCH} per GPU x {NUM_DEVICES} GPUs)"
)
print(f"Training batches shape: {x_batches_multi.shape}")
print()
print("One training batch: sharded along the batch dimension")
jax.debug.visualize_array_sharding(x_batches_multi[0])
print()
print("Weight w_block: replicated on all GPUs")
jax.debug.visualize_array_sharding(params_multi["w_block"])
Функция jax.debug.visualize_array_sharding выводит текстовую сетку, показывающую, какой графический процессор отвечает за какую часть массива. На узле с двумя графическими процессорами первые строки вывода должны выглядеть так:
Global batch: 2048 examples (1024 per GPU x 2 GPUs) Training batches shape: (8, 2048, 784)
Ниже вы должны увидеть пакет данных, представленный в виде двух блоков, расположенных друг над другом, по одному с меткой для каждого графического процессора, а также w_block представленный в виде единого блока с аннотациями, указывающими на оба графических процессора, сегментированные данные и реплицированные веса.
7. Выполните тот же шаг с использованием JIT-компиляции на сегментированных массивах.
Код этапа обучения остается неизменным. Он использует тот же скомпилированный train_step что и для базового варианта с одной видеокартой, работающего с сегментированными входными данными.
Когда JAX обнаруживает, что пакет данных распределен между графическими процессорами и параметры реплицируются, он автоматически выполняет следующие действия:
- Выполняет прямой проход по фрагменту данных каждого графического процессора.
- Вычисляет градиенты для каждого сегмента.
- Вставляет алгоритм редукции для усреднения градиентов по всем графическим процессорам.
- Параметры обновляются одинаково на каждом графическом процессоре.
Вы не пишете никакого кода, отвечающего за обмен данными. Параллелизм обеспечивается исключительно за счет расположения массивов.
result_multi = benchmark_training(
train_step,
params_multi,
opt_state_multi,
x_batches_multi,
y_batches_multi,
)
show_table(
["Metric", "Value"],
[
("GPUs used", NUM_DEVICES),
("Global batch", GLOBAL_BATCH),
("Per-GPU batch", PER_DEVICE_BATCH),
("Throughput", f"{result_multi['examples_per_sec']:,.0f} examples/sec"),
("Step time", f"{result_multi['ms_per_step']:.2f} ms"),
("Final loss", f"{result_multi['final_loss']:.4f}"),
("Final accuracy", f"{100 * result_multi['final_accuracy']:.1f}%"),
],
title=f"Data-parallel training on {NUM_DEVICES} GPUs",
)
В результате вы должны получить таблицу, похожую на базовую, теперь с отображением как глобального пакета данных, так и пакета данных для каждого графического процессора. Не сравнивайте показатели пропускной способности визуально — на следующем шаге это будет сделано правильно, и коэффициент — единственный показатель, имеющий значение.
8. Сравните пропускную способность при работе с одной и несколькими видеокартами.
На этом этапе пакет данных для каждого графического процессора одинаков в обоих запусках, а в многопроцессорном режиме обрабатывается больше примеров за шаг, поскольку каждый графический процессор получает свой собственный сегмент.
Вы добавляете графический процессор и одновременно увеличиваете рабочую нагрузку, а затем задаетесь вопросом, успевает ли пропускная способность соответствовать этому уровню. Это не тот же вопрос, что и «завершится ли обработка фиксированного пакета данных в два раза быстрее». Здесь под «быстрее» подразумевается более высокая пропускная способность обучения в примерах в секунду.
speed_ratio = result_multi["examples_per_sec"] / result_1gpu["examples_per_sec"]
show_table(
["", "1 GPU", f"{NUM_DEVICES} GPUs", "Throughput ratio"],
[
("Per-GPU batch", PER_DEVICE_BATCH, PER_DEVICE_BATCH, "same"),
("Global batch", PER_DEVICE_BATCH, GLOBAL_BATCH, f"{NUM_DEVICES}x"),
(
"Examples/sec",
f"{result_1gpu['examples_per_sec']:,.0f}",
f"{result_multi['examples_per_sec']:,.0f}",
f"{speed_ratio:.2f}x",
),
(
"ms/step",
f"{result_1gpu['ms_per_step']:.2f}",
f"{result_multi['ms_per_step']:.2f}",
"",
),
],
title="Throughput: same per-GPU batch",
aligns=["left", "right", "right", "right"],
)
show_bars(
[
("1 GPU", result_1gpu["examples_per_sec"]),
(f"{NUM_DEVICES} GPUs", result_multi["examples_per_sec"]),
],
"Training throughput (examples/sec)",
"examples/s",
)
Оцените результат честно.
Следующий фрагмент кода основывается на том, что вы фактически измерили. Запустите его и посмотрите, что получится.
step_ratio = result_multi["ms_per_step"] / result_1gpu["ms_per_step"]
if speed_ratio >= 1.0:
message = (
f"The multi-GPU run is faster for this Fashion-MNIST workload: "
f"throughput improves by {speed_ratio:.2f}x. Each GPU still processes "
f"{PER_DEVICE_BATCH} examples, while the global batch increases from "
f"{PER_DEVICE_BATCH} to {GLOBAL_BATCH}. Step time changes by {step_ratio:.2f}x, "
f"so the larger batch translates into higher examples/sec."
)
else:
message = (
f"This run is still communication-bound: throughput changes by {speed_ratio:.2f}x. "
f"Increase BLOCK_REPEATS or PER_DEVICE_BATCH to give each GPU more local work."
)
border_color = "#1a7f37" if speed_ratio >= 1.0 else "#d1242f"
display(HTML(
"<div style='font-family: system-ui; max-width: 900px; "
f"border-left: 4px solid {border_color}; padding: 10px 12px; "
"background: #f6f8fa; margin: 12px 0;'>"
f"{html.escape(message)}"
"</div>"
))
Если коэффициент равен или превышает 1,0, каждый графический процессор выполняет одинаковую локальную нагрузку, и время шага увеличивается меньше, чем время обработки пакета. Если коэффициент ниже 1,0, выполнение ограничено объемом обмена данными, и градиентное сокращение по пути PHB который вы видели при проверке топологии, обходится дороже, чем дополнительный графический процессор.
9. Получите явный контроль с помощью shard_map.
Стандартный подход охватывает большинство задач параллельной обработки данных. Однако иногда требуется точно контролировать, какие вычисления выполняет каждый графический процессор. shard_map позволяет написать функцию, работающую с массивами для каждого сегмента , и использовать явные объединения для обмена данными между устройствами.
Внутри функции shard_map :
- Каждый графический процессор получает свой локальный сегмент, например,
(1024, 784) -
in_specsопределяет способ нарезки входных данных. -
out_specsопределяет способ сборки выходных данных. -
jax.lax.pmean(x, 'data')усредняет значениеxпо всем графическим процессорам вдоль оси'data'
Обратите внимание, что теперь вызовы jax.lax.pmean представляют собой все операции reduce, которые jax.jit добавил для вас на предыдущем шаге.
@partial(
jax.shard_map,
mesh=mesh,
in_specs=(P(), P("data", None), P("data",)),
out_specs=(P(), P(), P()),
)
def compute_grads_shardmap(params, x_shard, y_shard):
(loss, metrics), grads = jax.value_and_grad(loss_with_metrics, has_aux=True)(
params,
(x_shard, y_shard),
)
grads = jax.lax.pmean(grads, "data")
loss = jax.lax.pmean(loss, "data")
accuracy = jax.lax.pmean(metrics["accuracy"], "data")
return grads, loss, accuracy
@jax.jit
def train_step_explicit(params, opt_state, batch):
x, y = batch
grads, loss, accuracy = compute_grads_shardmap(params, x, y)
updates, opt_state = optimizer.update(grads, opt_state, params)
params = optax.apply_updates(params, updates)
return params, opt_state, {"loss": loss, "accuracy": accuracy}
Обновление оптимизатора происходит вне shard_map . Градиенты уже усреднены на момент их получения, а параметры реплицируются, поэтому каждая видеокарта применяет идентичное обновление.
Теперь сравним его с автоматической версией:
params_explicit = jax.device_put(init_params(seed=1), replicated)
opt_state_explicit = optimizer.init(params_explicit)
result_explicit = benchmark_training(
train_step_explicit,
params_explicit,
opt_state_explicit,
x_batches_multi,
y_batches_multi,
)
show_table(
["Approach", "Examples/sec", "ms/step"],
[
(
"jit on sharded arrays",
f"{result_multi['examples_per_sec']:,.0f}",
f"{result_multi['ms_per_step']:.2f}",
),
(
"shard_map explicit",
f"{result_explicit['examples_per_sec']:,.0f}",
f"{result_explicit['ms_per_step']:.2f}",
),
],
title="Automatic vs explicit data parallelism",
aligns=["left", "right", "right"],
)
drop_device_refs(
"params_multi",
"opt_state_multi",
"params_explicit",
"opt_state_explicit",
"x_batches_multi",
"y_batches_multi",
)
Вы должны увидеть две строки, описывающие одно и то же вычисление, выраженное двумя разными способами. Воспринимайте это как подтверждение того, что явная версия воспроизводит автоматическую, а не как гонку — они выполняют одну и ту же работу и приходят к одному и тому же результату.
10. Проведите глобальное сканирование размера партии.
Параллельная обработка данных позволяет масштабировать глобальный размер пакета в зависимости от количества графических процессоров. Большие пакеты компенсируют накладные расходы на запуск ядра и повышают эффективность использования графических процессоров до тех пор, пока узким местом не станет память или обмен данными на каждом устройстве.
Приведенный ниже тест проверяет несколько глобальных размеров пакетов данных на всех графических процессорах. Размеры, которые не делятся нацело на количество ваших графических процессоров, пропускаются, а пакет данных, в котором произошел сбой, например, из-за нехватки памяти, сообщается без остановки цикла.
BATCH_SIZES = [256, 512, 1024, 2048, 4096]
scaling_results = []
for bs in BATCH_SIZES:
if bs % NUM_DEVICES != 0:
print(f"Skipping global batch {bs}: not divisible by {NUM_DEVICES} GPUs.")
continue
try:
x_bs, y_bs = make_fashion_batches(bs)
x_bs = jax.device_put(x_bs, all_data_sharding)
y_bs = jax.device_put(y_bs, all_label_sharding)
params_bs = jax.device_put(init_params(seed=1), replicated)
opt_bs = optimizer.init(params_bs)
result = benchmark_training(
train_step,
params_bs,
opt_bs,
x_bs,
y_bs,
)
scaling_results.append(
{
"batch_size": bs,
"per_device": bs // NUM_DEVICES,
"examples_per_sec": result["examples_per_sec"],
"ms_per_step": result["ms_per_step"],
}
)
except Exception as e:
print(f"Batch size {bs}: {e}")
finally:
drop_device_refs("x_bs", "y_bs", "params_bs", "opt_bs", "result")
Это самый долго выполняющийся фрагмент кода в практическом задании, где каждый размер пакета запускает собственную компиляцию, прогрев и повторные операции по расписанию. Теперь постройте график полученных данных:
show_table(
["Global batch", "Per GPU", "Examples/sec", "ms/step"],
[
(
r["batch_size"],
r["per_device"],
f"{r['examples_per_sec']:,.0f}",
f"{r['ms_per_step']:.2f}",
)
for r in scaling_results
],
title=f"Batch-size scaling on {NUM_DEVICES} GPUs",
aligns=["right", "right", "right", "right"],
)
fig, ax = plt.subplots(figsize=(8, 5))
batches = [r["batch_size"] for r in scaling_results]
throughputs = [r["examples_per_sec"] for r in scaling_results]
ax.plot(
batches,
throughputs,
"o-",
color="#0969da",
linewidth=2,
markersize=8,
)
ax.set_xlabel("Global batch size")
ax.set_ylabel("Examples per second")
ax.set_title(f"Throughput vs batch size — {NUM_DEVICES} GPUs data-parallel")
ax.set_xscale("log", base=2)
ax.set_xticks(batches)
ax.set_xticklabels([str(b) for b in batches])
ax.grid(True, alpha=0.25)
fig.tight_layout()
plt.show()
Для каждого завершенного пакета данных вы должны получить одну строку таблицы и одну точку на кривой. Вы увидите, что пропускная способность возрастает по мере увеличения размера пакета и амортизации фиксированных накладных расходов на каждом шаге, а затем выравнивается, как только графические процессоры насыщаются или начинает доминировать алгоритм all-reduce. Точка выравнивания является свойством данной модели на данном межсоединении, и это значение важно знать, прежде чем масштабировать систему на большее количество графических процессоров.
11. Уборка
Удалите рабочую нагрузку Jupyter, включая балансировщик нагрузки и постоянный том:
kubectl delete -f deploy/jupyter.yaml
Уничтожьте кластер, пул узлов, VPC и учетную запись службы:
cd terraform
terraform destroy
При появлении запроса введите yes , затем подтвердите, что ничего не осталось:
gcloud container clusters list
gcloud compute instances list
Оба поля должны быть пустыми для этого проекта. Если вы создали проект только для этой серии, вы можете удалить весь проект из консоли Cloud .
12. Поздравляем!
Вы перенесли цикл обучения JAX с одного графического процессора на два, изменив местоположение массивов, а не переписав этап обучения.
Что вы узнали
- Как
Mesh(devices, axis_names)сопоставляет физические графические процессоры с логической сеткой с именованными осями, и что эти имена вы можете выбрать сами. - Как
PartitionSpecопределяет, какие измерения массива соответствуют каким осям сетки —P('data', None)сегментирует измерение пакета и дублирует признаки - Как
NamedSharding(mesh, spec)объединяет mesh и spec в план размещения дляjax.device_put - Как
jax.debug.visualize_array_shardingпоказывает, какой графический процессор содержит какой фрагмент данных, и почему её следует запускать после каждого изменения расположения массива. - Как работает автоматическая параллелизация:
jax.jitна шардированных входных данных автоматически вставляет вычисления для всего процесса Reduce и для каждого шарда без изменения кода. - Как
shard_mapобеспечивает явный контроль над каждым сегментом, используяjax.lax.pmeanдля усреднения градиентов, когда необходимо настроить шаблон обмена данными. - Как работает масштабирование размера пакета: большие глобальные пакеты могут повысить пропускную способность до тех пор, пока использование графического процессора, памяти или обмена данными не станет узким местом.
Следующие шаги
- Codelab 7: Обучение трансформера от начала до конца с помощью Flax NNX и Orbax сочетает механизм внимания из Codelab 5 с обучением на нескольких графических процессорах из этого Codelab.
- Увеличьте значение
BLOCK_REPEATSилиPER_DEVICE_BATCHчтобы предоставить каждому графическому процессору больше локальной работы, затем повторно запустите этап сравнения и понаблюдайте за изменением коэффициента пропускной способности. - Увеличьте пул узлов до
g2-standard-48с 4 L4 — установитеgpu_count = 4вterraform.tfvarsиnvidia.com/gpu: "4"вdeploy/jupyter.yaml— и повторно запустите сканирование размера пакета на четырех устройствах. - Попробуйте использовать двухмерную сетку с осью
model, расположенной рядом сdata, и распределитеw_blockвдоль неё, вместо того чтобы дублировать её.