РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
646 of 769 menu

Функция export.load

Функция export.load загружает программу, ранее экспортированную с помощью export.export или export.save. Она предназначена для восстановления модели в формате ExportedProgram, который включает в себя граф вычислений, параметры и буферы модели. Первым параметром функция принимает путь к файлу. Вторым параметром можно указать устройства (например, 'cpu' или 'cuda') для загрузки.

Синтаксис

torch.export.load(file_path, device=None)

Пример

Давайте сохраним модель с помощью export.save и затем загрузим её с помощью export.load:

import torch from torch.export import export, save, load class MyModule(torch.nn.Module): def forward(self, x): return x + 1 model = MyModule() example_inputs = (torch.randn(3, 4),) ep = export(model, example_inputs) save(ep, "model.pt") loaded_ep = load("model.pt") print(loaded_ep.module()(torch.ones(3, 4)))

Результат выполнения кода:

tensor([ [2., 2., 2., 2.], [2., 2., 2., 2.], [2., 2., 2., 2.], ])

Пример

Загрузка экспортированной программы на устройство CUDA (если доступно):

import torch from torch.export import export, save, load class MyModule(torch.nn.Module): def forward(self, x): return x * 2 model = MyModule() example_inputs = (torch.randn(2, 3),) ep = export(model, example_inputs) save(ep, "model_cuda.pt") device = "cuda" if torch.cuda.is_available() else "cpu" loaded_ep = load("model_cuda.pt", device=device) print(loaded_ep.module()(torch.ones(2, 3, device=device)))

Результат выполнения кода (на CPU):

tensor([ [2., 2., 2.], [2., 2., 2.], ])

Пример

Загрузка модели с сохранёнными параметрами и применение к новым данным:

import torch from torch.export import export, save, load class LinearModule(torch.nn.Module): def __init__(self): super().__init__() self.linear = torch.nn.Linear(5, 3) def forward(self, x): return self.linear(x) model = LinearModule() example_inputs = (torch.randn(2, 5),) ep = export(model, example_inputs) save(ep, "linear_model.pt") loaded_ep = load("linear_model.pt") t = torch.randn(2, 5) res = loaded_ep.module()(t) print(res)

Результат выполнения кода:

tensor([ [ 0.2679, -0.6827, 0.4435], [ 0.0570, -0.7008, -0.1617], ], grad_fn=<AddmmBackward0>)

Смотрите также

  • функцию export.save,
    которая сохраняет экспортированную программу в файл
  • функцию export.export,
    которая создаёт экспортированную программу из модели
  • функцию load,
    которая загружает сохранённое состояние модели (state_dict)
  • функцию save,
    которая сохраняет объект модели в файл
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить