Функция mps.synchronize
Функция mps.synchronize принудительно синхронизирует все операции,
выполняемые на MPS-устройстве (Metal Performance Shaders). Она блокирует
выполнение основного потока до тех пор, пока все запущенные на устройстве
операции не будут завершены. Функция не принимает параметров и полезна
для отладки, измерения времени выполнения операций и обеспечения
корректного порядка выполнения операций при работе с MPS.
Синтаксис
torch.mps.synchronize()
Пример
Давайте создадим тензор на MPS-устройстве и выполним несколько операций, а затем синхронизируем их:
import torch
if torch.mps.is_available():
device = torch.device('mps')
t = torch.tensor([1, 2, 3, 4, 5], device=device)
t = t * 2
t = t + 10
torch.mps.synchronize()
print(t)
else:
print("MPS device not available")
Результат выполнения кода:
tensor([12, 14, 16, 18, 20], device='mps:0')
Пример
Давайте измерим время выполнения операций на MPS-устройстве с синхронизацией:
import torch
import time
if torch.mps.is_available():
device = torch.device('mps')
t = torch.randn(1000, 1000, device=device)
start = time.time()
res = t.matmul(t.t())
torch.mps.synchronize()
end = time.time()
print(f"Execution time: {end - start:.4f} seconds")
print(f"Result shape: {res.shape}")
else:
print("MPS device not available")
Результат выполнения кода:
"Execution time: 0.0234 seconds"
"Result shape: torch.Size([1000, 1000])"
Пример
Давайте убедимся, что операции выполняются асинхронно без синхронизации, и сравним с синхронизацией:
import torch
import time
if torch.mps.is_available():
device = torch.device('mps')
t = torch.randn(5000, 5000, device=device)
# Без синхронизации
start = time.time()
res1 = t @ t.T
end = time.time()
print(f"Without sync: {end - start:.4f}s")
# С синхронизацией
start = time.time()
res2 = t @ t.T
torch.mps.synchronize()
end = time.time()
print(f"With sync: {end - start:.4f}s")
else:
print("MPS device not available")
Результат выполнения кода:
"Without sync: 0.0010s"
"With sync: 0.0156s"
Смотрите также
-
функцию
mps.is_available,
которая проверяет доступность MPS-устройства -
функцию
mps.is_built,
которая проверяет поддержку MPS в сборке PyTorch -
функцию
synchronize,
которая синхронизирует операции на CUDA-устройстве -
функцию
empty_cache,
которая освобождает неиспользуемую память на устройстве