Spec-Zone.ru › PyTorch 2

torch.jit.freeze

torch.jit.freeze(mod, preserved_attrs=None, optimize_numerics=True) [source]

Замораживание ScriptModule позволит создать его копию и попытаться встроить подмодули, параметры и атрибуты скопированного модуля в качестве констант в график TorchScript IR. По умолчанию будут сохранены forward, а также атрибуты и методы, указанные в preserved_attrs. Кроме того, любые атрибуты, изменённые в сохранённом методе, также будут сохранены.

Замораживание в настоящее время поддерживает только ScriptModules, находящиеся в режиме eval.

Замораживание применяет общие оптимизации, которые ускорят вашу модель независимо от машины. Для дальнейшей оптимизации с использованием параметров, специфичных для сервера, выполните optimize_for_inference после замораживания.

Параметры
  • mod (ScriptModule) – модуль, подлежащий замораживанию
  • preserved_attrs (Необязательный[Список[str]]) – список атрибутов, которые необходимо сохранить помимо метода forward. Атрибуты, изменённые в сохранённых методах, также будут сохранены.
  • optimize_numerics (bool) – Если True, будет выполнена серия оптимизационных проходов, которые не сохраняют численность точно. Полные сведения об оптимизации можно найти по ссылке torch.jit.run_frozen_optimizations.
Возвращает

Замороженный ScriptModule.

Пример (замораживание простого модуля с параметром):

    def forward(self, input):
        output = self.weight.mm(input)
        output = self.linear(output)
        return output

scripted_module = torch.jit.script(MyModule(2, 3).eval())
frozen_module = torch.jit.freeze(scripted_module)
# parameters have been removed and inlined into the Graph as constants
assert len(list(frozen_module.named_parameters())) == 0
# See the compiled graph as Python code
print(frozen_module.code)

Пример (замораживание модуля с сохранёнными атрибутами)

    def forward(self, input):
        self.modified_tensor += 1
        return input + self.modified_tensor

scripted_module = torch.jit.script(MyModule2().eval())
frozen_module = torch.jit.freeze(scripted_module, preserved_attrs=["version"])
# we've manually preserved `version`, so it still exists on the frozen module and can be modified
assert frozen_module.version == 1
frozen_module.version = 2
# `modified_tensor` is detected as being mutated in the forward, so freezing preserves
# it to retain model semantics
assert frozen_module(torch.tensor(1)) == torch.tensor(12)
# now that we've run it once, the next result will be incremented by one
assert frozen_module(torch.tensor(1)) == torch.tensor(13)

Примечание

Поддерживается также замораживание атрибутов подмодулей: frozen_module = torch.jit.freeze(scripted_module, preserved_attrs=[“submodule.version”])

Примечание

Если вы не уверены, почему атрибут не встраивается в качестве константы, вы можете запустить dump_alias_db на frozen_module.forward.graph, чтобы убедиться, что замораживание обнаружило, что атрибут изменяется.

Примечание

Поскольку замораживание делает веса постоянными и удаляет иерархию модулей, to и другие методы nn.Module для изменения устройства или типа данных больше не работают. В качестве обходного решения можно переназначить устройства, указав map_location в torch.jit.load, однако устройство-специфическая логика может быть встроена в модель.

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.jit.freeze.html

Spec-Zone.ru

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