Масштабирование и запуск JAX-моделей на графических процессорах NVIDIA.

  1. Что лучше всего проверить в первую очередь перед отладкой производительности JAX GPU?

  2. Если jax.default_backend() выводит cpu , а jax.devices() показывает только CpuDevice , какова наиболее вероятная причина?

  3. Почему первый вызов функции jax.jit может быть намного медленнее, чем последующие?

  4. Какова распространённая причина случайной перекомпиляции в JAX?

  5. Почему внутри функции @jax.jit условие if x > 0: на отслеживаемом массиве вызывает ошибку?

  6. Почему необходимо вызывать block_until_ready() перед остановкой таймера во время работы с графическим процессором?

  7. Вызов функции float(loss) на каждом шаге внутри цикла обучения может замедлить его, потому что это:

  8. Почему в процессе обучения данные преобразуются в пакеты фиксированного размера, а неполный итоговый пакет отбрасывается?

  9. Как обновляются параметры на этапе обучения JAX?

  10. Что требуется для установки параметра implementation="cudnn" в jax.nn.dot_product_attention ?

  11. Почему наивный механизм внимания использует больше памяти, чем ядро ​​cuDNN с интегрированным модулем на длинных последовательностях?

  12. Что же на самом деле создает параллелизм при параллельном обучении с использованием JAX?

  13. Что означает PartitionSpec('data', None) для двумерного массива на сетке с осью 'data' ?

  14. Как преобразователь из Урока 7 получает причинно-следственное (только обратное) внимание?

  15. Что сохраняется, а что нет при сохранении модели NNX с помощью Orbax StandardCheckpointer ?

  16. Какую проблему решает AOT-компиляция ( lower()compile() ) при обслуживании?

  17. Какой путь развертывания создает переносимый артефакт, который работает в любой среде выполнения JAX без исходного кода модели?