Компиляция в PyTorch
Функция compile оборачивает
модуль и может собрать более
быстрый вариант вызова на
том устройстве, где модуль
работает. Первый прогон иногда
долго готовит граф.
Соберём крошечный модуль, обернём его и один раз подадим вход на процессоре:
import torch
import torch.nn as nn
class TinyShift(nn.Module):
def forward(self, x):
return x + 1
model = torch.compile(TinyShift())
input_row = torch.tensor([1.0, 2.0])
output = model(input_row)
print(output.shape) # выведет torch.Size([2])
Форма ответа совпадает с обычным модулем: меняется способ выполнения, а не размерность результата. Ускорение зависит от железа и размера модели.
Оберните модуль, который
удваивает вход, подайте
ряд [2.0, 3.0]
и выведите форму ответа.
Создайте линейный слой
с одним входом и двумя
выходами, оберните его
компилятором, подайте
число 1.0 и выведите
форму результата.
Оберните модуль, который
возвращает вход без
изменений, прогоните
таблицу [[1.0, 2.0]]
и выведите форму выхода.