Функция set_float32_matmul_precision
Функция set_float32_matmul_precision управляет точностью вычислений при выполнении матричных операций (например, torch.mm, torch.bmm, torch.matmul) для тензоров типа float32 на CUDA-устройствах. Она позволяет выбрать между более высокой производительностью и повышенной точностью. Первым и единственным параметром функция принимает строку, указывающую режим точности: "highest", "high" или "medium".
Синтаксис
torch.set_float32_matmul_precision(precision)
Режимы точности
Функция поддерживает три режима точности:
"highest" - максимальная точность. Использует алгоритм tf32 только для тензоров с размерностью не менее 3, для двумерных матриц использует float32. Обеспечивает наилучшую точность за счёт производительности.
"high" - высокая точность. По умолчанию использует tf32 для всех матричных умножений. Это режим по умолчанию в PyTorch.
"medium" - средняя точность. Использует tf32 для всех матричных операций, что даёт максимальную производительность на CUDA-устройствах с поддержкой tf32.
Пример 1
Давайте посмотрим текущий режим точности с помощью функции get_float32_matmul_precision:
import torch
current_precision = torch.get_float32_matmul_precision()
print(current_precision)
Результат выполнения кода:
"highest"
Пример 2
Установим режим точности "medium" для повышения производительности на CUDA:
import torch
torch.set_float32_matmul_precision("medium")
print(torch.get_float32_matmul_precision())
Результат выполнения кода:
"medium"
Пример 3
Установим максимальную точность для критических вычислений:
import torch
torch.set_float32_matmul_precision("highest")
print(torch.get_float32_matmul_precision())
t = torch.tensor([[1.0, 2.0], [3.0, 4.0]])
res = torch.mm(t, t)
print(res)
Результат выполнения кода:
"highest"
tensor([
[ 7.0000, 10.0000],
[15.0000, 22.0000],
])
Пример 4
Сравним точность вычислений в разных режимах на примере умножения матриц:
import torch
torch.manual_seed(0)
t = torch.randn(1000, 1000, device="cuda")
torch.set_float32_matmul_precision("highest")
res1 = torch.matmul(t, t)
torch.set_float32_matmul_precision("medium")
res2 = torch.matmul(t, t)
diff = torch.abs(res1 - res2).max().item()
print(diff)
Результат выполнения кода:
0.0001220703125
Смотрите также
-
функцию
get_float32_matmul_precision,
которая возвращает текущий режим точности матричных вычислений -
функцию
is_tf32_supported,
которая проверяет поддержку формата tf32 на текущем устройстве -
функцию
allow_tf32,
которая управляет использованием tf32 в cuDNN -
функцию
is_available,
которая проверяет доступность CUDA-устройств