РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
⊗pytoSpHbTr 82 of 82 menu
◀

Трассировка модуля в 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]] и выведите форму ответа на том же входе.

Создайте линейный слой с двумя входами и одним выходом, зафиксируйте его по строке из двух чисел и выведите форму выхода обёртки.

←
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить