Запустите свою первую JAX-программу на графических процессорах NVIDIA с помощью GKE.

1. Введение

Курс обучения JAX на GPU. Лабораторная работа 1: Начало работы с JAX на GPU.

JAX — это библиотека Python для высокопроизводительных численных вычислений. На первый взгляд она похожа на NumPy, но по сути она отслеживает ваши функции Python, компилирует их с помощью XLA и запускает результат на ускорителях, таких как графические процессоры NVIDIA.

В этом практическом занятии вы создадите кластер Google Kubernetes Engine с графическими процессорами NVIDIA L4, используя Terraform, запустите JupyterLab внутри официального контейнера NVIDIA JAX на этом узле с графическим процессором и напишете свой первый JAX-вычисление. К концу у вас будет рабочая среда, на основе которой будет строиться вся остальная часть этой серии из восьми частей.

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

  • Создайте кластер GKE Standard с пулом из двух узлов NVIDIA L4 GPU, используя Terraform.
  • Разверните JupyterLab на узле с графическим процессором из образа контейнера NVIDIA JAX.
  • Проверьте работу графического процессора от начала до конца с помощью nvidia-smi и jax.devices()
  • Напишите код для работы с массивами JAX с использованием jax.numpy и убедитесь, что результат хранится на графическом процессоре.
  • Примените три преобразования, определяющие JAX: jax.jit , jax.grad и jax.vmap

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

  • Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
  • Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
  • Знание базовых языков Python и NumPy. Опыт работы с CUDA или Kubernetes не требуется.

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

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

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

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

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

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

gcloud config set project <YOUR_PROJECT_ID>

Все действия на этом этапе выполняются в Cloud Shell , в которой уже установлены gcloud , kubectl , terraform и git .

Включите необходимые API.

Включить все API, необходимые для этого практического занятия, одной командой:

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

Подтвердите свою квоту на видеокарту.

В конфигурации Terraform запрашиваются две видеокарты NVIDIA L4. Убедитесь, что у вас есть необходимая квота:

gcloud compute regions describe us-central1 \
  --format="value(quotas.filter(metric:NVIDIA_L4_GPUS).limit)"

Вы должны увидеть значение 2 или выше. Если вы видите 0 , запросите увеличение квоты, прежде чем продолжить.

Клонируйте репозиторий мастерской.

Модуль Terraform и манифест Kubernetes находятся в репозитории мастерской:

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

Вам понадобятся две следующие директории:

  • terraform/ содержит кластер GKE Standard, VPC, учетную запись службы узлов и пул узлов L4 GPU.
  • deploy/jupyter.yaml содержит PersistentVolumeClaim, под JupyterLab и службу балансировки нагрузки.

3. Создайте кластер графических процессоров с помощью Terraform.

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

Настройте свой проект

Скопируйте файл с примерами переменных и укажите в нем путь к вашему проекту:

cd terraform
cp terraform.tfvars.example terraform.tfvars

Отредактируйте terraform.tfvars и установите project_id . Значения по умолчанию для всего остального соответствуют этому руководству:

project_id   = "<YOUR_PROJECT_ID>"
region       = "us-central1"
zone         = "us-central1-a"
cluster_name = "jax-gpu-cluster"
machine_type = "g2-standard-24"
gpu_type     = "nvidia-l4"
gpu_count    = 2

Поймите, что вы создаёте.

Перед применением ознакомьтесь с определением пула узлов в main.tf Именно эта часть преобразует обычный узел в узел с графическим процессором:

resource "google_container_node_pool" "gpu" {
  name     = "gpu-pool"
  location = var.zone
  cluster  = google_container_cluster.primary.name

  node_count = 1

  node_config {
    machine_type = var.machine_type

    guest_accelerator {
      type  = var.gpu_type   # nvidia-l4
      count = var.gpu_count  # 2

      gpu_driver_installation_config {
        gpu_driver_version = "DEFAULT"
      }
    }

    disk_size_gb = 100
    disk_type    = "pd-balanced"
    # ...
  }
}

Важны две детали. Во-первых, machine_type и gpu_count должны совпадать: g2-standard-24 поставляется ровно с двумя графическими процессорами L4, а g2-standard-48 — с четырьмя. Во-вторых, gpu_driver_installation_config делает узел пригодным для использования — GKE устанавливает соответствующий драйвер NVIDIA, поэтому вашему Pod нужно только подключить пользовательские библиотеки CUDA.

Применять

terraform init
terraform apply

Проверьте план и введите yes . Создание кластера и выделение пула узлов займут около 10 минут . Сейчас самое время ознакомиться с информацией на будущее.

После завершения получите учетные данные кластера, чтобы kubectl мог взаимодействовать с новым кластером:

$(terraform output -raw get_credentials_command)

Убедитесь, что узел оснащен графическими процессорами (GPU).

kubectl get nodes -o custom-columns=\
NAME:.metadata.name,GPU:.status.allocatable.nvidia\\.com/gpu

Вы должны увидеть результат, похожий на следующий:

NAME                                            GPU
gke-jax-gpu-cluster-gpu-pool-3f21a0b4-k7wq      2

4. Разверните JupyterLab на узле с графическим процессором.

Теперь у вас есть узел с графическим процессором, но на нём ничего не запущено. Манифест в deploy/jupyter.yaml планирует запуск Pod-а, который запрашивает оба графических процессора и запускает JupyterLab из официального контейнера NVIDIA JAX.

Примените манифест

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

В результате создаются три объекта:

  • jax-workspace-pvc — это постоянный том объемом 50 ГБ, монтируемый в /workspace , чтобы ваши ноутбуки сохраняли работоспособность после перезапуска Pod.
  • jax-jupyter — это Pod, который запускает nvcr.io/nvidia/jax:26.04-maxtext-py3 и запрашивает nvidia.com/gpu: "2" .
  • jax-jupyter-svc — это балансировщик нагрузки, который предоставляет доступ к JupyterLab через порт 8884.

Запрос к графическому процессору — это важная строка:

resources:
  limits:
    nvidia.com/gpu: "2"
    memory: "48Gi"
    cpu: "12"

nvidia.com/gpu — это расширенный ресурс, рекламируемый плагином устройств GKE. Kubernetes планирует этот Pod только на том узле, который может его удовлетворить, именно так Pod попадает в ваш пул узлов с графическими процессорами.

Дождитесь готовности капсулы.

Образ контейнера большой, и Pod также устанавливает JupyterLab с помощью pip при запуске, поэтому первая загрузка занимает несколько минут:

kubectl get pod jax-jupyter -w

Дождитесь, пока STATUS не отобразится Running , затем нажмите Ctrl+C .

Получите URL-адрес и токен JupyterLab.

Получите внешний IP-адрес сервиса:

kubectl get svc jax-jupyter-svc -w

Подождите, пока EXTERNAL-IP не изменится. Перейдите по указанному адресу, затем нажмите Ctrl+C .

JupyterLab выводит одноразовый токен авторизации в лог пода:

kubectl logs jax-jupyter | grep -o 'token=[a-z0-9]*' | head -1

Откройте http:// :8884 в вашем браузере и вставьте токен, когда появится соответствующий запрос.

Создайте блокнот

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

5. Убедитесь, что JAX видит графический процессор.

Прежде чем писать какой-либо JAX-код, убедитесь, что оборудование видно изнутри контейнера. Если этот шаг не удастся, дальнейшие действия работать не будут.

Проверьте оборудование.

Запустите nvidia-smi из ноутбука:

!nvidia-smi

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

Теперь запросите вычислительные возможности — двухзначное число, определяющее поколение оборудования. В последующих практических занятиях используются функции, зависящие от этого параметра: для алгоритма cuDNN с механизмом внимания требуется версия 8.0 или новее, а для FP8 — версия 9.0 или новее.

import subprocess


def get_compute_capability() -> tuple[int, int]:
    """Query the compute capability of the first visible GPU."""
    out = subprocess.check_output(
        ["nvidia-smi", "--query-gpu=compute_cap", "--format=csv,noheader"],
        text=True,
    )
    major, minor = out.strip().split("\n")[0].split(".")
    return int(major), int(minor)


SM_MAJOR, SM_MINOR = get_compute_capability()
print(f"Detected compute capability: SM {SM_MAJOR}.{SM_MINOR}")

if SM_MAJOR < 7:
    print("WARNING: this course assumes SM 7.0+ (Volta or newer).")
else:
    print("GPU is compatible with this course.")

В L4 используется графический процессор Ada Lovelace, поэтому вы должны увидеть SM 8.9 .

Убедитесь, что JAX обнаружил графический процессор.

import jax
import jax.numpy as jnp

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

print(f"JAX version:        {jax.__version__}")
print(f"Default backend:    {jax.default_backend()}")
print(f"Available devices:  {devices}")

assert gpu_devices, f"No GPU backend found. Available devices: {devices}"
print(f"GPU devices:        {gpu_devices}")

Вы должны увидеть результат, похожий на следующий:

JAX version:        0.7.2
Default backend:    gpu
Available devices:  [CudaDevice(id=0), CudaDevice(id=1)]
GPU devices:        [CudaDevice(id=0), CudaDevice(id=1)]

Первый import jax занимает несколько секунд, поскольку JAX инициализирует среду выполнения CUDA и проверяет устройства.

