Spec-Zone.ru › PyTorch 2.14

torch.nn.attention.varlen

Создано: 14 окт. 2025 г. | Последнее обновление: 10 мар. 2026 г.

Реализация attention с переменной длиной последовательности с использованием Flash Attention.

Этот модуль предоставляет высокоуровневый интерфейс Python для attention с переменной длиной последовательности, вызывающий оптимизированные ядра Flash Attention.

torch.nn.attention.varlen.varlen_attn(query, key, value, cu_seq_q, cu_seq_k, max_q, max_k, *, return_aux=None, scale=None, window_size=(-1, -1), enable_gqa=False, seqused_k=None, block_table=None, num_splits=None) [исходный код]

Вычисляет attention с переменной длиной последовательности с использованием Flash Attention.

Эта функция похожа на scaled_dot_product_attention, но оптимизирована для последовательностей переменной длины с использованием тензоров кумулятивных позиций последовательностей.

Параметры:
  • query (Tensor) – Тензор запросов; форма (Tq,Hq,D)(T_q, H_q, D)
  • key (Tensor) – Тензор ключей; форма (Tk,Hkv,D)(T_k, H_{kv}, D) или (total_pages,page_size,Hkv,D)(\text{total\_pages}, \text{page\_size}, H_{kv}, D), если передан block_table.
  • value (Tensor) – Тензор значений; форма (Tk,Hkv,D)(T_k, H_{kv}, D) или (total_pages,page_size,Hkv,D)(\text{total\_pages}, \text{page\_size}, H_{kv}, D), если передан block_table.
  • cu_seq_q (Tensor) – Кумулятивные позиции последовательностей для запросов; форма (N+1,)(N+1,)
  • cu_seq_k (Tensor) – Кумулятивные позиции последовательностей для ключей/значений; форма (N+1,)(N+1,)
  • max_q (int) – Максимальная длина последовательности запросов в пакете.
  • max_k (int) – Максимальная длина последовательности ключей/значений в пакете.
  • return_aux (Optional[AuxRequest]) – Если значение не равно None и return_aux.lse равно True, также возвращает тензор logsumexp.
  • scale (float, optional) – Коэффициент масштабирования оценок attention
  • window_size (tuple[int, int], optional) – Размер окна для attention со скользящим окном в формате (left, right). Используйте (-1, -1) для полного attention (по умолчанию), (-1, 0) для причинно-следственного attention или (W, 0) для причинно-следственного attention со скользящим окном размера W.
  • enable_gqa (bool) – Если установлено в True, включает Grouped Query Attention (GQA) и позволяет использовать для key/value меньше голов, чем для query. Каждая голова KV разделяется группой из Hq/HkvH_q / H_{kv} голов query, поэтому HqH_q должно делиться нацело на HkvH_{kv}. По умолчанию — False.
  • seqused_k (Tensor, optional) – Количество допустимых токенов KV для каждого элемента пакета; форма (N,)(N,). Если задано, в attention участвуют только первые seqused_k[i] токенов последовательности key/value для элемента пакета i. Полезно при декодировании с KV-кэшем, когда слот кэша больше фактической последовательности. Только для инференса (не поддерживается при обратном проходе).
  • block_table (Tensor, optional) –

    Таблица блоков для страничного KV-кэша; форма (N,max_pages_per_seq)(N, \text{max\_pages\_per\_seq}), тип данных int32. Требуется seqused_k. Только для инференса (не поддерживается при обратном проходе).

    Когда передан block_table, key и value представляют собой «пул» страниц с токенами данных KV, причём страницы могут относиться к любой последовательности или порядку. block_table сопоставляет логические фрагменты каждой последовательности с физическими страницами в этом пуле.

    seqused_k[i] сообщает ядру, сколько токенов в последовательности i фактически допустимы, поскольку последняя страница обычно заполнена лишь частично.

  • num_splits (int, optional) – Количество разбиений для split-KV. Задайте 1, чтобы отключить split-KV, что обеспечивает инвариантность к составу пакета. Split-KV распараллеливает измерение последовательности key/value по нескольким блокам потоков и объединяет частичные результаты. Решение о разбиении зависит от max_k (самой длинной последовательности в пакете), поэтому изменение состава пакета может изменить порядок редукции и привести к различающимся результатам с плавающей точкой для одной и той же последовательности. Если эта функция отключена, для данной последовательности гарантируются побитово идентичные результаты независимо от того, какие другие последовательности входят в пакет, однако при небольшом количестве запросов загрузка GPU будет ниже. Если задано None (по умолчанию), ядро выбирает значение автоматически.
Возвращает:

Выходной тензор, полученный в результате вычисления attention; форма (Tq,Hq,D)(T_q, H_q, D).

If return_aux is not None and return_aux.lse is True:

lse (Tensor): Log-sum-exp оценок attention; форма (Hq,Tq)(H_q, T_q).

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

output (Tensor)

Обозначения форм:
  • NN: Размер пакета
  • TqT_q: Общее количество токенов запросов в пакете (сумма длин всех последовательностей запросов)
  • TkT_k: Общее количество токенов ключей/значений в пакете (сумма длин всех последовательностей ключей/значений)
  • HqH_q: Количество голов attention для запросов
  • HkvH_{kv}: Количество голов attention для ключей/значений (равно HqH_q, если GQA не включён)
  • DD: Размерность головы

Пример:

>>> batch_size, max_seq_len, embed_dim, num_heads = 2, 512, 1024, 16
>>> head_dim = embed_dim // num_heads
>>> seq_lengths = []
>>> for _ in range(batch_size):
...     length = torch.randint(1, max_seq_len // 64 + 1, (1,)).item() * 64
...     seq_lengths.append(min(length, max_seq_len))
>>> seq_lengths = torch.tensor(seq_lengths, device="cuda")
>>> total_tokens = seq_lengths.sum().item()
>>>
>>> # Create packed query, key, value tensors
>>> query = torch.randn(
...     total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda"
... )
>>> key = torch.randn(
...     total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda"
... )
>>> value = torch.randn(
...     total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda"
... )
>>>
>>> # Build cumulative sequence tensor
>>> cu_seq = torch.zeros(batch_size + 1, device="cuda", dtype=torch.int32)
>>> cu_seq[1:] = seq_lengths.cumsum(0)
>>> max_len = seq_lengths.max().item()
>>>
>>> # Call varlen_attn
>>> output = varlen_attn(
...     query, key, value, cu_seq, cu_seq, max_len, max_len
... )
torch.nn.attention.varlen.varlen_attn_out(out, query, key, value, cu_seq_q, cu_seq_k, max_q, max_k, *, return_aux=None, scale=None, window_size=(-1, -1), enable_gqa=False, seqused_k=None, block_table=None, num_splits=None) [исходный код]

Вычисляет attention с переменной длиной последовательности с использованием Flash Attention и предварительно выделенного выходного тензора.

То же, что и varlen_attn(), но записывает результат attention в переданный тензор out вместо выделения нового.

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

Tensor | tuple[Tensor, Tensor]

class torch.nn.attention.varlen.AuxRequest(lse=False) [исходный код]

Запрос на вычисление вспомогательных выходных данных из varlen_attn.

Каждое поле представляет собой логическое значение, указывающее, следует ли вычислять соответствующие вспомогательные выходные данные.

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/nn.attention.varlen.html

Spec-Zone.ru

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