Функция multi_head_attention_forward
Функция F.multi_head_attention_forward реализует механизм многоголового внимания, который является ключевым компонентом архитектуры Transformer. Она принимает запросы (query), ключи (key) и значения (value), а также множество дополнительных параметров, таких как количество голов, вероятность dropout и маски. Функция возвращает тензор с результатами внимания и тензор весов внимания.
Синтаксис
torch.nn.functional.multi_head_attention_forward(
query,
key,
value,
embed_dim_to_check,
num_heads,
in_proj_weight,
in_proj_bias,
bias_k,
bias_v,
add_zero_attn,
dropout_p,
out_proj_weight,
out_proj_bias,
training=True,
key_padding_mask=None,
need_weights=True,
attn_mask=None,
use_separate_proj_weight=False,
q_proj_weight=None,
k_proj_weight=None,
v_proj_weight=None,
static_k=None,
static_v=None
)
Описание параметров
Основные параметры функции:
-
query- тензор запросов размерности (L, N, E), где L - длина последовательности запросов, N - размер батча, E - размерность эмбеддинга -
key- тензор ключей размерности (S, N, E), где S - длина последовательности ключей -
value- тензор значений размерности (S, N, E) -
num_heads- количество голов внимания -
dropout_p- вероятность dropout (0.0 - отключить) -
key_padding_mask- маска для игнорирования позиций в ключах -
attn_mask- маска для предотвращения внимания к определённым позициям
Пример
Базовый пример использования функции для вычисления многоголового внимания с двумя головами:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
# Создаем входные данные
N = 2 # размер батча
E = 8 # размерность эмбеддинга
L = 4 # длина запросов
S = 6 # длина ключей/значений
num_heads = 2
query = torch.randn(L, N, E)
key = torch.randn(S, N, E)
value = torch.randn(S, N, E)
# Создаем веса проекций
in_proj_weight = torch.randn(3 * E, E)
in_proj_bias = torch.randn(3 * E)
out_proj_weight = torch.randn(E, E)
out_proj_bias = torch.randn(E)
# Вычисляем внимание
attn_output, attn_weights = F.multi_head_attention_forward(
query,
key,
value,
embed_dim_to_check=E,
num_heads=num_heads,
in_proj_weight=in_proj_weight,
in_proj_bias=in_proj_bias,
bias_k=None,
bias_v=None,
add_zero_attn=False,
dropout_p=0.0,
out_proj_weight=out_proj_weight,
out_proj_bias=out_proj_bias,
training=False,
need_weights=True
)
print(attn_output.shape)
print(attn_weights.shape)
Результат выполнения кода:
torch.Size([4, 2, 8])
torch.Size([2, 4, 6])
Пример
Использование маски внимания для предотвращения взаимодействия между определёнными позициями:
import torch
import torch.nn.functional as F
torch.manual_seed(42)
N = 1
E = 6
L = 3
S = 5
num_heads = 3
query = torch.randn(L, N, E)
key = torch.randn(S, N, E)
value = torch.randn(S, N, E)
in_proj_weight = torch.randn(3 * E, E)
in_proj_bias = torch.randn(3 * E)
out_proj_weight = torch.randn(E, E)
out_proj_bias = torch.randn(E)
# Создаем маску (запрещаем внимание к последним двум позициям)
attn_mask = torch.zeros(L, S)
attn_mask[:, -2:] = float('-inf')
attn_output, attn_weights = F.multi_head_attention_forward(
query,
key,
value,
embed_dim_to_check=E,
num_heads=num_heads,
in_proj_weight=in_proj_weight,
in_proj_bias=in_proj_bias,
bias_k=None,
bias_v=None,
add_zero_attn=False,
dropout_p=0.0,
out_proj_weight=out_proj_weight,
out_proj_bias=out_proj_bias,
training=False,
need_weights=True,
attn_mask=attn_mask
)
print(attn_weights)
Результат выполнения кода:
tensor([[[0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[2.5772e-01, 2.7042e-01, 2.2919e-01, 2.2134e-01, 2.2134e-02],
[0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00]]])
Пример
Использование маски ключей для игнорирования заполняющих токенов:
import torch
import torch.nn.functional as F
torch.manual_seed(123)
N = 2
E = 4
L = 3
S = 5
num_heads = 2
query = torch.randn(L, N, E)
key = torch.randn(S, N, E)
value = torch.randn(S, N, E)
in_proj_weight = torch.randn(3 * E, E)
in_proj_bias = torch.randn(3 * E)
out_proj_weight = torch.randn(E, E)
out_proj_bias = torch.randn(E)
# Создаем маску ключей (True - игнорировать позицию)
key_padding_mask = torch.zeros(N, S, dtype=torch.bool)
key_padding_mask[0, -2:] = True # игнорируем последние 2 позиции в первом примере
key_padding_mask[1, -3:] = True # игнорируем последние 3 позиции во втором примере
attn_output, attn_weights = F.multi_head_attention_forward(
query,
key,
value,
embed_dim_to_check=E,
num_heads=num_heads,
in_proj_weight=in_proj_weight,
in_proj_bias=in_proj_bias,
bias_k=None,
bias_v=None,
add_zero_attn=False,
dropout_p=0.0,
out_proj_weight=out_proj_weight,
out_proj_bias=out_proj_bias,
training=False,
need_weights=True,
key_padding_mask=key_padding_mask
)
print(attn_output.shape)
print(attn_weights[0, 0, -2:])
Результат выполнения кода:
torch.Size([3, 2, 4])
tensor([0.0000, 0.0000])
Смотрите также
-
функцию
scaled_dot_product_attention,
которая реализует масштабированное скалярное произведение внимания -
функцию
linear,
которая применяет линейное преобразование к входным данным -
функцию
dropout,
которая применяет регуляризацию dropout -
функцию
softmax,
которая вычисляет распределение вероятностей для весов внимания