Многоголовое внимание
-
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)[source] -
Позволяет модели совместно обращать внимание на информацию из разных подпространств представлений, как описано в статье: Attention Is All You Need.
Многоголовое внимание определяется как:
где .
forward()будет использовать специальную оптимизированную реализацию, если выполнены все следующие условия:- вычисляется самовнимание (т.е.,
query,key, иvalueявляются одним и тем же тензором. Это ограничение будет ослаблено в будущем.) - либо вычисление градиента отключено (используя
torch.inference_modeилиtorch.no_grad) или ни один из тензорных аргументовrequires_grad - тренировка отключена (используя
.eval()) - значение дропаута равно 0
-
add_bias_kvравноFalse -
add_zero_attnравноFalse -
batch_firstравноTrueи входной тензор является пакетированным -
kdimиvdimравныembed_dim - передаётся не более одного из
key_padding_maskилиattn_mask - если передаётся NestedTensor, ни
key_padding_mask, ниattn_maskне передаются
Если используется оптимизированная реализация, для
query/key/valueможно передать NestedTensor, чтобы более эффективно представить заполнение, чем с помощью маски заполнения. В этом случае будет возвращён NestedTensor, и ожидается дополнительное ускорение, пропорциональное доли входных данных, являющейся заполнением.- Параметры:
-
- embed_dim – Общая размерность модели.
-
num_heads – Количество параллельных головок внимания. Обратите внимание, что
embed_dimбудет разделено поnum_heads(т.е. каждая головка будет иметь размерembed_dim // num_heads). -
dropout – Вероятность дропаута для
attn_output_weights. По умолчанию:0.0(дропаут отсутствует). -
bias – Если указано, добавляет смещение к слоям проекции входных/выходных данных. По умолчанию:
True. -
add_bias_kv – Если указано, добавляет смещение к последовательностям ключей и значений в измерении 0. По умолчанию:
False. -
add_zero_attn – Если указано, добавляет новый пакет нулей к последовательностям ключей и значений в измерении 1. По умолчанию:
False. -
kdim – Общее количество признаков для ключей. По умолчанию:
None(используетkdim=embed_dim). -
vdim – Общее количество признаков для значений. По умолчанию:
None(используетvdim=embed_dim). -
batch_first – Если
True, то входные и выходные тензоры предоставляются как (пакет, последовательность, признак). По умолчанию:False(последовательность, пакет, признак).
Примеры:
>>> 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)[source] -
- Параметры:
-
-
query (Тензор) – Векторы запросов формы для неразмеченного входа, при
batch_first=Falseили приbatch_first=True, где — длина целевой последовательности, — размер пакета, а — размерность вектора запросаembed_dim. Запросы сравниваются с парами «ключ-значение» для получения выходных данных. Для получения дополнительной информации см. «Attention Is All You Need». -
key (Тензор) – Векторы ключей формы для неразмеченного входа, при
batch_first=Falseили приbatch_first=True, где — длина источника последовательности, — размер пакета, а — размерность вектора ключаkdim. Для получения дополнительной информации см. «Attention Is All You Need». -
value (Тензор) – Векторы значений формы для неразмеченного входа, при
batch_first=Falseили приbatch_first=True, где — длина источника последовательности, — размер пакета, а — размерность вектора значенийvdim. Для получения дополнительной информации см. «Attention Is All You Need». -
key_padding_mask (Необязательный[Тензор]) – Если указано, маска формы , указывающая, какие элементы внутри
keyследует игнорировать для целей внимания (т.е. рассматривать как «заполнение»). Для неразмеченногоquery, форма должна быть . Поддерживаются двоичные и байтовые маски. Для двоичной маски значениеTrueуказывает, что соответствующее значениеkeyбудет проигнорировано для целей внимания. Для плавающей маски она будет напрямую добавлена к соответствующему значениюkey. -
need_weights (логическое) – Если указано, возвращает
attn_output_weightsпомимоattn_outputs. По умолчанию:True. -
attn_mask (Необязательный[Тензор]) – Если указано, 2D или 3D маска, предотвращающая внимание к определённым позициям. Должна иметь форму или , где — размер пакета, — длина целевой последовательности, а — длина последовательности источника. 2D маска будет транслироваться по пакету, а 3D маска позволит иметь различную маску для каждого элемента пакета. Поддерживаются двоичные, байтовые и плавающие маски. Для двоичной маски значение
Trueуказывает, что соответствующая позиция не может участвовать во внимании. Для байтовой маски ненулевое значение указывает, что соответствующая позиция не может участвовать во внимании. Для плавающей маски значения маски будут добавлены к весу внимания.
-
query (Тензор) – Векторы запросов формы для неразмеченного входа, при
-
-
average_attn_weights (bool) – Если значение равно true, указывает, что возвращаемые
attn_weightsдолжны быть усреднены по всем головкам. В противном случаеattn_weightsпредоставляются отдельно для каждой головки. Обратите внимание, что этот флаг действует только приneed_weights=True. Значение по умолчанию:True(т.е. усреднение весов по всем головкам)
-
average_attn_weights (bool) – Если значение равно true, указывает, что возвращаемые
- Тип возвращаемого значения:
- Выходные данные:
-
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 - Выходные данные внимания, имеющие форму , когда вход не является пакетным, когда
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.MultiheadAttention.html