Класс MultiheadAttention
Класс MultiheadAttention из модуля nn реализует
механизм многоголового внимания. Он вычисляет внимание между запросами,
ключами и значениями, разделяя их на несколько голов для
одновременного изучения разных представлений.
Основные параметры конструктора:
-
embed_dim- размерность входных признаков -
num_heads- количество голов внимания -
batch_first- еслиTrue, размерность батча будет первой -
dropout- вероятность обнуления для регуляризации -
bias- использовать ли смещения в линейных слоях
Синтаксис
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)
Пример
Создадим слой многоголового внимания с размерностью 256 и 8 головами, а затем применим его к случайным входным данным:
import torch
torch.manual_seed(0)
embed_dim = 256
num_heads = 8
seq_len = 10
batch_size = 4
mha = torch.nn.MultiheadAttention(embed_dim, num_heads,
batch_first=True)
query = torch.randn(batch_size, seq_len, embed_dim)
key = torch.randn(batch_size, seq_len, embed_dim)
value = torch.randn(batch_size, seq_len, embed_dim)
res, attn_weights = mha(query, key, value)
print(res.shape)
Результат выполнения кода:
torch.Size([4, 10, 256])
Пример
Используем маску для запрета доступа к будущим позициям при обработке последовательности:
import torch
torch.manual_seed(0)
embed_dim = 64
num_heads = 4
seq_len = 5
batch_size = 2
mha = torch.nn.MultiheadAttention(embed_dim, num_heads,
batch_first=True)
query = torch.randn(batch_size, seq_len, embed_dim)
key = torch.randn(batch_size, seq_len, embed_dim)
value = torch.randn(batch_size, seq_len, embed_dim)
attn_mask = torch.triu(torch.ones(seq_len, seq_len),
diagonal=1).bool()
res, attn_weights = mha(query, key, value,
attn_mask=attn_mask)
print(attn_weights.shape)
Результат выполнения кода:
torch.Size([2, 4, 5, 5])
Пример
Применим слой внимания с добавлением смещений для ключей и значений и нулевого внимания:
import torch
torch.manual_seed(0)
embed_dim = 128
num_heads = 4
batch_size = 3
seq_len_q = 8
seq_len_kv = 6
mha = torch.nn.MultiheadAttention(embed_dim, num_heads,
add_bias_kv=True, add_zero_attn=True,
batch_first=True)
query = torch.randn(batch_size, seq_len_q, embed_dim)
key = torch.randn(batch_size, seq_len_kv, embed_dim)
value = torch.randn(batch_size, seq_len_kv, embed_dim)
res, attn_weights = mha(query, key, value)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 8, 128])
Смотрите также
-
класс
Transformer,
который использует многоголовое внимание внутри себя -
класс
TransformerEncoderLayer,
который включает слой внимания в свой состав -
класс
Linear,
который применяется для преобразования входных данных -
класс
Dropout,
используемый для регуляризации внутри механизма внимания