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

Функция jit.freeze

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

Синтаксис

torch.jit.freeze(module)

Пример

Давайте заморозим простую модель линейной регрессии:

import torch import torch.nn as nn class LinearModel(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(10, 5) def forward(self, x): return self.fc(x) model = LinearModel() scripted_model = torch.jit.script(model) frozen_model = torch.jit.freeze(scripted_model) print(type(frozen_model))

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

<class 'torch.jit._freeze.FrozenScriptModule'>

Пример

Давайте сравним скорость работы обычной модели и замороженной:

import torch import torch.nn as nn import time torch.manual_seed(0) class SimpleNet(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(100, 50) self.fc2 = nn.Linear(50, 10) def forward(self, x): x = torch.relu(self.fc1(x)) return self.fc2(x) model = SimpleNet() scripted_model = torch.jit.script(model) frozen_model = torch.jit.freeze(scripted_model) x = torch.randn(100, 100) start = time.time() res = model(x) print("Обычная модель:", time.time() - start) start = time.time() res = frozen_model(x) print("Замороженная модель:", time.time() - start)

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

Обычная модель: 0.0012 Замороженная модель: 0.0008

Пример

Давайте используем замороженную модель для инференса:

import torch import torch.nn as nn torch.manual_seed(0) model = nn.Linear(5, 3) scripted_model = torch.jit.script(model) frozen_model = torch.jit.freeze(scripted_model) x = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0]) res = frozen_model(x) print(res)

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

tensor([ 0.3431, -0.3410, -0.2866], grad_fn=<AddBackward0>)

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

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