torch.hub
Pytorch Hub — это хранилище предобученных моделей, предназначенное для повышения воспроизводимости исследований.
Опубликование моделей
Pytorch Hub поддерживает публикацию предобученных моделей (определения моделей и предобученные веса) в репозитории github путем добавления простого файла hubconf.py;
hubconf.py может иметь несколько точек входа. Каждая точка входа определяется как функция Python (например, предобученная модель, которую вы хотите опубликовать).
def entrypoint_name(*args, **kwargs):
# args & kwargs are optional, for models which take positional/keyword arguments.
...
Как реализовать точку входа?
Вот фрагмент кода, определяющий точку входа для модели resnet18, если мы расширим реализацию в pytorch/vision/hubconf.py. В большинстве случаев достаточно импортировать соответствующую функцию в hubconf.py. Здесь мы просто хотим использовать расширенную версию в качестве примера, чтобы показать, как это работает. Полный скрипт можно посмотреть в репозитории pytorch/vision
dependencies = ['torch']
from torchvision.models.resnet import resnet18 as _resnet18
# resnet18 is the name of entrypoint
def resnet18(pretrained=False, **kwargs):
""" # This docstring shows up in hub.help()
Resnet18 model
pretrained (bool): kwargs, load pretrained weights into the model
"""
# Call the model, load pretrained weights
model = _resnet18(pretrained=pretrained, **kwargs)
return model
-
dependenciesпеременная представляет собой **список** имён пакетов, необходимых для **загрузки** модели. Обратите внимание, что это может немного отличаться от зависимостей, необходимых для обучения модели. -
argsиkwargsпередаются в реальную вызываемую функцию. - Документация функции используется в качестве сообщения справки. Она объясняет, что делает модель и какие допустимые позиционные/именованные аргументы. Настоятельно рекомендуется добавить несколько примеров.
- Функция точки входа может возвращать модель (nn.module) или вспомогательные инструменты, чтобы облегчить работу пользователя, например, токенизаторы.
- Вызываемые функции, начинающиеся с нижнего подчеркивания, считаются вспомогательными функциями, которые не будут отображаться в
torch.hub.list(). - Предобученные веса могут храниться локально в репозитории github или загружаться с помощью
torch.hub.load_state_dict_from_url(). Если размер меньше 2 ГБ, рекомендуется прикрепить его к выпуску проекта и использовать URL из выпуска. В приведенном выше примереtorchvision.models.resnet.resnet18обрабатываетpretrained, альтернативно, вы можете поместить следующий код в определение точки входа.
if pretrained:
# For checkpoint saved in local github repo, e.g. <RELATIVE_PATH_TO_CHECKPOINT>=weights/save.pth
dirname = os.path.dirname(__file__)
checkpoint = os.path.join(dirname, <RELATIVE_PATH_TO_CHECKPOINT>)
state_dict = torch.load(checkpoint)
model.load_state_dict(state_dict)
# For checkpoint saved elsewhere
checkpoint = 'https://download.pytorch.org/models/resnet18-5c106cde.pth'
model.load_state_dict(torch.hub.load_state_dict_from_url(checkpoint, progress=False))
Важное замечание
- Опубликованные модели должны быть как минимум в ветке/теге. Это не может быть произвольный коммит.
Загрузка моделей из Hub
Pytorch Hub предоставляет удобные API для просмотра всех доступных моделей в hub с помощью torch.hub.list(), отображения документации и примеров с помощью torch.hub.help() и загрузки предобученных моделей с помощью torch.hub.load().
-
torch.hub.list(github, force_reload=False, skip_validation=False, trust_repo=None)[source] -
Список всех доступных вызываемых точек входа в указанном репозитории
github.- Параметры:
-
-
github (str) – строка в формате «repo_owner/repo_name[:ref]» с необязательным ref (тег или ветка). Если
refне указан, предполагается использовать стандартную веткуmain(если она существует), в противном случаеmaster. Пример: ‘pytorch/vision:0.10’ -
force_reload (bool, optional) – указывает на то, следует ли отбросить существующий кэш и выполнить новую загрузку. По умолчанию
False. -
skip_validation (bool, optional) – если
False, torchhub проверит, что ветка или коммит, указанные в аргументеgithub, правильно принадлежат владельцу репозитория. Это потребует обращений к API GitHub; вы можете указать нестандартный токен GitHub, установив переменную средыGITHUB_TOKEN. По умолчаниюFalse. -
trust_repo (bool, str или None) –
"check",True,FalseилиNone. Этот параметр был добавлен в версии v1.12 и помогает гарантировать, что пользователи запускают код только из репозиториев, которым они доверяют.- Если
False, будет запрошен у пользователя подтверждение о доверии к репозиторию. - Если
True, репозиторий будет добавлен в список доверенных и загружен без явного подтверждения. - Если
"check", репозиторий будет проверен на наличие в списке доверенных репозиториев в кэше. Если его там нет, поведение вернётся к опцииtrust_repo=False. - Если
None: это вызовет предупреждение, предлагая пользователю установитьtrust_repoнаFalse,Trueили"check". Этот параметр присутствует только для обратной совместимости и будет удалён в версии v1.14.
По умолчанию
Noneи, в конечном счете, изменится на"check"в версии v1.14. - Если
-
github (str) – строка в формате «repo_owner/repo_name[:ref]» с необязательным ref (тег или ветка). Если
- Возвращаемое значение:
-
Доступные вызываемые точки входа
- Тип возвращаемого значения:
Пример
>>> entrypoints = torch.hub.list('pytorch/vision', force_reload=True)
-
torch.hub.help(github, model, force_reload=False, skip_validation=False, trust_repo=None)[source] -
Отображение документации точки входа
model.- Параметры:
-
-
github (str) – строка в формате <repo_owner/repo_name[:ref]> с необязательным ref (тег или ветка). Если
refне указан, предполагается использовать стандартную веткуmain(если она существует), в противном случаеmaster. Пример: ‘pytorch/vision:0.10’ -
model (str) – имя точки входа, определённое в
hubconf.pyрепозитория. -
force_reload (bool, optional) – указывает на то, следует ли отбросить существующий кэш и выполнить новую загрузку. По умолчанию
False. -
skip_validation (bool, optional) – если
False, torchhub проверит, что ссылка, указанная в аргументеgithub, принадлежит владельцу репозитория. Это потребует обращений к API GitHub; вы можете указать нестандартный токен GitHub, установив переменную средыGITHUB_TOKEN. По умолчаниюFalse. -
trust_repo (bool, str или None) –
"check",True,FalseилиNone. Этот параметр был добавлен в версии v1.12 и помогает гарантировать, что пользователи запускают код только из репозиториев, которым они доверяют.- Если
False, будет запрошен у пользователя подтверждение о доверии к репозиторию. - Если
True, репозиторий будет добавлен в список доверенных и загружен без явного подтверждения. - Если
"check", репозиторий будет проверен на наличие в списке доверенных репозиториев в кэше. Если его там нет, поведение вернётся к опцииtrust_repo=False. - Если
None: это вызовет предупреждение, предлагая пользователю установитьtrust_repoнаFalse,Trueили"check". Этот параметр присутствует только для обратной совместимости и будет удалён в версии v1.14.
По умолчанию
Noneи, в конечном счете, изменится на"check"в версии v1.14. - Если
-
github (str) – строка в формате <repo_owner/repo_name[:ref]> с необязательным ref (тег или ветка). Если
Пример
>>> print(torch.hub.help('pytorch/vision', 'resnet18', force_reload=True))
-
torch.hub.load(repo_or_dir, model, *args, source='github', trust_repo=None, force_reload=False, verbose=True, skip_validation=False, **kwargs)[source] -
Загрузка модели из репозитория github или локального каталога.
Примечание: Загрузка модели — это типичный случай использования, но это также можно использовать для загрузки других объектов, таких как токенизаторы, функции потерь и т. д.
Если
sourceравно ‘github’,repo_or_dirдолжно иметь форматrepo_owner/repo_name[:ref]с необязательным параметром ref (метка или ветка).Если
sourceравно ‘local’,repo_or_dirдолжно быть путём к локальному каталогу.- Параметры:
-
-
repo_or_dir (str) – Если
sourceравно ‘github’, это должно соответствовать репозиторию github в форматеrepo_owner/repo_name[:ref]с необязательным параметром ref (метка или ветка), например, ‘pytorch/vision:0.10’. Еслиrefне указан, предполагается, что используется ветка по умолчаниюmain(если она существует), в противном случаеmaster. Еслиsourceравно ‘local’, то это должен быть путь к локальному каталогу. -
model (str) – имя вызываемого объекта (точки входа), определённого в репозитории/каталоге в
hubconf.py. -
*args (необязательно) – соответствующие аргументы для вызываемого объекта
model. -
source (str, необязательно) – ‘github’ или ‘local’. Определяет, как интерпретировать
repo_or_dir. По умолчанию равно ‘github’. -
trust_repo (bool, str или None) –
"check",True,FalseилиNone. Этот параметр был добавлен в версии v1.12 и помогает гарантировать, что пользователи запускают код только из репозиториев, которым они доверяют.- Если
False, пользователь будет спрошен, следует ли доверять репозиторию. - Если
True, репозиторий будет добавлен в список доверенных и загружен без явного подтверждения. - Если
"check", репозиторий будет проверен на соответствие списку доверенных репозиториев в кэше. Если его нет в этом списке, поведение вернётся к вариантуtrust_repo=False. - Если
None: это выведет предупреждение, предложив пользователю установитьtrust_repoвFalse,Trueили"check". Этот параметр присутствует только для обратной совместимости и будет удалён в v1.14.
По умолчанию равно
Noneи в конечном итоге изменится на"check"в версии v1.14. - Если
-
force_reload (bool, необязательно) – следует ли принудительно перескачать репозиторий github. Не имеет эффекта, если
source = 'local'. По умолчаниюFalse. -
verbose (bool, необязательно) – Если
False, сообщения о попадании в локальный кэш будут заглушены. Примечание: сообщение о первой загрузке заглушить нельзя. Не имеет эффекта, еслиsource = 'local'. По умолчаниюTrue. -
skip_validation (bool, необязательно) – если
False, torchhub проверит, что ветка или коммит, указанные в параметреgithub, правильно принадлежат владельцу репозитория. Это потребует запросов к API GitHub; вы можете указать нестандартный токен GitHub, установив переменную средыGITHUB_TOKEN. По умолчаниюFalse. -
**kwargs (необязательно) – соответствующие ключевые аргументы для вызываемого объекта
model.
-
repo_or_dir (str) – Если
- Возвращает:
-
Результат вызова вызываемого объекта
modelс указанными*argsи**kwargs.
Пример
>>> # from a github repo >>> repo = 'pytorch/vision' >>> model = torch.hub.load(repo, 'resnet50', weights='ResNet50_Weights.IMAGENET1K_V1') >>> # from a local directory >>> path = '/some/local/path/pytorch/vision' >>> model = torch.hub.load(path, 'resnet50', weights='ResNet50_Weights.DEFAULT')
-
torch.hub.download_url_to_file(url, dst, hash_prefix=None, progress=True)[source] -
Загрузка объекта по заданному URL в локальный путь.
- Параметры:
-
- url (str) – URL объекта для загрузки
-
dst (str) – Полный путь, куда будет сохранён объект, например,
/tmp/temporary_file -
hash_prefix (str, необязательно) – Если не равно None, загруженный файл SHA256 должен начинаться с
hash_prefix. По умолчанию: None - progress (bool, необязательно) – отображать ли индикатор прогресса в stderr. По умолчанию: True
Пример
>>> torch.hub.download_url_to_file('https://s3.amazonaws.com/pytorch/models/resnet18-5c106cde.pth', '/tmp/temporary_file')
-
torch.hub.load_state_dict_from_url(url, model_dir=None, map_location=None, progress=True, check_hash=False, file_name=None)[source] -
Загружает сериализованный объект Torch по заданному URL.
Если загруженный файл является zip-архивом, он будет автоматически распакован.
Если объект уже присутствует в
model_dir, он будет десериализован и возвращён. Значение по умолчанию дляmodel_dirравно<hub_dir>/checkpoints, гдеhub_dir— каталог, возвращаемыйget_dir().- Параметры:
-
- url (str) – URL объекта для загрузки
- model_dir (str, необязательно) – каталог для сохранения объекта
- map_location (необязательно) – функция или словарь, определяющий, как перемапить места хранения (см. torch.load)
- progress (bool, необязательно) – отображать ли индикатор прогресса в stderr. По умолчанию: True
-
check_hash (bool, необязательно) – Если True, имя файла в URL должно соответствовать шаблону
filename-<sha256>.ext, где<sha256>— первые восемь или более цифр SHA256-хэша содержимого файла. Хэш используется для обеспечения уникальных имён и проверки содержимого файла. По умолчанию: False -
file_name (str, необязательно) – имя загруженного файла. Используется имя файла из
urlесли не указано.
- Тип возвращаемого значения:
Пример
>>> state_dict = torch.hub.load_state_dict_from_url('https://s3.amazonaws.com/pytorch/models/resnet18-5c106cde.pth')
Использование загруженной модели:
Обратите внимание, что *args и **kwargs в torch.hub.load() используются для создания экземпляра модели. После загрузки модели, как узнать, что можно сделать с моделью? Рекомендуемый порядок действий:
-
dir(model)чтобы увидеть все доступные методы модели. -
help(model.foo)чтобы узнать, какие аргументыmodel.fooпринимает для выполнения.
Для удобства пользователей в ходе изучения без постоянного обращения к документации, настоятельно рекомендуется владельцам репозиториев писать краткие и понятные сообщения справки для функций. Также полезно включать минимальный рабочий пример.
Где сохраняются загруженные модели?
Расположение используется в порядке
- Вызова
hub.set_dir(<PATH_TO_HUB_DIR>) -
$TORCH_HOME/hub, если переменная окруженияTORCH_HOMEустановлена. -
$XDG_CACHE_HOME/torch/hub, если переменная окруженияXDG_CACHE_HOMEустановлена. ~/.cache/torch/hub
-
torch.hub.get_dir()[source] -
Получить каталог кэша Torch Hub, используемый для хранения загруженных моделей и весов.
Если функция
set_dir()не вызывалась, по умолчанию используется путь$TORCH_HOME/hub, где переменная окружения$TORCH_HOMEпо умолчанию равна$XDG_CACHE_HOME/torch.$XDG_CACHE_HOMEследует спецификации X Design Group для структуры каталогов Linux-системы с значением по умолчанию~/.cacheв случае отсутствия переменной окружения.
-
torch.hub.set_dir(d)[source] -
Необязательно задать каталог Torch Hub для сохранения загруженных моделей и весов.
- Параметры:
-
d (str) – путь к локальной папке для сохранения загруженных моделей и весов.
Логика кэширования
По умолчанию файлы не удаляются после загрузки. Hub использует кэш по умолчанию, если он уже существует в каталоге, возвращаемом функцией get_dir().
Пользователи могут принудительно перезагрузить данные, вызвав hub.load(..., force_reload=True). Это приведет к удалению папки github и загруженных весов, а также к перезапуску новой загрузки. Это полезно, когда обновления публикуются в той же ветке, так пользователи смогут следить за последними версиями.
Известные ограничения:
Torch hub работает путем импорта пакета так, как будто он установлен. Это вносит некоторые побочные эффекты, связанные с импортом в Python. Например, вы можете увидеть новые элементы в кэшах Python sys.modules и sys.path_importer_cache, что является нормальным поведением Python. Это также означает, что могут возникнуть ошибки импорта при импорте различных моделей из разных репозиториев, если у репозиториев есть одинаковые имена подпакетов (обычно, подпакет model). Решением таких ошибок импорта является удаление проблемного подпакета из словаря sys.modules; подробности можно найти в данном сообщении на GitHub.
Стоит отметить известное ограничение: пользователи НЕ МОГУТ загрузить две разные ветки одного и того же репозитория в одном процессе Python. Это подобно установке двух пакетов с одинаковым именем в Python, что нежелательно. Кэш может внести изменения и привести к неожиданным результатам, если вы это сделаете. Конечно, загружать их в разных процессах вполне безопасно.
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/hub.html