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

Класс RecursiveScriptModule

Класс RecursiveScriptModule представляет собой скомпилированную версию модуля PyTorch, полученную с помощью TorchScript. Экземпляры этого класса создаются автоматически при использовании функций torch.jit.script или torch.jit.trace. Такой модуль можно сохранять на диск, загружать обратно и выполнять без доступа к исходному Python-коду. Основные методы класса: save для сохранения на диск, load для загрузки, а также прямой вызов объекта для выполнения прямого прохода.

Синтаксис

# Создание RecursiveScriptModule из существующего модуля script_module = torch.jit.script(model) # Сохранение на диск script_module.save("model.pt") # Загрузка с диска loaded = torch.jit.load("model.pt") # Выполнение прямого прохода output = script_module(input_tensor)

Пример

Давайте создадим простую модель, скомпилируем её и выполним прямой проход:

import torch # Создаём простую линейную модель model = torch.nn.Linear(5, 3) model.eval() # Компилируем в TorchScript script_module = torch.jit.script(model) # Создаём входной тензор x = torch.randn(2, 5) # Выполняем прямой проход output = script_module(x) print(output)

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

tensor([ [-0.0895, -0.5332, -0.8045], [-0.5116, -0.6450, -0.0992], ], grad_fn=<AddmmBackward0>)

Пример

Рассмотрим сохранение скомпилированной модели на диск и её последующую загрузку:

import torch # Создаём и компилируем модель model = torch.nn.Sequential( torch.nn.Linear(10, 20), torch.nn.ReLU(), torch.nn.Linear(20, 5) ) model.eval() script_module = torch.jit.script(model) # Сохраняем на диск script_module.save("model.pt") # Загружаем с диска loaded_module = torch.jit.load("model.pt") # Проверяем работу загруженной модели x = torch.randn(1, 10) output = loaded_module(x) print(output.shape)

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

torch.Size([1, 5])

Пример

RecursiveScriptModule поддерживает доступ к вложенным подмодулям и их параметрам через точечную нотацию:

import torch class MyModule(torch.nn.Module): def __init__(self): super().__init__() self.fc1 = torch.nn.Linear(10, 20) self.fc2 = torch.nn.Linear(20, 5) def forward(self, x): x = torch.relu(self.fc1(x)) return self.fc2(x) model = MyModule() model.eval() script_module = torch.jit.script(model) # Доступ к параметрам подмодулей weight = script_module.fc1.weight bias = script_module.fc2.bias print(weight.shape) print(bias.shape)

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

torch.Size([20, 10]) torch.Size([5])

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

  • функцию torch.jit.script,
    которая компилирует модуль или функцию в TorchScript
  • функцию torch.jit.trace,
    которая компилирует модуль на основе примера входных данных
  • функцию torch.save,
    которая сохраняет обычный объект PyTorch в файл
  • функцию torch.jit.load,
    которая загружает скомпилированный модуль с диска
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить