Экспорт и запуск обученной модели JAX.

1. Введение

Обучение работе с Jax на графических процессорах. Лабораторная работа 8: Обслуживание и дальнейшие шаги.

В этом практическом занятии вы начинаете с обученной контрольной точки и проходите весь конвейер вывода JAX от начала до конца: от прямого прохода, скомпилированного с помощью JIT, до компиляции «на лету» (AOT), экспорта в JAX с помощью jax.export и сохранения модели TensorFlow с помощью jax2tf . Каждый путь измеряется, и все четыре предназначены для получения одинаковых прогнозов.

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

  • Пересоберите трансформер Codelab 7 и загрузите в него обученные веса с помощью Orbax.
  • Оберните прямой проход в jax.jit и сравните задержку первого вызова с задержкой вызова из кэша.
  • Удалите холодный старт с помощью AOT-компиляции ( lower() , затем compile() ) и прочтите IR-файл StableHLO.
  • Измерьте пропускную способность прямой передачи в токенах/сек для четырех размеров пакетов.
  • Создайте переносимый артефакт с помощью jax.export , затем десериализуйте его и вызовите.
  • Преобразуйте модель в TensorFlow SavedModel с помощью jax2tf и сравните все четыре пути.

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

  • Проект Google Cloud с включенной оплатой и кредитами для семинара или резервированием, покрывающим использование графического процессора.
  • Квота на использование как минимум двух видеокарт NVIDIA L4 в выбранном вами регионе ( как проверить квоту на видеокарты )
  • Выполнение заданий с 1 по 7 в рамках практических работ или использование эквивалентной среды JAX GPU с поддержкой CUDA.
  • Завершённое прохождение Codelab 7, контрольная точка Orbax которого загружается в этом лабораторном задании.

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

Четыре пути от обучения к служению

В этом практическом руководстве сравниваются четыре практических способа переноса обученной модели JAX на сервер, в зависимости от целевой платформы развертывания:

Путь

Формат

Цель подачи

Когда использовать

jax.jit

Кэшированный исполняемый файл, находящийся в процессе выполнения

Сервер на Python (FastAPI, Flask)

Простейший путь обслуживания JAX с низкой задержкой

AOT compile

Предварительно скомпилированный исполняемый файл, работающий в процессе выполнения.

Запуск/прогрев сервера Python

Избегайте задержки компиляции при первом запросе.

jax.export

Сериализованный JAX-экспорт с использованием StableHLO + метаданные

Совместимая среда выполнения JAX для экспортируемой платформы(платформ)

JAX-нативный портативный артефакт

jax2tf

TF SavedModel

TF Serving, TFX

экосистема TensorFlow

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

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:// :8884 , вставьте токен и создайте новый блокнот Python 3 в /workspace . Каждый блок кода из этого практического занятия помещается в ячейку этого блокнота.

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

!pip install --quiet flax orbax-checkpoint tensorflow

Пакеты flax и orbax-checkpoint обычно поставляются в контейнере NVIDIA JAX, поэтому они, как правило, ничего не делают. Если TensorFlow еще не установлен, pip загружает его.

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

Импортируйте JAX, Flax NNX и Orbax, убедитесь, что отображается хотя бы один графический процессор, и определите два вспомогательных класса, которые используются в остальной части практического задания для блокировки результатов и отрисовки таблиц результатов.

import os
os.environ["LD_LIBRARY_PATH"] = "/usr/local/nvidia/lib64:" + os.environ.get("LD_LIBRARY_PATH", "")

import html
import pathlib
import time
import warnings

from IPython.display import HTML, display
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
from flax import nnx
import orbax.checkpoint as ocp


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"GPU devices:     {gpu_devices}")

assert len(gpu_devices) >= 1, (
    f"This lesson needs at least 1 GPU. Found {len(gpu_devices)}. "
    f"Available devices: {devices}"
)


def block_tree(tree):
    return jax.block_until_ready(tree)


