Основные правила использования
Рассмотрим функцию, содержащую функции Токамакса, работающую на графическом процессоре 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 гарантирует два аспекта сериализации:
- Десериализованная функция, сериализованная для конкретного устройства, гарантированно будет работать именно на том устройстве, для которого она была сериализована.
- 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) .