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

Функция jit.save

Функция jit.save сохраняет скомпилированную модель PyTorch в файл на диске. Первым параметром функция принимает модель, созданную с помощью jit.script или jit.trace. Вторым параметром передаётся путь к файлу для сохранения. Функция сохраняет как архитектуру модели, так и её параметры, что позволяет загрузить модель без исходного кода.

Синтаксис

torch.jit.save(model, file_path)

Пример

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

import torch import torch.nn as nn class SimpleModel(nn.Module): def __init__(self): super(SimpleModel, self).__init__() self.fc = nn.Linear(10, 5) def forward(self, x): return self.fc(x) model = SimpleModel() scripted_model = torch.jit.script(model) torch.jit.save(scripted_model, 'model.pt') print("Model saved successfully")

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

"Model saved successfully"

Пример

Давайте сохраним модель, созданную с помощью jit.trace:

import torch import torch.nn as nn class TraceModel(nn.Module): def forward(self, x): return x * 2 + 1 model = TraceModel() example_input = torch.randn(1, 3) traced_model = torch.jit.trace(model, example_input) torch.jit.save(traced_model, 'traced_model.pt') print("Traced model saved")

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

"Traced model saved"

Пример

Давайте сохраним модель с использованием дополнительных параметров, например с включением дополнительной информации:

import torch import torch.nn as nn class ExtraModel(nn.Module): def __init__(self): super(ExtraModel, self).__init__() self.linear = nn.Linear(5, 2) def forward(self, x): return self.linear(x) model = ExtraModel() scripted_model = torch.jit.script(model) extra_files = {'metadata': 'model_info'} torch.jit.save(scripted_model, 'model_with_meta.pt', _extra_files=extra_files) print("Model with metadata saved")

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

"Model with metadata saved"

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

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