def show_table(headers, rows, title=None, aligns=None):
    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("\n".join(parts)))

Вы должны увидеть версию JAX, gpu в качестве бэкенда по умолчанию и список, содержащий как минимум одно устройство CUDA.

Убедитесь, что у вас есть обученный контрольный пункт.

В этом практическом занятии загружается контрольная точка Orbax, которая была записана в /tmp/jax-course/l7-checkpoints/trained при обучении трансформера от начала до конца с помощью Flax NNX и Orbax. Поэтому убедитесь, что эта директория всё ещё существует, прежде чем что-либо собирать.

import pathlib
ckpt_dir = pathlib.Path("/tmp/jax-course/l7-checkpoints")
assert (ckpt_dir / "trained").exists(), (
    f"No checkpoint at {ckpt_dir / 'trained'}. Run codelab 7 first."
)
print(f"Found checkpoint: {ckpt_dir / 'trained'}")

Вы должны увидеть на экране путь к контрольной точке.

3. Перестройте модель и загрузите контрольную точку.

Контрольная точка Orbax хранит значения параметров, а не класс модели, который их создал. Для ее восстановления сначала необходима та же архитектура, поэтому на этом шаге TinyTransformer из Codelab 7 объявляется заново, а затем в него вводятся сохраненные веса.

Переосмыслить архитектуру

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

VOCAB_SIZE = 256
D_MODEL = 256
NUM_HEADS = 4
FFN_DIM = 1024
NUM_LAYERS = 4
MAX_SEQ_LEN = 256


def causal_sdpa(query, key, value, **_):
    return jax.nn.dot_product_attention(query, key, value, is_causal=True)


class TransformerBlock(nnx.Module):
    def __init__(self, d_model: int, num_heads: int, ffn_dim: int, rngs: nnx.Rngs):
        self.ln1 = nnx.LayerNorm(d_model, rngs=rngs)
        self.attn = nnx.MultiHeadAttention(
            num_heads=num_heads,
            in_features=d_model,
            decode=False,
            attention_fn=causal_sdpa,
            rngs=rngs,
        )
        self.ln2 = nnx.LayerNorm(d_model, rngs=rngs)
        self.fc_up = nnx.Linear(d_model, ffn_dim, rngs=rngs)
        self.fc_down = nnx.Linear(ffn_dim, d_model, rngs=rngs)

    def __call__(self, x):
        x = x + self.attn(self.ln1(x))
        h = jax.nn.gelu(self.fc_up(self.ln2(x)))
        x = x + self.fc_down(h)
        return x


class TinyTransformer(nnx.Module):
    def __init__(self, vocab_size: int, d_model: int, num_heads: int,
                 ffn_dim: int, num_layers: int, max_seq_len: int,
                 rngs: nnx.Rngs):
        self.token_embed = nnx.Embed(vocab_size, d_model, rngs=rngs)
        self.pos_embed = nnx.Embed(max_seq_len, d_model, rngs=rngs)
        self.blocks = nnx.List([
            TransformerBlock(d_model, num_heads, ffn_dim, rngs=rngs)
            for _ in range(num_layers)
        ])
        self.final_norm = nnx.LayerNorm(d_model, rngs=rngs)
        self.lm_head = nnx.Linear(d_model, vocab_size, use_bias=False, rngs=rngs)

    def __call__(self, tokens):
        B, T = tokens.shape
        x = self.token_embed(tokens) + self.pos_embed(jnp.arange(T))
        for block in self.blocks:
            x = block(x)
        x = self.final_norm(x)
        return self.lm_head(x)


param_count = sum(x.size for x in jax.tree.leaves(nnx.state(TinyTransformer(
    VOCAB_SIZE, D_MODEL, NUM_HEADS, FFN_DIM, NUM_LAYERS, MAX_SEQ_LEN,
    rngs=nnx.Rngs(0),
), nnx.Param)))
print(f"TinyTransformer: {param_count:,} parameters")

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

