torch.nn.utils.rnn.pad_packed_sequence
-
torch.nn.utils.rnn.pad_packed_sequence(sequence, batch_first=False, padding_value=0.0, total_length=None)[source] -
Заполняет упакованную партию последовательностей переменной длины.
Это обратная операция к
pack_padded_sequence().Данные возвращаемого тензора будут иметь размер
T x B x *, гдеT— длина самой длинной последовательности, аB— размер пакет. Еслиbatch_firstравно True, данные будут транспонированы в форматB x T x *.Пример
>>> from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence >>> seq = torch.tensor([[1,2,0], [3,0,0], [4,5,6]]) >>> lens = [2, 1, 3] >>> packed = pack_padded_sequence(seq, lens, batch_first=True, enforce_sorted=False) >>> packed PackedSequence(data=tensor([4, 1, 3, 5, 2, 6]), batch_sizes=tensor([3, 2, 1]), sorted_indices=tensor([2, 0, 1]), unsorted_indices=tensor([1, 2, 0])) >>> seq_unpacked, lens_unpacked = pad_packed_sequence(packed, batch_first=True) >>> seq_unpacked tensor([[1, 2, 0], [3, 0, 0], [4, 5, 6]]) >>> lens_unpacked tensor([2, 1, 3])Примечание
total_lengthполезно для реализации шаблонаpack sequence -> recurrent network -> unpack sequenceвModule, обернутого вDataParallel. Подробности см. в этом разделе FAQ.- Параметры:
-
- sequence (PackedSequence) – пакет для заполнения
-
batch_first (bool, необязательно) – если
True, вывод будет в форматеB x T x *. - padding_value (float, необязательно) – значения для заполненных элементов.
-
total_length (int, необязательно) – если не
None, вывод будет заполнен до длиныtotal_length. Этот метод выброситValueError, еслиtotal_lengthменьше максимальной длины последовательности вsequence.
- Возвращаемое значение:
-
Кортеж тензоров, содержащий заполненную последовательность и тензор, содержащий список длин каждой последовательности в пакете. Элементы пакета будут переупорядочены так, как они были упорядочены изначально, когда пакет был передан в
pack_padded_sequenceилиpack_sequence. - Тип возвращаемого значения:
© 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.rnn.pad_packed_sequence.html