Как все части складываются воедино

JAX сам по себе никогда не взаимодействует напрямую с графическим процессором. Он создает программу , описывающую ваши вычисления, и передает ее в стек:

Слой

Роль

ДЖАКС

Преобразует вашу функцию Python в промежуточное представление.

XLA

Компилирует этот промежуточный код в оптимизированный код для графического процессора.

cuDNN, cuBLAS, NCCL

Библиотеки NVIDIA, в которые XLA обращается для выполнения сверток, GEMM и коллективных операций, используют различные алгоритмы.

Драйвер и среда выполнения CUDA

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

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

6. Напишите код для работы с массивами JAX на графическом процессоре.

Самый быстрый способ освоиться с JAX — заметить, как много в нём общего с NumPy. Применяются те же конструкторы и правила широковещательной рассылки. Меняется лишь местоположение массива и способ выполнения вычислений.

Сравните NumPy и JAX бок о бок.

import numpy as np

# NumPy: runs on the CPU, stored in host memory
x_np = np.arange(8, dtype=np.float32)
y_np = np.sin(x_np) ** 2 + np.cos(x_np) ** 2
print(f"NumPy result:  {y_np}")
print(f"NumPy device:  CPU (host memory)")
print()

# JAX: same code, different array library
x = jnp.arange(8, dtype=jnp.float32)
y = jnp.sin(x) ** 2 + jnp.cos(x) ** 2
print(f"JAX result:    {y}")
print(f"JAX device:    {y.device}")

# Sanity check: the two answers should agree
np.testing.assert_allclose(y_np, np.asarray(y), atol=1e-6)
print()
print("NumPy and JAX agree.")

Вы должны увидеть результат, похожий на следующий:

JAX result:    [1. 1. 1. 1. 1. 1. 1. 1.]
JAX device:    cuda:0

Три момента, на которые следует обратить внимание:

  1. Код идентичен, за исключением замены np на jnp .
  2. y.device сообщает об устройстве CUDA — JAX автоматически разместил массив на графическом процессоре, поскольку это бэкенд по умолчанию.
  3. JAX возвращает собственный тип массива ( jax.Array ), а не массив NumPy. Вызов np.asarray(y) запускает передачу данных с графического процессора на центральный процессор .

Обратите внимание на третий пункт. Каждое взаимодействие между графическим процессором и хостом занимает время, а вывод массива JAX требует синхронизации, поскольку Python должен получить значение для его отображения. Это хорошо для небольшого примера; настоящая проблема внутри цикла с таймером. Рабочее правило таково: создавайте массивы с помощью jnp , работайте с ними с помощью jnp и преобразуйте в NumPy только тогда, когда вам действительно нужно посмотреть на значения. В Codelab 3 показано, как обнаружить случайные передачи в профиле.

7. Примените JIT, grad и vmap.

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

jax.jit компилирует вашу функцию

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

Первый вызов выполняется медленно, потому что он компилируется. Все последующие вызовы выполняются быстро.

import time


def f(x):
    """Compose tanh, sin, and log1p so XLA has multiple ops to fuse when jitted."""
    return jnp.tanh(x) * jnp.sin(x) + jnp.log1p(x * x)


x = jnp.arange(1_000_000, dtype=jnp.float32)

# Eager: one kernel launch per operation
_ = f(x).block_until_ready()  # warm up
t0 = time.perf_counter()
for _ in range(10):
    y = f(x).block_until_ready()
eager_ms = (time.perf_counter() - t0) * 1000 / 10
print(f"Eager:            {eager_ms:6.3f} ms / call")

# Compiled: optimized executable, often with fused operations
f_jit = jax.jit(f)
_ = f_jit(x).block_until_ready()  # first call compiles
t0 = time.perf_counter()
for _ in range(10):
    y = f_jit(x).block_until_ready()
jit_ms = (time.perf_counter() - t0) * 1000 / 10
print(f"jax.jit (cached): {jit_ms:6.3f} ms / call")
print(f"Speedup:          {eager_ms / jit_ms:6.1f}x")

Точное ускорение зависит от размера и структуры ваших вычислений, но закономерность универсальна: JAX с предварительным включением удобен, а скомпилированный JAX работает быстро.

jax.grad автоматически производит дифференцирование

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

def loss(w, x, y):
    """Mean squared error of `w*x` vs `y`; scalar loss for the `jax.grad` demo below."""
    pred = w * x
    return jnp.mean((pred - y) ** 2)


w = jnp.array(0.5)
xs = jnp.array([1.0, 2.0, 3.0, 4.0])
ys = jnp.array([2.0, 4.0, 6.0, 8.0])

# grad returns a function with the same signature, differentiating w.r.t. the first argument
dloss_dw = jax.grad(loss)