Восстановите тренировочный вес

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

ckpt_dir = pathlib.Path("/tmp/jax-course/l7-checkpoints")

model = TinyTransformer(
    VOCAB_SIZE, D_MODEL, NUM_HEADS, FFN_DIM, NUM_LAYERS, MAX_SEQ_LEN,
    rngs=nnx.Rngs(0),
)
model_params = nnx.state(model, nnx.Param)
abstract_params = jax.tree.map(
    lambda x: jax.ShapeDtypeStruct(x.shape, x.dtype), model_params
)

checkpointer = ocp.StandardCheckpointer()
restored_params = checkpointer.restore(ckpt_dir / "trained", abstract_params)

# Move to a single GPU
single_device = gpu_devices[0]
restored_params = jax.device_put(restored_params, single_device)
nnx.update(model, restored_params)

print(f"\u2705 Checkpoint loaded from {ckpt_dir / 'trained'}")
print(f"Parameters: {sum(x.size for x in jax.tree.leaves(restored_params)):,}")

Вызов jax.device_put имеет значение. Сервер использует один графический процессор, поэтому параметры перемещаются в gpu_devices[0] перед тем, как попасть в модель.

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

4. Вариант 1: запуск с помощью jax.jit

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

Для подачи данных запеките веса в замыкание. nnx.split разделяет модель на статическое определение графа и динамическое состояние, а скомпилированная функция захватывает оба состояния и принимает в качестве аргумента только tokens . Это делает скомпилированную функцию самодостаточной.

graphdef, model_state = nnx.split(model)

@jax.jit
def predict_jit(tokens):
    m = nnx.merge(graphdef, model_state)
    return m(tokens)

dummy_input = jnp.zeros((1, MAX_SEQ_LEN), dtype=jnp.int32)

start = time.perf_counter()
logits = block_tree(predict_jit(dummy_input))
first_call_ms = (time.perf_counter() - start) * 1000

times = []
for _ in range(100):
    start = time.perf_counter()
    logits = block_tree(predict_jit(dummy_input))
    times.append((time.perf_counter() - start) * 1000)

avg_ms = np.mean(times)

show_table(
    ["", "Latency (ms)"],
    [
        ("First call (compile + execute)", f"{first_call_ms:,.1f}"),
        ("Subsequent calls (avg of 100)", f"{avg_ms:.2f}"),
        ("Speedup", f"{first_call_ms / avg_ms:.0f}\u00d7"),
    ],
    title="JIT inference latency",
    aligns=["left", "right"],
)

print(f"\nOutput shape: {logits.shape} (batch=1, seq={MAX_SEQ_LEN}, vocab={VOCAB_SIZE})")

Вы должны увидеть таблицу задержек JIT-вычислений, в которой первый вызов выполняется значительно медленнее, чем в среднем последующие 100, а затем Output shape: (1, 256, 256) .

5. Вариант 2: устранение холодного старта с помощью AOT-компиляции.

jax.jit компилируется при первом вызове, что хорошо во время разработки и является проблемой при запуске. AOT-компиляция разделяет этот единственный шаг на отдельные этапы:

  1. Вы пишете обычную функцию Python/JAX, например, функцию прямого прохода модели.
  2. JAX отслеживает выполнение функции для конкретных форм и типов данных входных данных и преобразует её в промежуточное представление компилятора с помощью функции lower()
  3. StableHLO описывает вычисления с использованием аппаратно-независимых операций, таких как dot , reshape и reduce .
  4. XLA оптимизирует StableHLO и создает исполняемый файл, специфичный для конкретного устройства, для GPU, TPU или CPU с помощью compile() .

Промежуточное представление (IR) StableHLO — это представление вычислений JAX на уровне компилятора после того, как код Python был преобразован в переносимую, аппаратно-независимую программу. Оно описывает такие операции, как умножение матриц, изменение формы, сокращение и управление потоком выполнения, в форме, которую XLA может компилировать для различных бэкендов, включая GPU, TPU и CPU.

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

