torch.func.stack_module_state
-
torch.func.stack_module_state(models) → params, buffers -
Подготавливает список torch.nn.Modules для объединения с помощью
vmap().Принимая на вход список
Mnn.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(model[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"Предупреждение
Все модули, которые объединяются, должны быть одинаковыми (кроме значений их параметров/буферов). Например, они должны быть в одном режиме (обучения или оценки).
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.stack_module_state.html