Spec-Zone.ru › PyTorch 2

torch.nn.functional.scaled_dot_product_attention

torch.nn.functional.scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None) → Tensor:

Вычисляет взвешенное скалярное произведение внимания для тензоров запроса, ключа и значения, используя необязательную маску внимания, если она передана, и применяет дропаут, если вероятность больше 0,0.

# 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) -> torch.Tensor:
    # Efficient implementation equivalent to the following:
    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)
    if is_causal:
        assert attn_mask is None
        temp_mask = torch.ones(L, S, dtype=torch.bool).tril(diagonal=0)
        attn_bias.masked_fill_(temp_mask.logical_not(), float("-inf"))
        attn_bias.to(query.dtype)

    if attn_mask is not None:
        if attn_mask.dtype == torch.bool:
            attn_mask.masked_fill_(attn_mask.logical_not(), float("-inf"))
        else:
            attn_bias += attn_mask
    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

Предупреждение

Эта функция находится в стадии бета-тестирования и может быть изменена.

Примечание

В настоящее время поддерживаются три реализации взвешенного скалярного произведения внимания:

  • FlashAttention: быстрое и экономичное по памяти точное внимание с учетом операций ввода-вывода
  • Эффективное по памяти внимание
  • Реализация в PyTorch, определённая в C++, соответствующая вышеприведённой формулировке

Функция может вызывать оптимизированные ядра для повышения производительности при использовании CUDA-бекенда. Для всех других бекендов будет использоваться реализация в PyTorch.

Все реализации включены по умолчанию. Взвешенное скалярное произведение внимания пытается автоматически выбрать наиболее оптимальную реализацию на основе входных данных. Для более точного управления используемой реализацией предоставляются следующие функции для включения и отключения реализаций. Предпочтительным механизмом является контекстный менеджер:

  • torch.backends.cuda.sdp_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.backends.cuda.sdp_kernel(). В случае отсутствия объединённой реализации будет выведено сообщение об ошибке с указанием причин.

Из-за особенностей объединения операций с плавающей точкой, результат этой функции может отличаться в зависимости от выбранного ядра бекенда. Реализация на C++ поддерживает torch.float64 и может использоваться, когда требуется более высокая точность. Для получения дополнительной информации см. Числовая точность

Примечание

В некоторых случаях, когда тензоры находятся на устройстве CUDA и используется CuDNN, этот оператор может выбрать недетерминированный алгоритм для повышения производительности. Если это нежелательно, можно попробовать сделать операцию детерминированной (возможно, с потерей производительности) путём установки torch.backends.cudnn.deterministic = True. Подробнее см. Воспроизводимость.

Параметры
  • query (Тензор) – Тензор запроса; форма (N,...,L,E)(N, ..., L, E).
  • key (Тензор) – Тензор ключа; форма (N,...,S,E)(N, ..., S, E).
  • value (Тензор) – Тензор значения; форма (N,...,S,Ev)(N, ..., S, Ev).
  • attn_mask (необязательный Тензор) – Маска внимания; форма (N,...,L,S)(N, ..., L, S). Поддерживаются два типа масок. Булевая маска, где значение True означает, что элемент *должен* участвовать во внимании. Маска с плавающей точкой того же типа, что и query, key, value, которая добавляется к оценке внимания.
  • dropout_p (число с плавающей точкой) – Вероятность дропаута; если больше 0,0, применяется дропаут
  • is_causal (логическое) – Если True, предполагается маска каузального внимания и генерируется ошибка, если установлены и attn_mask, и is_causal.
  • scale (необязательное python:число с плавающей точкой) – Коэффициент масштабирования, применяемый перед softmax. Если None, по умолчанию устанавливается 1E\frac{1}{\sqrt{E}}.
Возвращает

Вывод внимания; форма (N,...,L,Ev)(N, ..., L, Ev).

Тип возвращаемого значения

вывод (Тензор)

Легенда форм:
  • N:Размер пакета...:Любое число других размерностей пакета (необязательно)N: \text{Размер пакета} ... : \text{Любое число других размерностей пакета (необязательно)}
  • S:Длина исходной последовательностиS: \text{Длина исходной последовательности}
  • L:Длина целевой последовательностиL: \text{Длина целевой последовательности}
  • E:Размерность встраивания запроса и ключаE: \text{Размерность встраивания запроса и ключа}
  • Ev:Размерность встраивания значенияEv: \text{Размерность встраивания значения}

Примеры:

>>> # 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 torch.backends.cuda.sdp_kernel(enable_math=False):
>>>     F.scaled_dot_product_attention(query,key,value)

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.functional.scaled_dot_product_attention.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API