Spec-Zone.ru › PyTorch 1

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/1.13/generated/torch.jit.freeze.html

Spec-Zone.ru

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