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