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.
-
mod (
- Возвращает:
-
Замороженный
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