torch.nn.functional.scaled_dot_product_attention
-
torch.nn.functional.scaled_dot_product_attention()[source] -
- scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0,
-
is_causal=False, scale=None, enable_gqa=False) -> Tensor:
Вычисляет масштабированное скалярное произведение для тензоров query, key и value, используя при наличии необязательную маску внимания и применяя dropout, если указана вероятность больше 0.0. Необязательный аргумент scale можно указать только как именованный аргумент.
# Efficient implementation equivalent to the following: def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None, enable_gqa=False) -> torch.Tensor: L, S = query.size(-2), key.size(-2) scale_factor = 1 / math.sqrt(query.size(-1)) if scale is None else scale attn_bias = torch.zeros(L, S, dtype=query.dtype, device=query.device) if is_causal: assert attn_mask is None temp_mask = torch.ones(L, S, dtype=torch.bool, device=query.device).tril(diagonal=0) attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf")) if attn_mask is not None: if attn_mask.dtype == torch.bool: attn_bias.masked_fill_(attn_mask.logical_not(), float("-inf")) else: attn_bias = attn_mask + attn_bias if enable_gqa: key = key.repeat_interleave(query.size(-3)//key.size(-3), -3) value = value.repeat_interleave(query.size(-3)//value.size(-3), -3) attn_weight = query @ key.transpose(-2, -1) * scale_factor attn_weight += attn_bias attn_weight = torch.softmax(attn_weight, dim=-1) attn_weight = torch.dropout(attn_weight, dropout_p, train=True) return attn_weight @ valueПредупреждение
Эта функция находится в бета-версии и может измениться.
Предупреждение
Эта функция всегда применяет dropout в соответствии с указанным аргументом
dropout_p. Чтобы отключить dropout во время оценки, обязательно передайте значение0.0, если модуль, вызывающий функцию, не находится в режиме обучения.Например:
class MyModel(nn.Module): def __init__(self, p=0.5): super().__init__() self.p = p def forward(self, ...): return F.scaled_dot_product_attention(..., dropout_p=(self.p if self.training else 0.0))Примечание
Семантика булевой маски для
attn_maskпротивоположна семантикеkey_padding_maskуMultiheadAttention.В
scaled_dot_product_attention()Trueуказывает значения, которые должны участвовать во внимании.В
MultiheadAttentionTrueуказывает значения, которые следует замаскировать (заполнение).При переходе с MHA инвертируйте булеву маску (например, с помощью
~maskилиmask.logical_not()).Примечание
В настоящее время поддерживаются три реализации масштабированного скалярного произведения для внимания:
- FlashAttention-2: ускоренное внимание с улучшенным параллелизмом и распределением работы
- Эффективное по памяти внимание
- Реализация PyTorch на C++, соответствующая приведенной выше формуле
При использовании бэкенда CUDA функция может вызывать оптимизированные ядра для повышения производительности. Для всех остальных бэкендов будет использоваться реализация PyTorch.
Все реализации включены по умолчанию. Масштабированное скалярное произведение для внимания пытается автоматически выбрать наиболее оптимальную реализацию на основе входных данных. Для более точного управления используемой реализацией доступны следующие функции включения и отключения реализаций. Предпочтительным механизмом является менеджер контекста:
-
torch.nn.attention.sdpa_kernel(): менеджер контекста для включения или отключения любой из реализаций. -
torch.backends.cuda.enable_flash_sdp(): глобально включает или отключает FlashAttention. -
torch.backends.cuda.enable_mem_efficient_sdp(): глобально включает или отключает эффективное по памяти внимание. -
torch.backends.cuda.enable_math_sdp(): глобально включает или отключает реализацию PyTorch на C++.
У каждого объединенного ядра есть определенные ограничения на входные данные. Если пользователю требуется использовать конкретную объединенную реализацию, отключите реализацию PyTorch на C++ с помощью
torch.nn.attention.sdpa_kernel(). Если объединенная реализация недоступна, будет выдано предупреждение с причинами, по которым ее нельзя запустить.Из-за особенностей объединения операций с плавающей точкой результат этой функции может различаться в зависимости от выбранного ядра бэкенда. Реализация на C++ поддерживает torch.float64 и может использоваться, когда требуется более высокая точность. В математическом бэкенде все промежуточные значения сохраняются в torch.float, если входные данные имеют тип torch.half или torch.bfloat16.
Дополнительные сведения см. в разделе Численная точность
Grouped Query Attention (GQA) — экспериментальная функция. Она работает с FlashAttention, attention cuDNN и математическим ядром для тензоров CUDA. Эффективное по памяти внимание также поддерживает GQA на NVIDIA CUDA. GQA не поддерживает вложенные тензоры. Ограничения GQA:
- number_of_heads_query % number_of_heads_key_value == 0 и
- number_of_heads_key == number_of_heads_value
Примечание
В некоторых случаях при передаче тензоров на устройстве CUDA и использовании CuDNN этот оператор может выбрать недетерминированный алгоритм для повышения производительности. Если это нежелательно, можно попробовать сделать операцию детерминированной (возможно, ценой производительности), установив
torch.backends.cudnn.deterministic = True. Дополнительные сведения см. в разделе Воспроизводимость.- Параметры:
-
- query (Tensor) – Тензор query; форма .
- key (Tensor) – Тензор key; форма .
- value (Tensor) – Тензор value; форма .
- attn_mask (необязательный Tensor) – Маска внимания; ее форма должна допускать broadcasting к форме весов внимания, то есть . Поддерживаются два типа масок. Булева маска, в которой значение True указывает, что элемент должен участвовать во внимании. Маска с плавающей точкой того же типа, что и query, key и value, которая добавляется к оценке внимания.
- dropout_p (float) – Вероятность dropout; если она больше 0.0, применяется dropout
-
is_causal (bool) – Если задано значение true, маска внимания представляет собой нижнюю треугольную матрицу, когда маска является квадратной матрицей. Если матрица маски неквадратная, маска внимания имеет вид причинного смещения в верхнем левом углу из-за выравнивания (см.
torch.nn.attention.bias.CausalBias). Если заданы и attn_mask, и is_causal, возникает ошибка. - scale (необязательный python:float, только именованный аргумент) – Коэффициент масштабирования, применяемый перед softmax. Если значение равно None, по умолчанию используется .
- enable_gqa (bool) – Если задано значение True, включается Grouped Query Attention (GQA); по умолчанию установлено значение False.
- Возвращает:
-
Результат внимания; форма .
- Тип возвращаемого значения:
-
output (Tensor)
- Обозначения форм:
-
Примеры
>>> # Optionally use the context manager to ensure one of the fused kernels is run >>> query = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda") >>> key = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda") >>> value = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda") >>> with sdpa_kernel(backends=[SDPBackend.FLASH_ATTENTION]): >>> F.scaled_dot_product_attention(query,key,value)
>>> # Sample for GQA for llama3 >>> query = torch.rand(32, 32, 128, 64, dtype=torch.float16, device="cuda") >>> key = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda") >>> value = torch.rand(32, 8, 128, 64, dtype=torch.float16, device="cuda") >>> with sdpa_kernel(backends=[SDPBackend.MATH]): >>> F.scaled_dot_product_attention(query,key,value,enable_gqa=True)
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.functional.scaled_dot_product_attention.html