# Bake weights into a closure
def predict_closed(tokens):
    m = nnx.merge(graphdef, model_state)
    return m(tokens)

# Stage 1: Lower
abstract_tokens = jax.ShapeDtypeStruct((1, MAX_SEQ_LEN), jnp.int32)

lowered = predict_closed.lower(abstract_tokens)
print(f"Lowered to StableHLO ({len(lowered.as_text()):,} chars)")

# Stage 2: Compile
compiled = lowered.compile()
print(f"Compiled for: {jax.default_backend()}")

# Execute
start = time.perf_counter()
logits_aot = block_tree(compiled(dummy_input))
aot_first_ms = (time.perf_counter() - start) * 1000

times_aot = []
for _ in range(100):
    start = time.perf_counter()
    logits_aot = block_tree(compiled(dummy_input))
    times_aot.append((time.perf_counter() - start) * 1000)

avg_aot_ms = np.mean(times_aot)
max_diff_jit_aot = float(jnp.max(jnp.abs(logits - logits_aot)))

show_table(
    ["", "Latency (ms)"],
    [
        ("AOT first execution (no compile)", f"{aot_first_ms:.2f}"),
        ("AOT subsequent (avg of 100)", f"{avg_aot_ms:.2f}"),
        ("JIT first call (from above)", f"{first_call_ms:,.1f}"),
        ("Max |JIT − AOT|", f"{max_diff_jit_aot:.2e}"),
    ],
    title="AOT vs JIT latency",
    aligns=["left", "right"],
)

Обратите внимание, что lower() принимает jax.ShapeDtypeStruct , а не реальные данные. Для компиляции вам никогда не понадобится входной массив — только его форма и тип данных.

В таблице задержек AOT и JIT вы должны увидеть, что время первого выполнения AOT близко к значению в установившемся режиме AOT, а не к значению первого вызова JIT, и Max |JIT − AOT| фактически равно нулю.

Проверьте ИК-сигнал StableHLO.

lowered.as_text() отображает программу StableHLO, которая будет выполняться на устройстве. Это то же промежуточное представление, которое XLA использует на графических процессорах, тензорных процессорах и центральных процессорах, и оно полезно для отладки, анализа производительности и понимания того, что на самом деле видит компилятор.

Текст в функции stableHLO может быть очень большим, и одна строка может содержать длинную константу или атрибут. Чтобы обойти ограничение скорости передачи данных в Jupyter IOPub, следующий код сохраняет весь промежуточный результат на диск и выводит только ограниченный предварительный просмотр.

hlo_text = lowered.as_text()

hlo_path = pathlib.Path("/tmp/jax-course/l8-stablehlo.mlir")
hlo_path.parent.mkdir(parents=True, exist_ok=True)
hlo_path.write_text(hlo_text)

MAX_LINES = 20
MAX_CHARS_PER_LINE = 160

lines = hlo_text.splitlines()
preview_lines = []
for line in lines[:MAX_LINES]:
    if len(line) > MAX_CHARS_PER_LINE:
        preview_lines.append(line[:MAX_CHARS_PER_LINE] + " ... [line truncated]")
    else:
        preview_lines.append(line)

print(f"StableHLO program: {len(lines):,} lines, {len(hlo_text):,} chars")
print(f"Full StableHLO saved to: {hlo_path}")
print("=" * 60)
print("\n".join(preview_lines))
print(
    f"\n... ({max(len(lines) - MAX_LINES, 0):,} more lines; "
    "long lines are truncated in this preview)"
)

Вы должны увидеть количество строк и символов, отображаемое программой, путь к сохраненному файлу .mlir , а также первые 20 строк промежуточного представления (IR), за которыми следует количество пропущенных строк.

6. Измерьте пропускную способность пакетного вывода.

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

batch_sizes = [1, 4, 16, 64]
results = []

