Spec-Zone.ru › PyTorch 2.14

torch.func.stack_module_state

torch.func.stack_module_state(models) → params, buffers [исходный код]

Подготавливает список torch.nn.Modules для ансамблирования с помощью vmap().

Для списка M nn.Modules одного и того же класса возвращает два словаря, в которых объединены все их параметры и буферы, индексированные по имени. Объединённые параметры можно оптимизировать (то есть это новые листовые узлы в истории autograd, не связанные с исходными параметрами, которые можно напрямую передать оптимизатору).

Пример ансамблирования очень простой модели:

num_models = 5
batch_size = 64
in_features, out_features = 3, 3
models = [torch.nn.Linear(in_features, out_features) for i in range(num_models)]
data = torch.randn(batch_size, 3)


def wrapper(params, buffers, data):
    return torch.func.functional_call(models[0], (params, buffers), data)


params, buffers = stack_module_state(models)
output = vmap(wrapper, (0, 0, None))(params, buffers, data)

assert output.shape == (num_models, batch_size, out_features)

При наличии подмодулей используются соглашения об именовании state dict

import torch.nn as nn


class Foo(nn.Module):
    def __init__(self, in_features, out_features):
        super().__init__()
        hidden = 4
        self.l1 = nn.Linear(in_features, hidden)
        self.l2 = nn.Linear(hidden, out_features)

    def forward(self, x):
        return self.l2(self.l1(x))


num_models = 5
in_features, out_features = 3, 3
models = [Foo(in_features, out_features) for i in range(num_models)]
params, buffers = stack_module_state(models)
print(list(params.keys()))  # "l1.weight", "l1.bias", "l2.weight", "l2.bias"

Предупреждение

Все объединяемые модули должны быть одинаковыми (за исключением значений их параметров/буферов). Например, они должны находиться в одном режиме (training или eval).

Тип возвращаемого значения:

tuple[dict[str, Any], dict[str, Any]]

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.func.stack_module_state.html

Spec-Zone.ru

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