Функция hub.load
Функция hub.load загружает предобученную модель из указанного репозитория в PyTorch Hub.
Первым параметром передаётся строка с именем репозитория в формате "author/repo_name".
Вторым параметром - имя модели (entrypoint), определённое в репозитории.
Третьим параметром можно передать строку с тегом версии модели, если требуется конкретная версия.
Четвертым параметром передаются аргументы модели в виде именованных аргументов.
Синтаксис
torch.hub.load(repo, model, *args, **kwargs)
Пример
Загрузим предобученную модель ResNet-18 из официального репозитория PyTorch:
import torch
model = torch.hub.load('pytorch/vision', 'resnet18', pretrained=True)
print(type(model))
Результат выполнения кода:
<class 'torchvision.models.resnet.ResNet'>
Пример
Загрузим модель с указанием конкретной версии (тега) и передадим дополнительные аргументы:
import torch
model = torch.hub.load(
'pytorch/vision',
'resnet18',
pretrained=True,
progress=False
)
print(model)
Результат выполнения кода (сокращённый вывод):
ResNet(
(conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)
(bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
(relu): ReLU(inplace=True)
(maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1, dilation=1, ceil_mode=False)
(layer1): Sequential(
(0): BasicBlock(
...
Пример
Загрузим модель для классификации изображений и применим её к случайному тензору:
import torch
torch.manual_seed(0)
model = torch.hub.load('pytorch/vision', 'resnet18', pretrained=True)
model.eval()
# Создаём случайное изображение (batch_size=1, 3 канала, 224x224)
t = torch.randn(1, 3, 224, 224)
res = model(t)
print(res.shape)
Результат выполнения кода:
torch.Size([1, 1000])
Смотрите также
-
функцию
hub.list,
которая показывает доступные модели в репозитории -
функцию
hub.help,
которая выводит справку по модели и её аргументам -
функцию
hub.load_state_dict_from_url,
которая загружает состояние модели по URL -
функцию
save,
которая сохраняет модель или тензор в файл