Spec-Zone.ru › PyTorch 1

Растягивание

class torch.nn.Unflatten(dim, unflattened_size) [source]

Растягивает тензор по указанному измерению, расширяя его до желаемой формы. Для использования с Sequential.

  • dim указывает измерение входного тензора, которое нужно растянуть, и может быть либо int , либо str при использовании Tensor или NamedTensor соответственно.
  • unflattened_size — новая форма растянутого измерения тензора и может быть tuple целых чисел, или list целых чисел, или torch.Size для Tensor входных данных; NamedShape (кортеж (name, size) кортежей) для NamedTensor входных данных.
Форма:
  • Вход: (∗,Sdim,∗)(*, S_{\text{dim}}, *), где SdimS_{\text{dim}} — размер по размерности dim, а ∗* обозначает любое количество измерений, включая нулевое.
  • Выход: (∗,U1,...,Un,∗)(*, U_1, ..., U_n, *), где UU = unflattened_size и ∏i=1nUi=Sdim\prod_{i=1}^n U_i = S_{\text{dim}}.
Параметры:
  • dim (Union[int, str]) – Измерение для растягивания
  • unflattened_size (Union[torch.Size, Tuple, List, NamedShape]) – Новая форма растянутого измерения

Примеры

>>> input = torch.randn(2, 50)
>>> # With tuple of ints
>>> m = nn.Sequential(
>>>     nn.Linear(50, 50),
>>>     nn.Unflatten(1, (2, 5, 5))
>>> )
>>> output = m(input)
>>> output.size()
torch.Size([2, 2, 5, 5])
>>> # With torch.Size
>>> m = nn.Sequential(
>>>     nn.Linear(50, 50),
>>>     nn.Unflatten(1, torch.Size([2, 5, 5]))
>>> )
>>> output = m(input)
>>> output.size()
torch.Size([2, 2, 5, 5])
>>> # With namedshape (tuple of tuples)
>>> input = torch.randn(2, 50, names=('N', 'features'))
>>> unflatten = nn.Unflatten('features', (('C', 2), ('H', 5), ('W', 5)))
>>> output = unflatten(input)
>>> output.size()
torch.Size([2, 2, 5, 5])

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.Unflatten.html

Spec-Zone.ru

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