РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
445 of 769 menu

Функция 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,
    которая вычисляет распределение вероятностей для весов внимания
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить