MultiheadAttention
-
class torch.nn.MultiheadAttention(embed_dim, num_heads, dropout=0.0, bias=True, add_bias_kv=False, add_zero_attn=False, kdim=None, vdim=None, batch_first=False, device=None, dtype=None)[исходный код] -
Позволяет модели одновременно учитывать информацию из разных подпространств представлений.
Этот слой MultiheadAttention реализует исходную архитектуру, описанную в статье Attention Is All You Need. Этот слой предназначен для использования в качестве эталонной реализации, помогающей понять основы, поэтому по сравнению с более новыми архитектурами он включает лишь ограниченный набор функций. Учитывая стремительное развитие архитектур на основе трансформеров, рекомендуем изучить это руководство, чтобы создавать эффективные слои из базовых блоков ядра, или использовать библиотеки более высокого уровня из экосистемы PyTorch.
Механизм Multi-Head Attention определяется следующим образом:
где .
nn.MultiheadAttentionпри возможности использует оптимизированные реализацииscaled_dot_product_attention().Помимо поддержки новой функции
scaled_dot_product_attention(), для ускорения вывода MHA использует быстрый путь выполнения при выводе с поддержкой вложенных тензоров, если:- вычисляется внимание к самому себе (т. е.
query,keyиvalue— один и тот же тензор); - входные данные пакетированы (3D) с
batch_first==True - Автоматическое дифференцирование отключено (с помощью
torch.inference_modeилиtorch.no_grad) либо ни один аргумент-тензорrequires_grad - обучение отключено (с помощью
.eval()) -
add_bias_kv—False -
add_zero_attn—False -
kdimиvdimравныembed_dim - если передан вложенный тензор, то не переданы ни
key_padding_mask, ниattn_mask - автокастирование отключено
Если используется оптимизированная реализация быстрого пути для вывода, для
query/key/valueможно передать вложенный тензор, который представляет заполнение более эффективно, чем маска заполнения. В этом случае будет возвращен вложенный тензор, а ускорение будет пропорционально доле заполнения во входных данных.- Параметры:
-
- embed_dim – Общая размерность модели.
-
num_heads – Количество параллельных голов внимания. Обратите внимание, что
embed_dimбудет разделен междуnum_heads(т. е. размерность каждой головы составитembed_dim // num_heads). -
dropout – Вероятность исключения для
attn_output_weights. По умолчанию:0.0(без исключения). -
bias – Если задано, добавляет смещение к слоям проекции входных и выходных данных. По умолчанию:
True. -
add_bias_kv – Если задано, добавляет смещение к последовательностям ключей и значений по измерению dim=0. По умолчанию:
False. -
add_zero_attn – Если задано, добавляет новую порцию нулей к последовательностям ключей и значений по измерению dim=1. По умолчанию:
False. -
kdim – Общее количество признаков ключей. По умолчанию:
None(используетkdim=embed_dim). -
vdim – Общее количество признаков значений. По умолчанию:
None(используетvdim=embed_dim). -
batch_first – Если
True, входные и выходные тензоры задаются в формате (batch, seq, feature). По умолчанию:False(seq, batch, feature).
Примеры:
>>> multihead_attn = nn.MultiheadAttention(embed_dim, num_heads) >>> attn_output, attn_output_weights = multihead_attn(query, key, value)
-
forward(query, key, value, key_padding_mask=None, need_weights=True, attn_mask=None, average_attn_weights=True, is_causal=False)[исходный код] -
Вычисляет результаты внимания, используя эмбеддинги запроса, ключа и значения.
Поддерживает необязательные параметры заполнения, масок и весов внимания.
- Параметры:
-
-
query (Tensor) – Эмбеддинги запроса формы для непакетированного входа, при
batch_first=Falseили приbatch_first=True, где — длина целевой последовательности, — размер пакета, а — размерность эмбеддинга запросаembed_dim. Запросы сопоставляются с парами ключ-значение для получения результата. Подробнее см. в статье «Attention Is All You Need». -
key (Tensor) – Эмбеддинги ключа формы для непакетированного входа, при
batch_first=Falseили приbatch_first=True, где — длина исходной последовательности, — размер пакета, а — размерность эмбеддинга ключаkdim. Подробнее см. в статье «Attention Is All You Need». -
value (Tensor) – Эмбеддинги значения формы для непакетированного входа, при
batch_first=Falseили приbatch_first=True, где — длина исходной последовательности, — размер пакета, а — размерность эмбеддинга значенияvdim. Подробнее см. в статье «Attention Is All You Need». -
key_padding_mask (Tensor | None) – Если задана, маска формы , указывающая, какие элементы в
keyследует игнорировать при вычислении внимания (т. е. считать «заполнением»). Для непакетированныхqueryформа должна быть . Поддерживаются двоичные маски и маски с плавающей точкой. Для двоичной маски значениеTrueозначает, что соответствующее значениеkeyбудет игнорироваться при вычислении внимания. Для маски с плавающей точкой ее значение напрямую прибавляется к соответствующему значениюkey. -
need_weights (bool) – Если задано, возвращает
attn_output_weightsвместе сattn_outputs. Установитеneed_weights=False, чтобы использовать оптимизированнуюscaled_dot_product_attentionи добиться максимальной производительности MHA. По умолчанию:True. -
attn_mask (Tensor | None) – Если задана, двумерная или трехмерная маска, запрещающая обращать внимание на определенные позиции. Должна иметь форму или , где — размер пакета, — длина целевой последовательности, а — длина исходной последовательности. Двумерная маска будет распространена на весь пакет, а трехмерная маска позволяет задавать отдельную маску для каждого элемента пакета. Поддерживаются двоичные маски и маски с плавающей точкой. Для двоичной маски значение
Trueозначает, что внимание к соответствующей позиции запрещено. Для маски с плавающей точкой ее значения прибавляются к весам внимания. Если заданы одновременно attn_mask и key_padding_mask, их типы должны совпадать. -
average_attn_weights (bool) – Если значение истинно, возвращаемые
attn_weightsусредняются по головам. В противном случаеattn_weightsвозвращаются отдельно для каждой головы. Этот флаг действует только при условииneed_weights=True. По умолчанию:True(т. е. веса усредняются по головам). -
is_causal (bool) – Если задано, в качестве маски внимания применяется причинно-следственная маска. По умолчанию:
False. Предупреждение:is_causalслужит подсказкой о том, чтоattn_maskявляется причинно-следственной маской. Неверные подсказки могут привести к ошибочному выполнению, в том числе нарушить совместимость прямого и обратного проходов.
-
query (Tensor) – Эмбеддинги запроса формы для непакетированного входа, при
- Тип возвращаемого значения:
- Выходные данные:
-
-
attn_output — Результаты внимания формы , если вход не пакетирован, при
batch_first=Falseили приbatch_first=True, где — длина целевой последовательности, — размер пакета, а — размерность эмбеддингаembed_dim. -
attn_output_weights — Возвращается только при
need_weights=True. Еслиaverage_attn_weights=True, возвращаются веса внимания, усредненные по головам, формы для непакетированного входа или , где — размер пакета, — длина целевой последовательности, а — длина исходной последовательности. Еслиaverage_attn_weights=False, возвращаются веса внимания для каждой головы формы для непакетированного входа или .
Примечание
Аргумент
batch_firstигнорируется для непакетированных входных данных. -
attn_output — Результаты внимания формы , если вход не пакетирован, при
-
merge_masks(attn_mask, key_padding_mask, query)[исходный код] -
Определяет тип маски и при необходимости объединяет маски.
Если задана только одна маска, возвращаются эта маска и соответствующий ей тип. Если заданы обе маски, они расширяются до формы
(batch_size, num_heads, seq_len, seq_len), объединяются с помощью логической операцииor, а типу маски присваивается значение 2 :param attn_mask: маска внимания формы(seq_len, seq_len), тип маски 0 :param key_padding_mask: маска заполнения формы(batch_size, seq_len), тип маски 1 :param query: эмбеддинги запроса формы(batch_size, seq_len, embed_dim)- Возвращает:
-
объединенная маска mask_type: тип объединенной маски (0, 1 или 2)
- Тип возвращаемого значения:
-
merged_mask
- вычисляется внимание к самому себе (т. е.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.MultiheadAttention.html