Основные правила использования

Рассмотрим функцию, содержащую функции Токамакса, работающую на графическом процессоре H100:

import jax
import jax.numpy as jnp
import tokamax

def loss(x, scale):
  x = tokamax.layer_norm(
      x, scale=scale, offset=None, implementation="triton"
  )
  x = tokamax.dot_product_attention(x, x, x, implementation="xla_chunked")
  x = tokamax.layer_norm(x, scale=scale, offset=None, implementation=None)
  x = tokamax.dot_product_attention(x, x, x, implementation="mosaic")
  return jnp.sum(x)

f_grad = jax.jit(jax.grad(loss))

При использовании implementation=None Tokamax может выбрать наилучшую реализацию для каждой формы ядра. Он даже может выбирать разные реализации для прямого прохода и градиента. Кроме того, это всегда будет поддерживаться, поскольку он может вернуться к реализации XLA implementation='xla' .

Однако, возможно, вам потребуется выбрать конкретную реализацию ядра и потерпеть неудачу, если она не поддерживается. Например, implementation="mosaic" попытается использовать ядро ​​Pallas:Mosaic для графического процессора, если это возможно, и выбросит исключение, если это не поддерживается по какой-либо причине. Например, использование входных данных FP64 не поддерживается или используются более старые графические процессоры.

Оцените градиент

channels, seq_len, batch_size, num_heads = (64, 2048, 32, 16)
scale = jax.random.normal(jax.random.key(0), (channels,), dtype=jnp.float32)
x = jax.random.normal(
    jax.random.key(1),
    (batch_size, seq_len, num_heads, channels),
    dtype=jnp.bfloat16,
)

out = f_grad(x, scale)

Автотюнинг

Для достижения наилучшей производительности выполните автоматическую настройку всех ядер Tokamax в f_grad :

autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, x, scale)

autotune_result можно использовать в качестве менеджера контекста, применяя автоматически настроенные конфигурации для всех ядер Tokamax в f_grad :

with autotune_result:
  out_autotuned = f_grad(x, scale)

Для сериализации и повторного использования результата потенциально ресурсоемкого вызова tokamax.autotuning :

autotune_result_json: str = autotune_result.dumps()
autotune_result = tokamax.AutotuningResult.loads(autotune_result_json)

Пользователи могут самостоятельно настраивать свои ядра с помощью tokamax.autotune , унаследовав класс от tokamax.Op и переопределив метод tokamax.Op._get_autotuning_configs для определения пространства поиска для автоматической настройки.

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

Сериализация

Ядра могут быть сериализованы в StableHLO . Вызовы ядер — это пользовательские вызовы JAX, которые по умолчанию запрещены в jax.export , поэтому для разрешения экспорта всех ядер Tokamax необходимо использовать tokamax.DISABLE_JAX_EXPORT_CHECKS :

from jax import export

f_grad_exported = export.export(f_grad, disabled_checks=tokamax.DISABLE_JAX_EXPORT_CHECKS)(
    jax.ShapeDtypeStruct(x.shape, x.dtype),
    jax.ShapeDtypeStruct(scale.shape, scale.dtype),
)

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

  1. Десериализованная функция, сериализованная для конкретного устройства, гарантированно будет работать именно на том устройстве, для которого она была сериализована.
  2. Tokamax предоставляет те же гарантии совместимости, что и JAX : обратная совместимость в течение 6 месяцев.

Сравнительный анализ

Накладные расходы JAX Python часто намного превышают фактическое время выполнения ядра ускорителя. Это означает, что обычный подход с измерением времени выполнения jax.block_until_ready(f_grad(x, scale)) будет бесполезен. В Tokamax есть утилиты только для измерения времени выполнения ускорителя:


f_std, args = tokamax.benchmarking.standardize_function(f, kwargs={'x': x, 'scale': scale})
run = tokamax.benchmarking.compile_benchmark(f_std, args)
bench: tokamax.benchmarking.BenchmarkData = run(args)

Существуют различные методы измерения: например, на GPU используется профилировщик CUPTI , который можно указать с помощью run(args, method='cupti') . Он инструментирует ядро ​​и добавляет небольшие накладные расходы. Команда run(args, method=None) по умолчанию позволяет Tokamax выбрать метод и работает как для TPU, так и для GPU. Шум в бенчмарке можно уменьшить, увеличив количество итераций run(args, iterations=10) .