torch.set_float32_matmul_precision
-
torch.set_float32_matmul_precision(precision)[source] -
Устанавливает внутреннюю точность для операций умножения матриц с плавающей точкой float32.
Выполнение операций умножения матриц float32 с меньшей точностью может значительно повысить производительность, а в некоторых программах потеря точности имеет незначительное влияние.
Поддерживаются три значения:
- “highest”, для внутренних вычислений операций умножения матриц float32 используется тип данных float32 (24 бита мантиссы).
- “high”, для внутренних вычислений операций умножения матриц float32 используется тип данных TensorFloat32 (10 бит мантиссы) или каждый float32-число рассматривается как сумма двух bfloat16-чисел (приблизительно 16 бит мантиссы), если доступны соответствующие быстрые алгоритмы умножения матриц. В противном случае операции умножения матриц float32 вычисляются так, как если бы точность была «highest». Более подробную информацию о подходе с bfloat16 см. ниже.
- “medium”, для внутренних вычислений операций умножения матриц float32 используется тип данных bfloat16 (8 бит мантиссы), если доступен быстрый алгоритм умножения матриц, использующий этот тип данных внутри. В противном случае операции умножения матриц float32 вычисляются так, как если бы точность была «high».
При использовании точности «high», операции умножения float32 могут использовать алгоритм на основе bfloat16, который сложнее, чем просто усечение до меньшего количества бит мантиссы (например, 10 для TensorFloat32, 8 для bfloat16). См. [Henry2019] для полного описания этого алгоритма. Краткий обзор: на первом шаге мы понимаем, что одно число float32 можно представить как сумму трех bfloat16-чисел (потому что float32 имеет 24 бита мантиссы, а bfloat16 — 8, и оба имеют одинаковое количество бит порядка). Это означает, что произведение двух float32-чисел можно точно представить как сумму девяти произведений bfloat16-чисел. Затем мы можем пожертвовать точностью для скорости, отбросив некоторые из этих произведений. Алгоритм с точностью «high» сохраняет только три наиболее значимых произведения, что удобно исключает все произведения, включающие последние 8 бит мантиссы любого из входных данных. Это означает, что мы можем представить входные данные как сумму двух bfloat16-чисел, а не трех. Поскольку инструкции bfloat16 fused-multiply-add (FMA) обычно в 10 раз быстрее инструкций float32, выполнение трех умножений и двух сложений с точностью bfloat16 быстрее, чем выполнение одного умножения с точностью float32.
-
Henry2019
Примечание
Это не меняет тип данных результата операций умножения матриц float32, а управляет тем, как выполняется внутреннее вычисление операции умножения матриц.
Примечание
Это не меняет точность операций свёртки. Другие флаги, такие как
torch.backends.cudnn.allow_tf32, могут контролировать точность операций свёртки.Примечание
Этот флаг в настоящее время влияет только на один тип устройства с нативными драйверами: CUDA. Если установлены значения «high» или «medium», при вычислении умножения матриц float32 будет использоваться тип данных TensorFloat32, что эквивалентно установке
torch.backends.cuda.matmul.allow_tf32 = True. При установке «highest» (значение по умолчанию), для внутренних вычислений используется тип данных float32, что эквивалентно установкеtorch.backends.cuda.matmul.allow_tf32 = False.- Parameters
-
precision (str) – может быть установлен в «highest» (по умолчанию), «high» или «medium» (см. выше).
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.set_float32_matmul_precision.html