Функция hub.load_state_dict_from_url
Функция hub.load_state_dict_from_url загружает словарь состояний модели из удаленного репозитория по указанному URL. Первым параметром функция принимает строку с URL-адресом. Вторым параметром можно передать имя файла для сохранения. Функция также поддерживает параметры для проверки контрольной суммы и прогресса загрузки.
Синтаксис
torch.hub.load_state_dict_from_url(url, [model_dir], [map_location], [progress], [check_hash], [file_name])
Пример
Давайте загрузим словарь состояний модели из удаленного репозитория:
import torch
url = 'https://example.com/models/model.pt'
state_dict = torch.hub.load_state_dict_from_url(url)
print(type(state_dict))
Результат выполнения кода:
<class 'collections.OrderedDict'>
Пример
Давайте загрузим словарь состояний с указанием директории для сохранения и прогрессом загрузки:
import torch
url = 'https://example.com/models/model.pt'
state_dict = torch.hub.load_state_dict_from_url(
url,
model_dir='/path/to/models',
progress=True
)
print(type(state_dict))
Результат выполнения кода:
<class 'collections.OrderedDict'>
Пример
Давайте загрузим словарь состояний с проверкой контрольной суммы и указанием имени файла:
import torch
url = 'https://example.com/models/model.pt'
state_dict = torch.hub.load_state_dict_from_url(
url,
file_name='my_model.pt',
check_hash=True
)
print(type(state_dict))
Результат выполнения кода:
<class 'collections.OrderedDict'>
Пример
Давайте загрузим словарь состояний с перемещением на GPU:
import torch
url = 'https://example.com/models/model.pt'
state_dict = torch.hub.load_state_dict_from_url(
url,
map_location=torch.device('cuda')
)
print(type(state_dict))
Результат выполнения кода:
<class 'collections.OrderedDict'>