Spec-Zone.ru › PyTorch 2

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

Spec-Zone.ru

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