for bs in batch_sizes:
    tokens_batch = jnp.zeros((bs, MAX_SEQ_LEN), dtype=jnp.int32)

    # Compile for this batch size
    lowered_bs = predict_closed.lower(
        jax.ShapeDtypeStruct((bs, MAX_SEQ_LEN), jnp.int32),
    )
    compiled_bs = lowered_bs.compile()

    # Warmup
    block_tree(compiled_bs(tokens_batch))

    # Measure
    times_bs = []
    for _ in range(50):
        start = time.perf_counter()
        block_tree(compiled_bs(tokens_batch))
        times_bs.append((time.perf_counter() - start) * 1000)

    avg_bs = np.mean(times_bs)
    tokens_per_sec = (bs * MAX_SEQ_LEN) / (avg_bs / 1000)
    results.append((bs, f"{avg_bs:.2f}", f"{tokens_per_sec:,.0f}"))

show_table(
    ["Batch size", "Latency (ms)", "Tokens/sec"],
    results,
    title="Batched inference throughput",
    aligns=["right", "right", "right"],
)

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

7. Вариант 3: экспорт переносимого артефакта с помощью jax.export

Функция jax.export экспортирует скомпилированную JAX-функцию в объект Exported содержащий `StableHLO` и метаданные, необходимые для вызова этой функции из другого JAX-процесса. Сериализованные байты могут быть следующими:

  • Сохранено на диск и загружено в другом процессе.
  • Вызывается без исходного кода модели на Python.
  • Экспортируется для текущей платформы по умолчанию или для конкретных платформ с помощью аргумента platforms=[...]

Это путь развертывания, изначально предназначенный для JAX — зависимость от TensorFlow не требуется.

from jax import export

# Export the closure-based function
exported = export.export(predict_closed)(
    jax.ShapeDtypeStruct((1, MAX_SEQ_LEN), jnp.int32),
)

print(f"Exported function: {exported.fun_name}")
print(f"Input shapes:  {exported.in_avals}")
print(f"Output shapes: {exported.out_avals}")
print(f"Exported platforms: {exported.platforms}")

# Serialize to bytes
blob = exported.serialize()
export_path = pathlib.Path("/tmp/jax-course/exports")
export_path.mkdir(parents=True, exist_ok=True)

export_file = export_path / "tiny_transformer_jax_export.bin"
export_file.write_bytes(blob)
print()
print(f"Serialized to {export_file} ({len(blob):,} bytes, {len(blob) / 1024:.0f} KB)")

# Deserialize and call
rehydrated = export.deserialize(export_file.read_bytes())

test_input = jnp.zeros((1, MAX_SEQ_LEN), dtype=jnp.int32)
logits_exported = block_tree(rehydrated.call(test_input))
print()
print(f"✅ Deserialized call succeeded — output shape: {logits_exported.shape}")

# Verify outputs match
diff_export = float(jnp.max(jnp.abs(logits_aot - logits_exported)))
print(f"Max difference from AOT: {diff_export:.2e}")

rehydrated объект так и не был обработан TinyTransformer . Он был воссоздан из байтов на диске и по-прежнему возвращает те же логиты, что и является основной целью портативного артефакта.

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

8. Вариант 4: преобразование в сохраненную модель TensorFlow с помощью jax2tf

Если ваша инфраструктура для обслуживания приложений использует TensorFlow, например, TF Serving или конвейеры TFX, вы можете преобразовать функцию JAX в сохраненную модель TensorFlow. jax2tf по-прежнему находится в папке jax.experimental , но это стандартный путь взаимодействия JAX и TensorFlow.

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

Здесь важна одна деталь платформы. В этом практическом задании JAX работает на CUDA, но среда выполнения TensorFlow в контейнере JAX может выполнять SavedModel на ЦП. Чтобы избежать вызова модуля, экспортируемого в CUDA, из TensorFlow на ЦП, приведенный ниже код экспортирует модуль jax2tf для ("cpu",) . Если вы используете TensorFlow на графическом процессоре, экспортируйте модуль для ("cuda",) и используйте среду выполнения TensorFlow с поддержкой GPU/XLA.

