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.- Parameters
-
-
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 or None) –
"check",True,FalseилиNone. Этот параметр был добавлен в v1.12 и помогает гарантировать, что пользователи запускают код только из репозиториев, которым они доверяют.- Если
False, пользователь будет запрошен о доверии к репозиторию. - Если
True, репозиторий будет добавлен в список доверенных и загружен без явного подтверждения. - Если
"check", репозиторий будет проверен на предмет наличия в списке доверенных репозиториев в кеше. Если его нет в этом списке, поведение вернётся к вариантуtrust_repo=False. - Если
None: это выведет предупреждение, предложив пользователю установитьtrust_repoлибо наFalse,Trueили"check". Это сохранено только для обратной совместимости и будет удалено в версии 2.0.
По умолчанию
Noneи в конечном итоге изменится на"check"в версии 2.0. - Если
-
github (str) – строка в формате «repo_owner/repo_name[:ref]» с необязательным ref (метка или ветка). Если
- Returns
-
Доступные вызываемые точки входа
- Return type
Пример
>>> 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.- Parameters
-
-
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 or None) –
"check",True,FalseилиNone. Этот параметр был добавлен в v1.12 и помогает гарантировать, что пользователи запускают код только из репозиториев, которым они доверяют.- Если
False, пользователь будет запрошен о доверии к репозиторию. - Если
True, репозиторий будет добавлен в список доверенных и загружен без явного подтверждения. - Если
"check", репозиторий будет проверен на предмет наличия в списке доверенных репозиториев в кеше. Если его нет в этом списке, поведение вернётся к вариантуtrust_repo=False. - Если
None: это выведет предупреждение, предложив пользователю установитьtrust_repoлибо наFalse,Trueили"check". Это сохранено только для обратной совместимости и будет удалено в версии 2.0.
По умолчанию
Noneи в конечном итоге изменится на"check"в версии 2.0. - Если
-
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". Этот параметр присутствует только для обратной совместимости и будет удалён в версии v2.0.
По умолчанию равно
Noneи в конечном счете изменится на"check"в версии v2.0. - Если
-
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, скачанный файл должен начинаться с
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, weights_only=False)[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будет использовано, если не задано. -
weights_only (bool, необязательно) – Если True, загружаются только веса, а не сложные сериализованные объекты. Рекомендуется для недоверенных источников. См.
load()для более подробной информации.
- Тип возвращаемого значения
Пример
>>> 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 для сохранения загруженных моделей и весов.
- Parameters
-
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/2.1/hub.html