Spec-Zone.ru › PyTorch 1

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

Spec-Zone.ru

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