from jax.experimental import jax2tf
import tensorflow as tf
import shutil

# Capture model_state as a closure
def predict_for_tf(tokens):
    m = nnx.merge(graphdef, model_state)
    return m(tokens)

TF_EXPORT_PLATFORMS = ("cpu",)
tf_predict = jax2tf.convert(
    predict_for_tf,
    native_serialization_platforms=TF_EXPORT_PLATFORMS,
)

# Wrap in a tf.Module for SavedModel export
module = tf.Module()
module.predict = tf.function(
    tf_predict,
    input_signature=[tf.TensorSpec(shape=(1, MAX_SEQ_LEN), dtype=tf.int32)],
    autograph=False,
)

# TF Serving expects a versioned model directory
savedmodel_base_dir = pathlib.Path("/tmp/jax-course/exports") / "tiny_transformer_savedmodel"
savedmodel_dir = savedmodel_base_dir / "1"
if savedmodel_base_dir.exists():
    shutil.rmtree(savedmodel_base_dir)

tf.saved_model.save(module, str(savedmodel_dir))
print(f"✅ SavedModel saved to {savedmodel_dir}")
print(f"Exported for TensorFlow platform(s): {TF_EXPORT_PLATFORMS}")

# Verify the SavedModel path produces the same logits as the JAX path
with tf.device("/CPU:0"):
    tf_logits = module.predict(tf.zeros((1, MAX_SEQ_LEN), dtype=tf.int32))
diff_tf = np.max(np.abs(np.asarray(tf_logits) - np.asarray(logits_aot)))
print(f"Max difference from AOT: {diff_tf:.2e}")

# List saved files
for dirpath, _, filenames in os.walk(savedmodel_base_dir):
    for f in filenames:
        full = os.path.join(dirpath, f)
        size = os.path.getsize(full)
        print(f"  {os.path.relpath(full, savedmodel_base_dir):44s} {size:>10,} bytes")
print()
print(
    f"To serve: docker run -p 8501:8501 "
    f"--mount type=bind,source={savedmodel_base_dir},target=/models/transformer "
    "-e MODEL_NAME=transformer tensorflow/serving "
    "--xla_cpu_compilation_enabled=true"
)

Вы должны увидеть путь к каталогу SavedModel, кортеж экспортируемой платформы, максимальное отклонение от логитов AOT, которое фактически равно нулю, и список файлов, записанных в каталог с указанием версии.

