torch.nn.utils.parametrize.cached
-
torch.nn.utils.parametrize.cached()[source] -
Менеджер контекста, который активирует систему кэширования внутри параметризаций, зарегистрированных с помощью
register_parametrization().Значение параметризованных объектов вычисляется и кэшируется в первый раз, когда они требуются, когда активен этот менеджер контекста. Кэшированные значения удаляются при выходе из менеджера контекста.
Это полезно, когда параметризованный параметр используется более одного раза в прямом проходе. Примером этого является параметризация рекуррентного ядра RNN или при совместном использовании весов.
Самый простой способ активировать кэш — обернуть прямой проход нейронной сети
import torch.nn.utils.parametrize as P ... with P.cached(): output = model(inputs)при обучении и оценке. Можно также обернуть части модулей, которые многократно используют параметризованные тензоры. Например, цикл RNN с параметризованным рекуррентным ядром:
with P.cached(): for x in xs: out_rnn = self.rnn_cell(x, out_rnn)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.utils.parametrize.cached.html