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