Последнее, что выводит код, — это команда ` docker run ... tensorflow/serving . Она носит иллюстративный характер и показывает, как указать TensorFlow Serving путь к созданному вами каталогу с версиями. В этом практическом занятии вы её не запускаете.

9. Сравните четыре пути обслуживания.

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

show_table(
    ["Path", "Format", "Dependencies", "Serving target", "Portable"],
    [
        ("jax.jit", "Cached in-process executable", "JAX", "Python server", "No"),
        ("AOT compile", "Compiled in-process executable", "JAX", "Python server startup / warmup", "No"),
        ("jax.export", "Serialized JAX export", "JAX runtime", "Compatible runtime for exported platform(s)", "Yes"),
        ("jax2tf", "TF SavedModel", "TensorFlow", "TF Serving", "Yes"),
    ],
    title="Serving path comparison",
)

# Show file sizes
sizes = []
export_size = os.path.getsize(export_file)
sizes.append(("jax.export", f"{export_size:,} bytes", f"{export_size / 1024:.0f} KB"))

sm_size = sum(
    os.path.getsize(os.path.join(dirpath, f))
    for dirpath, _, filenames in os.walk(savedmodel_base_dir)
    for f in filenames
)
sizes.append(("jax2tf SavedModel", f"{sm_size:,} bytes", f"{sm_size / 1024:.0f} KB"))

show_table(
    ["Export", "Size (bytes)", "Size (KB)"],
    sizes,
    title="Export file sizes",
    aligns=["left", "right", "right"],
)

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

Выбирайте, исходя из ограничений развертывания, а не из архитектуры модели. Если JAX уже запущен в вашем процессе обслуживания, кратчайший путь — это jax.jit с предварительной загрузкой при запуске или компиляцией AOT. Если артефакт должен покинуть этот процесс, jax.export оставит вас внутри JAX, а jax2tf передаст вас в экосистему TensorFlow.

10. Уборка

Все, что было написано в этом практическом задании, хранится в каталоге /tmp на Pod'е и исчезает вместе с Pod'ом. Сначала скопируйте все, что хотите сохранить, выполнив следующую команду в Cloud Shell :

kubectl cp jax-jupyter:/tmp/jax-course/exports ./jax-course-exports

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

kubectl delete -f deploy/jupyter.yaml

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

cd terraform
terraform destroy

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

gcloud container clusters list
gcloud compute instances list

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

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

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

Что вы узнали

  • jax.jit — это самый простой путь: обернуть и вызвать. Первый вызов компилируется; последующие вызовы выполняются быстро. Хорошо подходит для серверной части на Python (FastAPI, Flask), где JAX уже установлен.
  • Компиляция AOT ( lower() затем compile() ) отделяет компиляцию от выполнения. Компиляция выполняется один раз при запуске, после чего приложение запускается без задержки компиляции первого запроса. lowered.as_text() отображает промежуточное представление StableHLO для отладки и анализа производительности.
  • jax.export сериализует JIT-скомпилированную функцию в JAX-нативный артефакт, содержащий StableHLO и метаданные вызова. Полученный файл может быть загружен и вызван совместимой средой выполнения JAX без исходного кода модели. По умолчанию экспортируется для текущей платформы; используйте platforms=[...] при необходимости указания конкретной целевой платформы.
  • jax2tf преобразует функцию в сохраненную модель TensorFlow. Используйте ее, если ваша инфраструктура обслуживания основана на TensorFlow. В текущих версиях JAX по умолчанию используется нативная сериализация; установите native_serialization_platforms в соответствии с тем, где TensorFlow будет выполнять модель.
  • Как восстановить контрольную точку Orbax в перестроенную архитектуру с помощью целей jax.ShapeDtypeStruct и как переместить параметры на одно устройство с помощью jax.device_put
  • Почему скорость передачи токенов в секунду не является авторегрессивной производительностью генерации, и почему для каждой новой формы входных данных требуется собственный скомпилированный исполняемый файл.

Краткий обзор курса

В восьми лабораториях вы найдете:

  1. L1–L3 : Настройка JAX, изучение jit компиляции и профилирование выполнения на GPU.
  2. L4 : Создал цикл обучения с нуля — оптимизатор, функция потерь, обновление градиента.
  3. L5 : Исследованы механизмы внимания — наивный, SDPA, объединенный механизм внимания cuDNN.
  4. L6 : Масштабирование до нескольких графических процессоров с параллельным распределением данных — Mesh, NamedSharding, параллельное сегментирование данных.
  5. L7 : Объединили все в языковую модель-трансформер — Flax NNX, Orbax, generation.
  6. L8 : Подготовлена ​​обученная модель для производства — JIT, AOT, jax.export , jax2tf

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

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

  • Добавьте декодирование в KV-кэше и пакетную генерацию, чтобы измерять реальную пропускную способность генерации, а не пропускную способность прямого прохода.
  • Обучение проводится на реальном токенизаторе и наборе данных, а не на используемом здесь игрушечном словаре.
  • Экспериментируйте с использованием смешанной точности и квантования, и измеряйте, а не предполагайте ускорение.
  • Масштабирование выходит за рамки параллельной обработки данных и включает параллельную обработку моделей и тензоров.
  • Создайте небольшую среду для обслуживания пользователей на основе одного из современных инструментов, включив в неё мониторинг и нагрузочные тесты.

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

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