Трансформер
-
class torch.nn.Transformer(d_model=512, nhead=8, num_encoder_layers=6, num_decoder_layers=6, dim_feedforward=2048, dropout=0.1, activation=<function relu>, custom_encoder=None, custom_decoder=None, layer_norm_eps=1e-05, batch_first=False, norm_first=False, bias=True, device=None, dtype=None)[source] -
Модель трансформера. Пользователь может изменять атрибуты по мере необходимости. Архитектура основана на статье «Attention Is All You Need». Ashish Vaswani, Noam Shazeer, Niki Parmar, Jakob Uszkoreit, Llion Jones, Aidan N Gomez, Lukasz Kaiser и Illia Polosukhin. 2017. Attention is all you need. В Advances in Neural Information Processing Systems, страницы 6000-6010.
- Параметры
-
- d_model (int) – количество ожидаемых признаков в входных данных кодировщика/декодировщика (по умолчанию=512).
- nhead (int) – количество голов в моделях многоголового внимания (по умолчанию=8).
- num_encoder_layers (int) – количество подслоёв кодировщика в кодировщике (по умолчанию=6).
- num_decoder_layers (int) – количество подслоёв декодировщика в декодировщике (по умолчанию=6).
- dim_feedforward (int) – размер модели нейронной сети прямого отображения (по умолчанию=2048).
- dropout (float) – значение дропаута (по умолчанию=0.1).
- activation (Union[str, Callable[[Tensor], Tensor]]) – функция активации промежуточного слоя кодировщика/декодировщика, может быть строкой (“relu” или “gelu”) или унарной функцией. По умолчанию: relu
- custom_encoder (Optional[Any]) – настраиваемый кодировщик (по умолчанию=None).
- custom_decoder (Optional[Any]) – настраиваемый декодировщик (по умолчанию=None).
- layer_norm_eps (float) – значение eps в компонентах нормализации слоёв (по умолчанию=1e-5).
-
batch_first (bool) – Если
True, то входные и выходные тензоры предоставляются как (batch, seq, feature). По умолчанию:False(seq, batch, feature). -
norm_first (bool) – если
True, слои кодировщика и декодировщика будут выполнять LayerNorms перед другими операциями внимания и прямого отображения, иначе после. По умолчанию:False(после). -
bias (bool) – Если установлено в
False, слоиLinearиLayerNormне будут учиться аддитивному смещению. По умолчанию:True.
- Примеры::
-
>>> transformer_model = nn.Transformer(nhead=16, num_encoder_layers=12) >>> src = torch.rand((10, 32, 512)) >>> tgt = torch.rand((20, 32, 512)) >>> out = transformer_model(src, tgt)
Примечание: Полный пример применения модуля nn.Transformer для модели языка слов доступен по адресу https://github.com/pytorch/examples/tree/master/word_language_model
-
forward(src, tgt, src_mask=None, tgt_mask=None, memory_mask=None, src_key_padding_mask=None, tgt_key_padding_mask=None, memory_key_padding_mask=None, src_is_causal=None, tgt_is_causal=None, memory_is_causal=False)[source] -
Принимает и обрабатывает замаскированные последовательности источника/цели.
- Параметры
-
- src (Tensor) – последовательность для кодировщика (требуется).
- tgt (Tensor) – последовательность для декодировщика (требуется).
- src_mask (Optional[Tensor]) – аддитивная маска для последовательности src (необязательно).
- tgt_mask (Optional[Tensor]) – аддитивная маска для последовательности tgt (необязательно).
- memory_mask (Optional[Tensor]) – аддитивная маска для выходных данных кодировщика (необязательно).
- src_key_padding_mask (Optional[Tensor]) – маска тензора для ключей src по каждой группе (необязательно).
- tgt_key_padding_mask (Optional[Tensor]) – маска тензора для ключей tgt по каждой группе (необязательно).
- memory_key_padding_mask (Optional[Tensor]) – маска тензора для ключей памяти по каждой группе (необязательно).
-
src_is_causal (Optional[bool]) – Если указано, применяет маску причинности как
src_mask. По умолчанию:None; пытается обнаружить маску причинности. Предупреждение:src_is_causalдаёт подсказку, чтоsrc_mask— это маска причинности. Неправильные подсказки могут привести к некорректной работе, включая совместимость вперёд и назад. -
tgt_is_causal (Optional[bool]) – Если указано, применяет маску причинности как
tgt_mask. По умолчанию:None; пытается обнаружить маску причинности. Предупреждение:tgt_is_causalдаёт подсказку, чтоtgt_mask— это маска причинности. Неправильные подсказки могут привести к некорректной работе, включая совместимость вперёд и назад. -
memory_is_causal (bool) – Если указано, применяет маску причинности как
memory_mask. По умолчанию:False. Предупреждение:memory_is_causalдаёт подсказку, чтоmemory_mask— это маска причинности. Неправильные подсказки могут привести к некорректной работе, включая совместимость вперёд и назад.
- Тип возвращаемого значения
- Форма:
-
- src: для неразбитого на пакеты ввода, если
batch_first=Falseили(N, S, E)еслиbatch_first=True. - tgt: для неразбитого на пакеты ввода, если
batch_first=Falseили(N, T, E)еслиbatch_first=True. - src_mask: или .
- tgt_mask: или .
- memory_mask: .
- src_key_padding_mask: для неразбитого на пакеты ввода в противном случае .
- tgt_key_padding_mask: для неразбитого на пакеты ввода в противном случае .
- memory_key_padding_mask: для неразбитого на пакеты ввода в противном случае .
Примечание: [src/tgt/memory]_mask гарантирует, что позиция i разрешена для внимания к не замаскированным позициям. Если предоставлен BoolTensor, позиции с
Trueне разрешено обращать внимание, в то время как значенияFalseостанутся неизменными. Если предоставлен FloatTensor, он будет добавлен к весу внимания. [src/tgt/memory]_key_padding_mask предоставляет указанные элементы в ключе, которые нужно игнорировать вниманием. Если предоставлен BoolTensor, позиции со значениемTrueбудут игнорироваться, а позиции со значениемFalseостанутся неизменными.- вывод: для неразбитого на пакеты ввода, если
batch_first=Falseили(N, T, E)еслиbatch_first=True.
Примечание: Из-за архитектуры многоголового внимания в модели преобразователя длина выходной последовательности преобразователя такая же, как длина входной последовательности (т. е. целевой) декодера.
где S - длина исходной последовательности, T - длина целевой последовательности, N - размер пакета, E - число признаков
- src: для неразбитого на пакеты ввода, если
Примеры
>>> output = transformer_model(src, tgt, src_mask=src_mask, tgt_mask=tgt_mask)
-
static generate_square_subsequent_mask(sz, device=device(type='cpu'), dtype=torch.float32)[source] -
Создайте квадратную причинную маску для последовательности. Замаскированные позиции заполнены float(‘-inf’). Незамаскированные позиции заполнены float(0.0).
- Возвращает
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.Transformer.html