Автонастройка и эвристика

Производительность ядра может существенно зависеть от специфических для реализации гиперпараметров (например, размеров тайлов, компоновки и т. д.), известных в Tokamax как операционные конфигурации (Op Configs). Оптимальные значения этих конфигураций должны определяться эмпирически для каждого интересующего входного параметра с помощью автонастройки . Tokamax предоставляет фреймворк как для выполнения автонастройки, так и для упрощения управления этим процессом и его результатами.

Общий обзор

Каждая реализация tokamax.Op определяет свой набор настраиваемых гиперпараметров и диапазон значений, которые может принимать каждая из этих конфигураций, в tokamax.Op._get_autotuning_configs ( пример ). Автонастройщик выполняет исчерпывающий поиск в этом пространстве, чтобы определить конфигурацию с наилучшей производительностью (наименьшее время выполнения). Результат может быть использован различными способами: непосредственно в вашей программе в качестве менеджера контекста, сериализован и сохранен в вашем личном кэше для последующего использования или как часть глобального кэша библиотеки.

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

Форматы ввода автотюнера

Вызываемая функция

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

# Autotune all Tokamax kernels in f, which can be a JAX function with multiple Tokamax ops and non-Tokamax ops as well. Assume f takes a dictionary of args as input.
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, **args)
# Best possible result
with autotune_result:
    out_autotuned = f(**args)

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

# This pullback function contains Tokamax VJP ops.
f_grad = jax.grad(f)
autotune_result = tokamax.autotune(f_grad)

with autotune_result:
  out = f_grad()

Последовательность BoundArguments

Когда вызывается операция Tokamax с набором входных аргументов, операция сначала «привязывается» к этим аргументам. Этот процесс канонизирует и проверяет входные данные, создавая объект BoundArgument , состоящий из самой операции и её входных аргументов. Автотюнер может принимать последовательность объектов BoundArguments в качестве входных данных и выполняет автонастройку каждого из них.

ragged_dot_ba = tokamax.PallasMosaicTpuRaggedDot.bind(x, y, group)
attention_ba = tokamax.PallasMosaicTpuAttention.bind(q, k, v)

autotune_result: tokamax.AutotuningResult = tokamax.autotune([ragged_dot_ba, attention_ba])

with autotune_result:
  ...

Сериализованный список BoundArguments

Скрипт xplane_to_bound_args.py принимает на вход прототип Xplane, сгенерированный XProf во время профилирования, и выдает JSON-файл с BoundArguments . Автонастройка может прочитать этот JSON-файл и выполнить автоматическую настройку всех операций Tokamax, которые отображаются в профиле.

Предположим, что в результате профилирования у вас есть файл xplane my_model.xplane.pb . Вы можете извлечь и автоматически настроить все используемые операции Tokamax следующим образом:

 python tokamax/_src/tools/xplane_to_bound_args.py \
   --xplane_file=my_model.xplane.pb \
   --output_file=/tmp/bound_args.json
import tokamax

bound_args = tokamax.autotuning.bound_args_from_json("/tmp/bound_args.json")
autotune_result = tokamax.autotune(bound_args)

Использование выходных данных автотюнера

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

Менеджер контекста

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

# Autotune all Tokamax kernels in f, which can be a JAX function with multiple Tokamax ops and non-Tokamax ops as well. Assume f takes a dictionary of args as input.
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, **args)
# Best possible result
with autotune_result:
    out_autotuned = f(**args)

Последовательный AutotuningResult

Возможно, вам потребуется сериализовать и повторно использовать autotune_result по одной из двух причин: (1) автонастройка может быть дорогостоящей с точки зрения вычислительного времени, в зависимости от размера пространства поиска; (2) время выполнения ядра может быть нестабильным, что приводит к различным «оптимальным» вариантам в разных запусках и может привести к численной недетерминированности.

autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, **args)
autotune_result_json: str = autotune_result.dumps()
autotune_result = tokamax.AutotuningResult.loads(autotune_result_json)
with autotune_result:
    out_autotuned = f(**args)

Несколько объектов AutotuningResult также можно объединить в один объект, прочитав результаты в виде строки JSON с помощью load .

autotune_result_1 = tokamax.AutotuningResult.load(path_to_file_1)
autotune_result_2 = tokamax.AutotuningResult.load(path_to_file_2)
merged_results = autotune_result_1 | autotune_result_2

with merged_results:
  ...