tf.keras.layers.Attention
| Просмотр исходного кода на GitHub |
Слой внимания с точечным произведением, также известный как внимание по Луонгу.
tf.keras.layers.Attention(
use_scale=False, score_mode='dot', **kwargs
)
Входными данными являются query тензор формы [batch_size, Tq, dim], value тензор формы [batch_size, Tv, dim] и key тензор формы [batch_size, Tv, dim]. Вычисление выполняется по следующим шагам:
- Вычисление оценок с формой
[batch_size, Tq, Tv]какquery-keyточечного произведения:scores = tf.matmul(query, key, transpose_b=True). - Использование оценок для вычисления распределения с формой
[batch_size, Tq, Tv]:distribution = tf.nn.softmax(scores). - Использование
distributionдля создания линейной комбинацииvalueс формой[batch_size, Tq, dim]:return tf.matmul(distribution, value).
| Аргументы | |
|---|---|
use_scale | Если True, создаст скалярную переменную для масштабирования оценок внимания. |
causal | Булево значение. Установите в True для самовнимания декодера. Добавляет маску, такую, что позиция i не может обращаться к позициям j > i. Это предотвращает поток информации из будущего в прошлое. По умолчанию False. |
dropout | Вещественное число от 0 до 1. Доля единиц, которые нужно отбросить для оценок внимания. По умолчанию 0,0. |
score_mode | Функция для вычисления оценок внимания, одна из {"dot", "concat"}. "dot" относится к точечному произведению между векторами запроса и ключа. "concat" относится к гиперболическому тангенсу конкатенации векторов запроса и ключа. |
Аргументы вызова:
-
inputs: Список следующих тензоров:- запрос: Запрос
Tensorформы[batch_size, Tq, dim]. - значение: Значение
Tensorформы[batch_size, Tv, dim]. - ключ: Необязательный ключ
Tensorформы[batch_size, Tv, dim]. Если не задан, будет использоватьсяvalueдля обоихkeyиvalue, что является наиболее распространенным случаем.
- запрос: Запрос
-
mask: Список следующих тензоров:- маска запроса: Булева маска
Tensorформы[batch_size, Tq]. Если задана, выход будет равен нулю в позициях, гдеmask==False. - маска значения: Булева маска
Tensorформы[batch_size, Tv]. Если задана, маска будет применяться таким образом, что значения в позициях, гдеmask==Falseне будут вносить вклад в результат.
- маска запроса: Булева маска
-
return_attention_scores: bool, еслиTrue, возвращает оценки внимания (после маскирования и применения softmax) как дополнительный аргумент вывода. -
training: булево значение Python, указывающее, должен ли слой работать в режиме обучения (добавляя дропаут) или в режиме вывода (без дропаута).
Вывод:
Выходные данные внимания формы [batch_size, Tq, dim]. [Необязательно] Оценки внимания после маскирования и применения softmax с формой [batch_size, Tq, Tv].
Значение query, value и key зависят от приложения. Например, в случае сходства текстов query — это вложения последовательностей первого текста, а value — это вложения последовательностей второго текста. key обычно является тем же тензором, что и value.
Вот пример использования Attention в сети CNN+Attention:
# Variable-length int sequences.
query_input = tf.keras.Input(shape=(None,), dtype='int32')
value_input = tf.keras.Input(shape=(None,), dtype='int32')
# Embedding lookup.
token_embedding = tf.keras.layers.Embedding(input_dim=1000, output_dim=64)
# Query embeddings of shape [batch_size, Tq, dimension].
query_embeddings = token_embedding(query_input)
# Value embeddings of shape [batch_size, Tv, dimension].
value_embeddings = token_embedding(value_input)
# CNN layer.
cnn_layer = tf.keras.layers.Conv1D(
filters=100,
kernel_size=4,
# Use 'same' padding so outputs have the same shape as inputs.
padding='same')
# Query encoding of shape [batch_size, Tq, filters].
query_seq_encoding = cnn_layer(query_embeddings)
# Value encoding of shape [batch_size, Tv, filters].
value_seq_encoding = cnn_layer(value_embeddings)
# Query-value attention of shape [batch_size, Tq, filters].
query_value_attention_seq = tf.keras.layers.Attention()(
[query_seq_encoding, value_seq_encoding])
# Reduce over the sequence axis to produce encodings of shape
# [batch_size, filters].
query_encoding = tf.keras.layers.GlobalAveragePooling1D()(
query_seq_encoding)
query_value_attention = tf.keras.layers.GlobalAveragePooling1D()(
query_value_attention_seq)
# Concatenate query and document encodings to produce a DNN input layer.
input_layer = tf.keras.layers.Concatenate()(
[query_encoding, query_value_attention])
# Add DNN layers, and create Model.
# ...
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/keras/layers/Attention