Метод mps
Метод mps класса Module перемещает все параметры и буферы модели на устройство MPS,
что позволяет выполнять вычисления на GPU устройств Apple Silicon (M1, M2, M3 и новее).
Метод не принимает параметров и возвращает сам модуль, что позволяет использовать цепочки вызовов.
Синтаксис
model.mps()
Пример
Давайте создадим простую линейную модель и переместим её на устройство MPS:
import torch
import torch.nn as nn
model = nn.Linear(10, 5)
print(model.weight.device)
model = model.mps()
print(model.weight.device)
Результат выполнения кода:
device(type='cpu')
device(type='mps', index=0)
Пример
Давайте создадим модель, переместим её на MPS и выполним прямой проход с тензорами, также расположенными на MPS:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Sequential(
nn.Linear(5, 10),
nn.ReLU(),
nn.Linear(10, 3)
)
model.mps()
x = torch.randn(4, 5).mps()
y = model(x)
print(y.device)
print(y.shape)
Результат выполнения кода:
device(type='mps', index=0)
torch.Size([4, 3])
Пример
Метод mps поддерживает цепочки вызовов, что позволяет сразу переместить модель на устройство MPS
после её создания:
import torch
import torch.nn as nn
model = nn.Linear(8, 4).mps()
print(next(model.parameters()).device)
x = torch.randn(2, 8).mps()
res = model(x)
print(res.device)
Результат выполнения кода:
device(type='mps', index=0)
device(type='mps', index=0)