Spec-Zone.ru › PyTorch 2.14

Автоматическая смешанная точность

Создано: 30 сентября 2025 г. | Последнее обновление: 30 сентября 2025 г.

Общие сведения

Автоматическая смешанная точность (AMP) позволяет использовать типы с плавающей точкой одинарной точности (32 бита) и половинной точности (16 бит) во время обучения или инференса.

Ключевые компоненты:

  • Автоматическое приведение типов: автоматически приводит операции к типам с меньшей точностью (например, float16 или bfloat16), чтобы повысить производительность при сохранении точности.
  • Масштабирование градиентов: динамически масштабирует градиенты при обратном распространении ошибки, предотвращая потерю значимости при обучении со смешанной точностью.

Проектирование

Стратегия приведения типов

CastPolicy используется для определения правил преобразования типов. Каждое значение перечисления представляет набор требований к преобразованию типов для группы операторов, обеспечивая единообразную обработку операций, в которых приоритет отдается точности или производительности.

Политика

Описание

lower_precision_fp

Перед выполнением операции привести все входные данные к lower_precision_fp.

fp32

Перед запуском операции привести все входные данные к at::kFloat.

fp32_set_opt_dtype

Выполнять в at::kFloat, учитывая указанный пользователем тип выходных данных, если он задан.

fp32_append_dtype

Добавить at::kFloat к аргументам и выполнить повторную диспетчеризацию на перегрузку, учитывающую тип данных

promote

Перед выполнением привести все входные данные к типу с наибольшей разрядностью.

Списки операторов

PyTorch определяет общий список операторов для каждой из перечисленных выше стратегий приведения типов в качестве справочного материала для разработчиков новых ускорителей.

Политика

Список операторов

lower_precision_fp

Ссылка на список

fp32

Ссылка на список

fp32_set_opt_dtype

Ссылка на список

fp32_append_dtype

Ссылка на список

promote

Ссылка на список

Реализация

Интеграция с Python

Реализуйте метод get_amp_supported_dtype, возвращающий типы данных, поддерживаемые новым ускорителем в контексте AMP.

1def get_amp_supported_dtype():
2    return [torch.float16, torch.bfloat16]
3
4

Интеграция с C++

В этом разделе показано, как AMP регистрирует ядра автоматического приведения типов для ключа диспетчеризации AutocastPrivateUse1.

  • Зарегистрируйте резервную реализацию, которая перенаправляет необработанные операции к их обычным реализациям.
  • Зарегистрируйте специальные ядра aten для AutocastPrivateUse1 с помощью вспомогательного макроса KERNEL_PRIVATEUSEONE, который сопоставляет операцию с нужной реализацией точности (с перечислением at::autocast::CastPolicy)
1TORCH_LIBRARY_IMPL(_, AutocastPrivateUse1, m) {
2  m.fallback(torch::CppFunction::makeFallthrough());
3}
 1TORCH_LIBRARY_IMPL(aten, AutocastPrivateUse1, m) {
 2  // lower_precision_fp
 3  KERNEL_PRIVATEUSEONE(mm, lower_precision_fp)
 4
 5  // fp32
 6  KERNEL_PRIVATEUSEONE(asin, fp32)
 7
 8  m.impl(
 9      TORCH_SELECTIVE_NAME("aten::binary_cross_entropy"),
10      TORCH_FN((&binary_cross_entropy_banned)));
11}

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/accelerator/amp.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API