Сравнительный анализ
Основы бенчмаркинга
Tokamax предоставляет интегрированную инфраструктуру для сравнительного анализа производительности. Вы можете оценить производительность вашей операции f(x) для заданного набора входных данных следующим образом:
# Conforming and initializing data to prepare for benchmarking
f_std, args = tokamax.standardize_function(f, kwargs={'x': x})
# Execute and measure timing
bench: tokamax.BenchmarkData = tokamax.benchmark(f_std, args)
standardize_function упрощает сложные функции, например, с аргументами, не являющимися массивами. Сначала она создает стандартную форму с одним аргументом ` args , который представляет собой список абстрактных или конкретных массивов jax.Array | jax.ShapeDtypeStruct . Затем она случайным образом инициализирует все абстрактные тензоры и возвращает стандартизированную f_std(args) только с конкретными аргументами в виде массивов. Ее можно легко скомпилировать с помощью JIT-компилятора, не беспокоясь о статических аргументах, таких как строки.
Полный список поддерживаемых каждой из этих функций параметров см. в соответствующих документах (docstrings); некоторые ключевые темы обсуждаются ниже.
Темы продвинутого бенчмаркинга
Выполнить итерации
tokamax.benchmark позволяет выбрать количество итераций; большее количество итераций обычно приводит к снижению шума измерений, например, tokamax.benchmark(f_std, args, iterations=num_iters). Однако, если количество итераций слишком велико за короткий промежуток времени, может сработать тепловое дросселирование, особенно для ресурсоемких ядер, что повлияет на время выполнения. Балансировка этих факторов часто является эмпирическим процессом. Предлагаемый подход заключается в выполнении небольшого количества итераций в каждом эксперименте с несколькими экспериментами, проводимыми с интервалами, ценой увеличения времени выполнения.
Метод сравнительного анализа
Накладные расходы JAX Python часто намного превышают фактическое время выполнения ядра ускорителя. Это означает, что обычный подход с измерением времени с помощью jax.block_until_ready(f(x)) будет бесполезен. benchmark позволяет выбрать базовую методологию измерения времени, используемую для бенчмаркинга, например, benchmark(f_std, args, iterations=num_iters, method=method)
Для ядер на базе TPU мы настоятельно рекомендуем использовать method=xprof_hermetic , который запускает профилировщик XProf и измеряет время выполнения на оборудовании. Этот метод практически не накладывает дополнительных затрат на инструментирование благодаря поддержке полного стека, включая оборудование и компилятор.
Для ядер GPU вы также можете использовать xprof_hermetic ; XProf, в свою очередь, использует API CUPTI от NVIDIA. Вы также можете напрямую вызвать таймер CUPTI с помощью method=cupti . Оба метода вводят некоторые переменные накладные расходы, обычно до 5%.
Распределение данных
Предыдущие исследования показали, что производительность может значительно варьироваться в зависимости от распределения данных из-за сложных взаимодействий на аппаратном уровне, связанных с энергопотреблением и тепловыделением. Для решения этой проблемы standardize_function инициализирует входные массивы способом, репрезентативным для реальных задач обучения, например, любой вещественный jax.ShapeDtypeStruct будет инициализирован случайным образом. Вы можете адаптировать это под свои нужды.
Режим бенчмаркинга
standardize_function позволяет выбрать mode=forward для выбора только прямого прохода, forward_res для выбора прямого прохода и вычисления остатков, vjp для выбора только функции VJP и forward_and_vjp для вычисления полного прямого и VJP прохода для бенчмаркинга. Обратите внимание, что бенчмаркинг только с использованием VJP может привести к ошибкам нехватки памяти (OOM), поскольку прямой проход вычисляется вне возвращаемой стандартизированной функции, а все промежуточные значения закладываются в HLO, который остается в памяти HBM.