Автоматическая смешанная точность
Создано: 30 сентября 2025 г. | Последнее обновление: 30 сентября 2025 г.
Общие сведения
Автоматическая смешанная точность (AMP) позволяет использовать типы с плавающей точкой одинарной точности (32 бита) и половинной точности (16 бит) во время обучения или инференса.
Ключевые компоненты:
- Автоматическое приведение типов: автоматически приводит операции к типам с меньшей точностью (например, float16 или bfloat16), чтобы повысить производительность при сохранении точности.
- Масштабирование градиентов: динамически масштабирует градиенты при обратном распространении ошибки, предотвращая потерю значимости при обучении со смешанной точностью.
Проектирование
Стратегия приведения типов
CastPolicy используется для определения правил преобразования типов. Каждое значение перечисления представляет набор требований к преобразованию типов для группы операторов, обеспечивая единообразную обработку операций, в которых приоритет отдается точности или производительности.
Политика | Описание |
|---|---|
| Перед выполнением операции привести все входные данные к |
| Перед запуском операции привести все входные данные к |
| Выполнять в |
| Добавить at::kFloat к аргументам и выполнить повторную диспетчеризацию на перегрузку, учитывающую тип данных |
| Перед выполнением привести все входные данные к типу с наибольшей разрядностью. |
Списки операторов
PyTorch определяет общий список операторов для каждой из перечисленных выше стратегий приведения типов в качестве справочного материала для разработчиков новых ускорителей.
Политика | Список операторов |
|---|---|
| |
| |
| |
| |
|
Реализация
Интеграция с 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