torch.set_float32_matmul_precision
-
torch.set_float32_matmul_precision(precision)[исходный код] -
Задаёт внутреннюю точность умножения матриц типа float32.
Выполнение умножения матриц типа float32 с более низкой точностью может значительно повысить производительность, а в некоторых программах потеря точности практически не влияет на результат.
Поддерживаются три режима:
- «highest»: при умножении матриц типа float32 для внутренних вычислений используется тип данных float32 (24 бита мантиссы, из которых 23 хранятся явно).
- «high»: при умножении матриц типа float32 используется либо тип данных TensorFloat32 (10 бит мантиссы хранятся явно), либо каждое число float32 представляется как сумма двух чисел bfloat16 (примерно 16 бит мантиссы, из которых 14 хранятся явно), если доступны соответствующие быстрые алгоритмы умножения матриц. В противном случае умножение матриц типа float32 выполняется так, как если бы точность была «highest». Подробнее о подходе с bfloat16 см. ниже.
- «medium»: при умножении матриц типа float32 для внутренних вычислений используется тип данных bfloat16 (8 бит мантиссы, из которых 7 хранятся явно), если доступен быстрый алгоритм умножения матриц, использующий этот тип данных для внутренних вычислений. В противном случае умножение матриц типа float32 выполняется так, как если бы точность была «high».
При использовании точности «high» для умножения чисел float32 может применяться алгоритм на основе bfloat16, более сложный, чем простое усечение мантиссы до некоторого меньшего числа бит (например, до 10 для TensorFloat32 или до 7 для явно хранящихся битов bfloat16). Полное описание этого алгоритма приведено в [Henry2019]. Кратко поясним его здесь. Сначала заметим, что любое число float32 можно точно представить как сумму трёх чисел bfloat16 (поскольку у float32 мантисса содержит 23 бита, у bfloat16 явно хранятся 7 бит, а число бит экспоненты у них одинаково). Это означает, что произведение двух чисел float32 можно точно выразить как сумму девяти произведений чисел bfloat16. Затем можно обменять точность на скорость, отбросив некоторые из этих произведений. Алгоритм точности «high» оставляет только три произведения с наибольшим значением, что позволяет исключить все произведения, включающие последние 8 бит мантиссы любого из исходных чисел. Это означает, что исходные числа можно представить как сумму двух чисел bfloat16, а не трёх. Поскольку инструкции fused-multiply-add (FMA) для bfloat16 обычно более чем в 10 раз быстрее, чем для float32, три умножения и два сложения с точностью bfloat16 выполняются быстрее, чем одно умножение с точностью float32.
Примечание
Это не меняет выходной тип данных умножения матриц типа 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.- Параметры:
-
precision (str) – может принимать значения «highest» (по умолчанию), «high» или «medium» (см. выше).
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.set_float32_matmul_precision.html