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) – Тензор запросов; форма
-
key (Tensor) – Тензор ключей; форма или , если передан
block_table. -
value (Tensor) – Тензор значений; форма или , если передан
block_table. - cu_seq_q (Tensor) – Кумулятивные позиции последовательностей для запросов; форма
- cu_seq_k (Tensor) – Кумулятивные позиции последовательностей для ключей/значений; форма
- 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 разделяется группой из голов query, поэтому должно делиться нацело на . По умолчанию — False.
-
seqused_k (Tensor, optional) – Количество допустимых токенов KV для каждого элемента пакета; форма . Если задано, в attention участвуют только первые
seqused_k[i]токенов последовательности key/value для элемента пакета i. Полезно при декодировании с KV-кэшем, когда слот кэша больше фактической последовательности. Только для инференса (не поддерживается при обратном проходе). -
block_table (Tensor, optional) –
Таблица блоков для страничного KV-кэша; форма , тип данных
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; форма .
-
If return_aux is not None and return_aux.lse is True: -
lse (Tensor): Log-sum-exp оценок attention; форма .
-
- Тип возвращаемого значения:
-
output (Tensor)
- Обозначения форм:
-
- : Размер пакета
- : Общее количество токенов запросов в пакете (сумма длин всех последовательностей запросов)
- : Общее количество токенов ключей/значений в пакете (сумма длин всех последовательностей ключей/значений)
- : Количество голов attention для запросов
- : Количество голов attention для ключей/значений (равно , если GQA не включён)
- : Размерность головы
Пример:
>>> 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вместо выделения нового.
-
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