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/2.1/generated/torch.nn.utils.parametrize.cached.html