Автонастройка и эвристика
Производительность ядра может существенно зависеть от специфических для реализации гиперпараметров (например, размеров тайлов, компоновки и т. д.), известных в 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:
...