Spec-Zone.ru › PyTorch 1

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.

Возвращаемое значение:

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

Тип возвращаемого значения:

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.

Параметры:
  • 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.

Пример

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

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

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 для сохранения загруженных моделей и весов.

Параметры:

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

Spec-Zone.ru

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