Класс MultiHeadAttention
Класс MultiHeadAttention представляет собой слой многоголового внимания. Он применяется для вычисления взвешенной суммы значений на основе запросов и ключей. Слой принимает на вход запросы, ключи и значения, а также может использовать маски для игнорирования определенных позиций. Основными параметрами являются: num_heads (количество голов внимания), key_dim (размерность пространства ключей и запросов), value_dim (размерность значений), dropout (доля отключаемых нейронов) и use_bias (использование смещения).
Синтаксис
tf.keras.layers.MultiHeadAttention(
num_heads,
key_dim,
value_dim=None,
dropout=0.0,
use_bias=True,
output_shape=None,
attention_axes=None,
kernel_initializer='glorot_uniform',
bias_initializer='zeros',
kernel_regularizer=None,
bias_regularizer=None,
activity_regularizer=None,
kernel_constraint=None,
bias_constraint=None,
**kwargs
)
Пример
Давайте создадим слой многоголового внимания и применим его к случайным тензорам запросов, ключей и значений:
import tensorflow as tf
tf.random.set_seed(0)
# Create a MultiHeadAttention layer
mha = tf.keras.layers.MultiHeadAttention(num_heads=2, key_dim=4)
# Create random query, key, value tensors
query = tf.random.normal((1, 5, 8))
key = tf.random.normal((1, 5, 8))
value = tf.random.normal((1, 5, 8))
# Apply attention
res = mha(query, value, key=key)
print(res.shape)
Результат выполнения кода:
(1, 5, 8)
Пример
Давайте используем слой многоголового внимания внутри функциональной модели Keras:
Результат выполнения кода (сокращенно):
"Model: \"model\""
"__________________________________________________________________________________________________"
" Layer (type) Output Shape Param # Connected to "
"=================================================================================================="
" input_1 (InputLayer) [(None, 5, 8)] 0 [] "
" "
" input_2 (InputLayer) [(None, 5, 8)] 0 [] "
" "
" multi_head_attention (Multi (None, 5, 8) 360 ['input_1[0][0]', "
" HeadAttention) 'input_2[0][0]'] "
" "
"=================================================================================================="
"Total params: 360"
"Trainable params: 360"
"Non-trainable params: 0"
"__________________________________________________________________________________________________"
Смотрите также
-
класс
Attention,
который реализует базовый механизм внимания -
класс
AdditiveAttention,
который реализует аддитивный механизм внимания -
класс
LSTM,
который представляет собой рекуррентный слой с долгой краткосрочной памятью -
класс
Dense,
который реализует полносвязный слой