print(f"loss(w=0.5):   {loss(w, xs, ys):.4f}")
print(f"dloss/dw:      {dloss_dw(w, xs, ys):.4f}")

# Sanity check against a finite-difference approximation
eps = 1e-3
fd = (loss(w + eps, xs, ys) - loss(w - eps, xs, ys)) / (2 * eps)
print(f"finite diff:   {fd:.4f}  (should match)")

Градиент отрицательный, что говорит оптимизатору о том, что увеличение w приведет к уменьшению функции потерь — и это совершенно верно, поскольку истинная зависимость y = 2x , а вы начали с w = 0.5 . В Codelab 4 показан полный цикл обучения, основанный на этом принципе.

jax.vmap выполняет векторизацию по всему пакету данных.

Графические процессоры предпочитают пакетную обработку. Наивный способ применения функции ко множеству входных данных — это цикл for Python, но он запускает ядра по одному и истощает ресурсы графического процессора. jax.vmap принимает функцию, написанную для одного примера , и возвращает версию, которая работает с пакетом данных , без цикла и без ручного изменения формы.

def predict(W, x):
    """Tanh of a single-example matrix-vector product; vmapped below to batch over many `x`."""
    # Single example: W is (out, in), x is (in,) -> result is (out,)
    return jnp.tanh(W @ x)


key_w, key_x = jax.random.split(jax.random.key(0))
W = jax.random.normal(key_w, (4, 3))
xs = jax.random.normal(key_x, (10, 3))  # batch of 10 examples

# Without vmap: a Python loop, one kernel launch per example
ys_loop = jnp.stack([predict(W, x) for x in xs])

# With vmap: batch over the leading axis of xs, share W across the batch
batched_predict = jax.vmap(predict, in_axes=(None, 0))
ys_vmap = batched_predict(W, xs)

print(f"ys_loop shape:  {ys_loop.shape}")
print(f"ys_vmap shape:  {ys_vmap.shape}")
np.testing.assert_allclose(np.asarray(ys_loop), np.asarray(ys_vmap), atol=1e-6)
print("vmap matches the explicit loop.")

Аргумент in_axes=(None, 0) означает: не выполнять пакетную обработку W (а широковещательную), а выполнять пакетную обработку xs вдоль оси 0. Результат идентичен циклу, но выполняется как единая пакетная операция на графическом процессоре.

Составьте их

Настоящая суперсила — это композиция. Преобразования суммируются:

fast_batched_grad = jax.jit(
    jax.vmap(jax.grad(loss), in_axes=(None, 0, 0))
)

Одна строка кода предоставляет вам скомпилированную, векторизованную, дифференцированную функцию, которая возвращает градиент для каждого примера в пакете (x, y) — это практически всё, что вам нужно для пакетного обучения. В Codelab 4 этот подход применяется на реальном наборе данных.

8. Уборка

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

kubectl delete -f deploy/jupyter.yaml

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

cd terraform
terraform destroy

При появлении запроса введите yes . Удаление займет несколько минут.

Наконец, убедитесь, что ничего не осталось:

gcloud container clusters list
gcloud compute instances list

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

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

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

Что вы узнали

  • Как настроить кластер GKE Standard с пулом узлов NVIDIA L4 GPU с помощью Terraform, включая блок gpu_driver_installation_config , который делает узел доступным для использования.
  • Как запланировать размещение Pod-а на узле GPU с помощью расширенного ресурса nvidia.com/gpu
  • Как работает стек JAX на GPU: трассировка JAX, компиляция XLA, а также выполнение cuDNN, cuBLAS и среды выполнения CUDA.
  • Как проверить среду GPU с помощью nvidia-smi , вычислительных возможностей, jax.devices() и jax.default_backend()
  • В чём jax.numpy соответствует NumPy, а в чём отличается: неизменяемые массивы, обновления .at[...] , типы данных по умолчанию 32-битные и стоимость передачи данных между хостами.
  • Как применять и компоновать jax.jit , jax.grad и jax.vmap

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

  • Практическое занятие 2: Управление компиляцией JAX с помощью jax.jit , где вы узнаете, почему первый вызов происходит медленно, что запускает перекомпиляцию и как поддерживать стабильность параметров компиляции.
  • Попробуйте установить JAX_PLATFORMS=cpu перед импортом JAX, чтобы принудительно запустить программу на ЦП, и сравните результаты измерения времени выполнения с данными jax.jit указанными выше.
  • Увеличьте количество узлов в пуле до 4 графических процессоров L4, установив machine_type = "g2-standard-48" и gpu_count = 4 , а также сопоставив nvidia.com/gpu в deploy/jupyter.yaml

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