Функция get_float32_matmul_precision
Функция get_float32_matmul_precision возвращает
строку, которая указывает текущий режим точности для
матричных умножений (matmul) с типом данных
float32. Режим определяет, будет ли PyTorch
использовать более быстрые вычисления с пониженной
точностью (например, tf32 на GPU) или сохранять
полную точность float32. Функция не принимает
никаких параметров.
Синтаксис
torch.get_float32_matmul_precision()
Пример
Давайте получим текущий режим точности по умолчанию:
import torch
precision = torch.get_float32_matmul_precision()
print(precision)
Результат выполнения кода:
"highest"
Пример
Давайте изменим режим точности с помощью функции
set_float32_matmul_precision и затем проверим его:
import torch
torch.set_float32_matmul_precision('high')
precision = torch.get_float32_matmul_precision()
print(precision)
Результат выполнения кода:
"high"
Пример
Давайте установим режим medium и проверим его:
import torch
torch.set_float32_matmul_precision('medium')
precision = torch.get_float32_matmul_precision()
print(precision)
Результат выполнения кода:
"medium"
Пример
Вот как режим точности влияет на выполнение матричного умножения:
import torch
torch.manual_seed(0)
torch.set_float32_matmul_precision('highest')
a = torch.randn(1000, 1000, dtype=torch.float32)
b = torch.randn(1000, 1000, dtype=torch.float32)
t1 = torch.matmul(a, b)
print(t1[0, 0])
Результат выполнения кода:
tensor(-11.6355)
Пример
Теперь установим режим high, который использует
tf32 для ускорения вычислений на современных GPU:
import torch
torch.manual_seed(0)
torch.set_float32_matmul_precision('high')
a = torch.randn(1000, 1000, dtype=torch.float32)
b = torch.randn(1000, 1000, dtype=torch.float32)
t2 = torch.matmul(a, b)
print(t2[0, 0])
Результат выполнения кода:
tensor(-11.6355)
Смотрите также
-
функцию
set_float32_matmul_precision,
которая устанавливает режим точности вычислений -
функцию
is_tf32_supported,
которая проверяет поддержку tf32 на текущем устройстве -
функцию
is_bf16_supported,
которая проверяет поддержку bf16 на текущем устройстве -
функцию
allow_tf32,
которая управляет использованием tf32 в cuDNN