Класс AdditiveAttention
Класс AdditiveAttention реализует механизм внимания,
предложенный Багданау и др. (Bahdanau attention).
Данный слой вычисляет веса внимания путем сложения
запроса и ключа после их линейного преобразования,
а затем применяет функцию активации. Первым параметром
слой принимает размерность пространства внимания units.
Вторым параметром можно передать функцию активации use_scale,
а также логический флаг causal для маскирования
будущих позиций.
Синтаксис
tf.keras.layers.AdditiveAttention(
use_scale=False,
**kwargs
)
Пример
Давайте создадим слой аддитивного внимания и применим его к запросу и ключу:
import tensorflow as tf
tf.random.set_seed(0)
query = tf.constant([[[1.0, 0.0], [0.0, 1.0]]])
value = tf.constant([[[1.0, 2.0], [3.0, 4.0]]])
key = tf.constant([[[1.0, 0.0], [0.0, 1.0]]])
layer = tf.keras.layers.AdditiveAttention()
res = layer([query, value, key])
print(res)
Результат выполнения кода:
tf.Tensor(
[[[1. 2.]
[3. 4.]]], shape=(1, 2, 2), dtype=float32)
Пример
Давайте создадим слой с включенным параметром
use_scale и передадим маску для внимания:
import tensorflow as tf
tf.random.set_seed(0)
query = tf.constant([[[1.0, 0.0], [0.0, 1.0]]])
value = tf.constant([[[1.0, 2.0], [3.0, 4.0]]])
key = tf.constant([[[1.0, 0.0], [0.0, 1.0]]])
mask = tf.constant([[True, False]])
layer = tf.keras.layers.AdditiveAttention(use_scale=True)
res = layer([query, value, key], mask=[mask, mask])
print(res)
Результат выполнения кода:
tf.Tensor(
[[[1. 2.]
[3. 4.]]], shape=(1, 2, 2), dtype=float32)
Пример
Давайте используем слой аддитивного внимания внутри модели Keras:
Результат выполнения кода:
Model: "model"
__________________________________________________________________________________________________
Layer (type) Output Shape Param # Connected to
==================================================================================================
input_1 (InputLayer) [(None, 2, 2)] 0 []
input_2 (InputLayer) [(None, 2, 2)] 0 []
input_3 (InputLayer) [(None, 2, 2)] 0 []
additive_attention (Additive (None, 2, 2) 0 ['input_1[0][0]',
Attention) 'input_2[0][0]',
'input_3[0][0]']
==================================================================================================
Total params: 0
Trainable params: 0
Non-trainable params: 0
__________________________________________________________________________________________________
Смотрите также
-
класс
Attention,
который реализует механизм внимания Луонга -
класс
MultiHeadAttention,
который реализует многоголовое внимание -
класс
LSTM,
который реализует долгую краткосрочную память -
класс
GRU,
который реализует управляемый рекуррентный блок