Spec-Zone.ru › PyTorch 2

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.

Returns

Доступные вызываемые точки входа

Return type

list

Пример

>>> 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.

Пример

>>> 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.
Возвращает

Результат вызова вызываемого объекта 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() для более подробной информации.
Тип возвращаемого значения

Dict[str, Any]

Пример

>>> 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API