torch.nn.attention.flex_attention
Создано: 16 июля 2024 г. | Последнее обновление: 7 августа 2026 г.
-
torch.nn.attention.flex_attention.flex_attention(query: Tensor, key: Tensor, value: Tensor, score_mod: Callable[[Tensor, Tensor, Tensor, Tensor, Tensor], Tensor] | None = None, block_mask: BlockMask | None = None, scale: float | None = None, enable_gqa: bool = False, return_lse: Literal[False] = False, kernel_options: FlexKernelOptions | None = None, *, return_aux: None = None) → Tensor[исходный код] - torch.nn.attention.flex_attention.flex_attention(query:Tensor, key:Tensor, value:Tensor, score_mod:Callable[[Tensor,Tensor,Tensor,Tensor,Tensor],Tensor]|None=None, block_mask:BlockMask|None=None, scale:float|None=None, enable_gqa:bool=False, return_lse:Literal[True]=False, kernel_options:FlexKernelOptions|None=None, *, return_aux:None=None) tuple[Tensor,Tensor]
- torch.nn.attention.flex_attention.flex_attention(query:Tensor, key:Tensor, value:Tensor, score_mod:Callable[[Tensor,Tensor,Tensor,Tensor,Tensor],Tensor]|None=None, block_mask:BlockMask|None=None, scale:float|None=None, enable_gqa:bool=False, return_lse:bool=False, kernel_options:FlexKernelOptions|None=None, *, return_aux:AuxRequest) tuple[Tensor,AuxOutput]
- torch.nn.attention.flex_attention.flex_attention(query:Tensor, key:Tensor, value:Tensor, score_mod:Callable[[Tensor,Tensor,Tensor,Tensor,Tensor],Tensor]|None=None, block_mask:BlockMask|None=None, scale:float|None=None, enable_gqa:bool=False, return_lse:Literal[True]=False, kernel_options:FlexKernelOptions|None=None, *, return_aux:AuxRequest) Never
-
Эта функция реализует масштабированное скалярное произведение для механизма внимания с произвольной функцией модификации оценок внимания, описанной в статье Flex Attention. См. также публикацию в блоге.
Эта функция вычисляет масштабированное скалярное произведение для механизма внимания между тензорами query, key и value с помощью пользовательской функции модификации оценок внимания. Функция модификации оценок внимания применяется после вычисления оценок внимания между тензорами query и key. Оценки внимания вычисляются следующим образом:
Функция
score_modдолжна иметь следующую сигнатуру:def score_mod( score: Tensor, batch: Tensor, head: Tensor, q_idx: Tensor, k_idx: Tensor ) -> Tensor:- Где:
-
-
score: скалярный тензор, представляющий оценку внимания, с тем же типом данных и на том же устройстве, что и тензоры query, key и value. -
batch,head,q_idx,k_idx: скалярные тензоры, указывающие соответственно индекс пакета, индекс головы query, индекс query и индекс key/value. Они должны иметь тип данныхtorch.intи находиться на том же устройстве, что и тензор оценки.
-
- Параметры:
-
- query (Tensor) – Тензор query; форма . Для типов данных FP8 для оптимальной производительности следует использовать построчную раскладку в памяти.
- key (Tensor) – Тензор key; форма . Для типов данных FP8 для оптимальной производительности следует использовать построчную раскладку в памяти.
- value (Tensor) – Тензор value; форма . Для типов данных FP8 для оптимальной производительности следует использовать столбцовую раскладку в памяти.
- score_mod (Optional[Callable]) – Функция для модификации оценок внимания. По умолчанию score_mod не применяется.
- block_mask (Optional[BlockMask]) – Объект BlockMask, управляющий шаблоном разреженности блоков внимания.
- scale (Optional[float]) – Коэффициент масштабирования, применяемый перед softmax. Если он не задан, используется значение по умолчанию .
- enable_gqa (bool) – Если установлено значение True, включает Grouped Query Attention (GQA) и распределяет головы key/value между головами query.
-
return_lse (bool) – Возвращать ли logsumexp оценок внимания. По умолчанию — False. Устарело: вместо этого используйте
return_aux=AuxRequest(lse=True). -
kernel_options (Optional[FlexKernelOptions]) – Параметры для управления поведением базовых ядер Triton. Доступные параметры и примеры использования см. в разделе
FlexKernelOptions. -
return_aux (Optional[AuxRequest]) – Определяет, какие вспомогательные выходные данные следует вычислить и вернуть. Если значение равно None, возвращается только результат механизма внимания. Используйте
AuxRequest(lse=True, max_scores=True), чтобы запросить оба вспомогательных выходных значения.
- Возвращает:
-
Результат механизма внимания; форма .
-
When return_aux is not None: -
aux (AuxOutput): вспомогательные выходные данные с заполненными запрошенными полями.
-
When return_aux is None (deprecated paths): -
lse (Tensor): log-sum-exp оценок внимания; форма . Возвращается только при
return_lse=True.
-
- Тип возвращаемого значения:
-
результат (Tensor)
- Обозначения размерностей:
-
Предупреждение
torch.nn.attention.flex_attention— экспериментальная функция в PyTorch. В будущей версии PyTorch ожидается более стабильная реализация. Подробнее о классификации функций: https://pytorch.org/blog/pytorch-feature-classification-changes/#prototype
-
class torch.nn.attention.flex_attention.AuxOutput(lse=None, max_scores=None)[исходный код] -
Вспомогательные выходные данные операции flex_attention.
Поля будут равны None, если они не запрошены; если поле запрошено, оно будет содержать тензор.
-
class torch.nn.attention.flex_attention.AuxRequest(lse=False, max_scores=False)[исходный код] -
Запрос вспомогательных выходных данных для вычисления функцией flex_attention.
Каждое поле — это логическое значение, указывающее, следует ли вычислять соответствующие вспомогательные выходные данные.
Утилиты BlockMask
-
torch.nn.attention.flex_attention.create_block_mask(mask_mod, B, H, Q_LEN, KV_LEN, device=None, BLOCK_SIZE=128, _compile=False, separate_full_blocks=True, compute_dq_write_order=False, dq_kv_order=True)[исходный код] -
Эта функция создает кортеж маски блоков из функции mask_mod.
- Параметры:
-
- mask_mod (Callable) – Функция mask_mod. Это вызываемый объект, задающий шаблон маскирования для механизма внимания. Он принимает четыре аргумента: b (размер пакета), h (число голов), q_idx (индекс запроса) и kv_idx (индекс ключа/значения). Функция должна возвращать булев тензор, указывающий, какие связи внимания разрешены (True), а какие замаскированы (False).
- B (int) – Размер пакета.
- H (int) – Число голов запросов.
- Q_LEN (int) – Длина последовательности запроса.
- KV_LEN (int) – Длина последовательности ключей/значений.
- device (str) – Устройство, на котором создается маска.
- BLOCK_SIZE (int or tuple[int, int]) – Размер блока для маски блоков. Если указано одно целое число, оно используется и для запроса, и для ключа/значения.
- separate_full_blocks (bool) – Если True, полностью незамаскированные блоки хранятся отдельно, чтобы ядра могли пропускать вызов mask_mod для этих блоков. Если False, все непустые блоки хранятся как частичные, а mask_mod применяется к каждому блоку.
- compute_dq_write_order (bool) – Если True, предварительно вычисляются метаданные порядка записи dQ, необходимые для детерминированного обратного прохода FLASH с разреженностью по блокам.
- dq_kv_order (bool) – Порядок планирования столбцов KV, используемый для детерминированного накопления dQ, когда compute_dq_write_order имеет значение True. False означает порядок n-блоков по возрастанию, а True — порядок по убыванию/SPT. В create_block_mask пока не поддерживаются явные расписания тензоров; они поддерживаются в BlockMask.from_kv_blocks для случаев, когда вызывающий код напрямую передает предварительно вычисленные метаданные порядка записи.
- Возвращает:
-
Объект BlockMask, содержащий сведения о маске блоков.
- Тип возвращаемого значения:
- Пример использования:
-
def causal_mask(b, h, q_idx, kv_idx): return q_idx >= kv_idx block_mask = create_block_mask(causal_mask, 1, 1, 8192, 8192, device="cuda") query = torch.randn(1, 1, 8192, 64, device="cuda", dtype=torch.float16) key = torch.randn(1, 1, 8192, 64, device="cuda", dtype=torch.float16) value = torch.randn(1, 1, 8192, 64, device="cuda", dtype=torch.float16) output = flex_attention(query, key, value, block_mask=block_mask)
-
torch.nn.attention.flex_attention.create_mask(mod_fn, B, H, Q_LEN, KV_LEN, device=None)[исходный код] -
Эта функция создает тензор маски из функции mod_fn.
- Параметры:
-
- mod_fn (Union[_score_mod_signature, _mask_mod_signature]) – Функция для изменения оценок внимания.
- B (int) – Размер пакета.
- H (int) – Число голов запросов.
- Q_LEN (int) – Длина последовательности запроса.
- KV_LEN (int) – Длина последовательности ключей/значений.
- device (str) – Устройство, на котором создается маска.
- Возвращает:
-
Тензор маски формы (B, H, M, N).
- Тип возвращаемого значения:
-
mask (Tensor)
-
torch.nn.attention.flex_attention.and_masks(*mask_mods)[исходный код] -
Возвращает mask_mod, представляющую собой пересечение заданных mask_mod
-
torch.nn.attention.flex_attention.or_masks(*mask_mods)[исходный код] -
Возвращает mask_mod, представляющую собой объединение заданных mask_mod
-
torch.nn.attention.flex_attention.noop_mask(batch, head, token_q, token_kv)[исходный код] -
Возвращает пустую операцию mask_mod
- Тип возвращаемого значения:
Параметры FlexKernel
-
class torch.nn.attention.flex_attention.FlexKernelOptions[исходный код] -
Параметры управления поведением ядер FlexAttention.
Эти параметры передаются базовым ядрам Triton для управления производительностью и численным поведением. Большинству пользователей не потребуется задавать эти параметры, поскольку автоматическая настройка по умолчанию обеспечивает хорошую производительность.
К параметрам можно добавлять префиксы
fwd_илиbwd_, чтобы применять их только к прямому или обратному проходу соответственно. Например:fwd_BLOCK_Mиbwd_BLOCK_M1.Примечание
В настоящее время мы не гарантируем обратную совместимость этих параметров. Тем не менее большинство из них остаются довольно стабильными с момента появления. Однако пока мы не считаем эту часть публичным API. Мы считаем, что хоть какая-то документация лучше, чем скрытые флаги, но в будущем можем изменить эти параметры.
- Пример использования:
-
# Using dictionary (backward compatible) kernel_opts = {"BLOCK_M": 64, "BLOCK_N": 64, "PRESCALE_QK": True} output = flex_attention(q, k, v, kernel_options=kernel_opts) # Using TypedDict (recommended for type safety) from torch.nn.attention.flex_attention import FlexKernelOptions kernel_opts: FlexKernelOptions = { "BLOCK_M": 64, "BLOCK_N": 64, "PRESCALE_QK": True, } output = flex_attention(q, k, v, kernel_options=kernel_opts) # Forward/backward specific options kernel_opts: FlexKernelOptions = { "fwd_BLOCK_M": 64, "bwd_BLOCK_M1": 32, "PRESCALE_QK": False, } output = flex_attention(q, k, v, kernel_options=kernel_opts)
-
BACKEND: NotRequired[Literal['AUTO', 'TRITON', 'FLASH', 'TRITON_DECODE']] -
Выбирает конкретную серверную часть ядра.
- Варианты:
-
- “AUTO”: Использовать текущие эвристики (обычно ядра на основе Triton с автоматическим выбором между flex_attention и flex_decoding)
- “TRITON”: Стандартное ядро flex_attention на Triton
- “TRITON_DECODE”: Ядро flex_decoding на Triton, доступное только для коротких последовательностей при определенных конфигурациях
- “FLASH”: Экспериментальный вариант: ядро Flash Attention (cute-dsl); пользователю необходимо установить flash
Этот параметр нельзя сочетать с устаревшими настройками, такими как
FORCE_USE_FLEX_ATTENTION. Если запрошенную серверную часть нельзя использовать, возникает ошибка. Значение по умолчанию: “AUTO”
-
BLOCKS_ARE_CONTIGUOUS: NotRequired[bool] -
Если True, гарантируется, что все блоки в маске смежны. Это позволяет оптимизировать обход блоков. Например, причинно-следственные маски удовлетворяют этому условию, а prefix_lm + скользящее окно — нет. Значение по умолчанию: False.
-
BLOCK_M: NotRequired[int] -
Размер блока потоков для измерения длины последовательности Q в прямом проходе. Должен быть степенью 2. Типичные значения: 16, 32, 64, 128. Значение по умолчанию определяется автоматической настройкой.
-
BLOCK_M1: NotRequired[int] -
Размер блока потоков для измерения Q в обратном проходе. Используйте как ‘bwd_BLOCK_M1’. Значение по умолчанию определяется автоматической настройкой.
-
BLOCK_M2: NotRequired[int] -
Размер блока потоков для второго измерения Q в обратном проходе. Используйте как ‘bwd_BLOCK_M2’. Значение по умолчанию определяется автоматической настройкой.
-
BLOCK_N: NotRequired[int] -
Размер блока потоков для измерения длины последовательности K/V в прямом проходе. Должен быть степенью 2. Типичные значения: 16, 32, 64, 128. Значение по умолчанию определяется автоматической настройкой.
-
BLOCK_N1: NotRequired[int] -
Размер блока потоков для измерения K/V в обратном проходе. Используйте как ‘bwd_BLOCK_N1’. Значение по умолчанию определяется автоматической настройкой.
-
BLOCK_N2: NotRequired[int] -
Размер блока потоков для второго измерения K/V в обратном проходе. Используйте как ‘bwd_BLOCK_N2’. Значение по умолчанию определяется автоматической настройкой.
-
FORCE_USE_FLEX_ATTENTION: NotRequired[bool] -
Если True, принудительно используется ядро flex attention вместо потенциально более оптимизированного ядра flex-decoding для коротких последовательностей. Этот параметр может быть полезен при отладке. Значение по умолчанию: False.
-
PRESCALE_QK: NotRequired[bool] -
Следует ли предварительно масштабировать QK на 1/sqrt(d) и изменить основание. Это немного быстрее, но может привести к большей численной погрешности. Значение по умолчанию: False.
-
ROWS_GUARANTEED_SAFE: NotRequired[bool] -
Если True, гарантируется, что хотя бы одно значение в каждой строке не замаскировано. Это позволяет пропустить проверки безопасности и повысить производительность. Устанавливайте это значение только в том случае, если вы уверены, что ваша маска гарантирует данное свойство. Например, причинно-следственное внимание гарантированно безопасно, поскольку для каждого запроса есть как минимум один ключ-значение, на который можно обратить внимание. Значение по умолчанию: False.
-
USE_TMA: NotRequired[bool] -
Использовать ли Tensor Memory Accelerator (TMA) на поддерживаемом оборудовании. Это экспериментальная функция, которая может работать не на всем оборудовании; в настоящее время она предназначена только для графических процессоров NVIDIA Hopper и новее. Значение по умолчанию: False.
-
WRITE_DQ: NotRequired[bool] -
Определяет, выполняется ли scatter-операция градиентов в цикле итераций DQ обратного прохода. Если задать False, операция будет выполняться в цикле DK, что в зависимости от конкретных score_mod и mask_mod может быть быстрее. Значение по умолчанию: True.
-
kpack: NotRequired[int] -
Параметр упаковки ядра, специфичный для ROCm.
-
matrix_instr_nonkdim: NotRequired[int] -
Измерение матричной инструкции, отличное от K, специфичное для ROCm.
-
num_stages: NotRequired[int] -
Число этапов конвейера в ядре CUDA. Большее значение может повысить производительность, но увеличивает использование разделяемой памяти. Значение по умолчанию определяется автоматической настройкой.
-
num_warps: NotRequired[int] -
Число варпов, используемых в ядре CUDA. Большее значение может повысить производительность, но увеличивает нагрузку на регистры. Значение по умолчанию определяется автоматической настройкой.
-
waves_per_eu: NotRequired[int] -
Число волн на исполнительный блок, специфичное для ROCm.
BlockMask
-
class torch.nn.attention.flex_attention.BlockMask(seq_lengths, kv_num_blocks, kv_indices, full_kv_num_blocks, full_kv_indices, q_num_blocks, q_indices, full_q_num_blocks, full_q_indices, BLOCK_SIZE=(128, 128), mask_mod=<function noop_mask>, *, dq_write_order=None, dq_write_order_full=None, dq_kv_order=None, dq_kv_order_spt=None)[исходный код] -
BlockMask — это наш формат для представления блочной разреженной маски внимания. Он представляет собой нечто среднее между BCSR и разреженным форматом.
Основы
Блочно-разреженная маска означает, что вместо представления разреженности отдельных элементов маски блок размером KV_BLOCK_SIZE x Q_BLOCK_SIZE считается разреженным, только если разрежен каждый элемент внутри этого блока. Это хорошо соответствует возможностям оборудования, которое обычно рассчитано на выполнение последовательных операций загрузки и вычислений.
Этот формат в первую очередь оптимизирован для 1. простоты и 2. эффективности ядра. Примечательно, что он не оптимизирован по размеру, поскольку размер этой маски всегда сокращается в KV_BLOCK_SIZE * Q_BLOCK_SIZE раз. Если размер имеет значение, тензоры можно уменьшить, увеличив размер блока.
Основные составляющие нашего формата:
num_blocks_in_row: Tensor[ROWS]: описывает количество блоков в каждой строке.
col_indices: Tensor[ROWS, MAX_BLOCKS_IN_COL]:
col_indices[i]— это последовательность позиций блоков для строки i. Значения этой строки послеcol_indices[i][num_blocks_in_row[i]]не определены.Например, чтобы восстановить исходный тензор из этого формата:
dense_mask = torch.zeros(ROWS, COLS) for row in range(ROWS): for block_idx in range(num_blocks_in_row[row]): dense_mask[row, col_indices[row, block_idx]] = 1Примечательно, что этот формат упрощает реализацию редукции вдоль строк маски.
Подробности
Для базового варианта нашего формата достаточно kv_num_blocks и kv_indices. Основная блочно-разреженная структура представлена не более чем 4 парами тензоров:
1. (kv_num_blocks, kv_indices): используются на прямом проходе механизма внимания, поскольку редукция выполняется вдоль измерения KV.
2. [НЕОБЯЗАТЕЛЬНО] (full_kv_num_blocks, full_kv_indices): необязательная составляющая, предназначенная исключительно для оптимизации. Оказывается, применение маскирования к каждому блоку обходится довольно дорого! Если нам точно известно, какие блоки являются «полными» и вообще не требуют маскирования, то можно пропустить применение mask_mod к этим блокам. Для этого пользователю необходимо выделить отдельный mask_mod из score_mod. Для причинных масок это ускоряет вычисления примерно на 15%.
3. [ГЕНЕРИРУЕТСЯ] (q_num_blocks, q_indices): необходимы для обратного прохода, поскольку для вычисления dKV требуется проходить по маске вдоль измерения Q. Генерируются автоматически на основе пункта 1.
4. [ГЕНЕРИРУЕТСЯ] (full_q_num_blocks, full_q_indices): аналогично предыдущему пункту, но для обратного прохода. Генерируются автоматически на основе пункта 2.
Дополнительные необязательные тензоры могут содержать детерминированные метаданные dQ для обратного прохода блочно-разреженного FLASH:
5. [НЕОБЯЗАТЕЛЬНО] dq_write_order: метаданные порядка записи для частичных блоков. Формируются функцией create_block_mask при compute_dq_write_order=True или напрямую передаются в BlockMask.from_kv_blocks вызывающими сторонами, которые предварительно вычислили эти метаданные.
6. [НЕОБЯЗАТЕЛЬНО] dq_write_order_full: метаданные порядка записи для полных блоков, формируемые или передаваемые тем же способом, что и dq_write_order.
7. [НЕОБЯЗАТЕЛЬНО] dq_kv_order: явный порядок планировщика KV, используемый для формирования метаданных порядка записи. В настоящее время create_block_mask принимает dq_kv_order типа boolean; BlockMask.from_kv_blocks также принимает тензор для вызывающих сторон, которые напрямую передают предварительно вычисленные метаданные порядка записи.
-
BLOCK_SIZE: tuple[int, int]
-
dq_kv_order: Tensor | None
-
dq_kv_order_spt: bool | None
-
dq_write_order: Tensor | None
-
dq_write_order_full: Tensor | None
-
classmethod from_kv_blocks(kv_num_blocks, kv_indices, full_kv_num_blocks=None, full_kv_indices=None, BLOCK_SIZE=128, mask_mod=None, seq_lengths=None, compute_q_blocks=True, *, dq_write_order=None, dq_write_order_full=None, dq_kv_order=None)[исходный код] -
Создаёт экземпляр BlockMask на основе сведений о блоках ключей и значений.
- Параметры:
-
- kv_num_blocks (Tensor) – Количество блоков kv в каждой строковой плитке Q_BLOCK_SIZE.
- kv_indices (Tensor) – Индексы блоков ключей и значений в каждой строковой плитке Q_BLOCK_SIZE.
- full_kv_num_blocks (Optional[Tensor]) – Количество полных блоков kv в каждой строковой плитке Q_BLOCK_SIZE.
- full_kv_indices (Optional[Tensor]) – Индексы полных блоков ключей и значений в каждой строковой плитке Q_BLOCK_SIZE.
- BLOCK_SIZE (Union[int, tuple[int, int]]) – Размер плиток KV_BLOCK_SIZE x Q_BLOCK_SIZE.
- mask_mod (Optional[Callable]) – Функция для изменения маски.
- dq_write_order (Optional[Tensor]) – Предварительно вычисленные детерминированные метаданные порядка записи dQ.
- dq_write_order_full (Optional[Tensor]) – Предварительно вычисленные детерминированные метаданные порядка записи dQ для полных блоков.
- dq_kv_order (Optional[Union[Tensor, bool]]) – Порядок планировщика столбцов KV, используемый для формирования dq_write_order. Значение bool выбирает встроенный порядок; тензор задаёт явную перестановку рангов планировщика в n-блоки.
- Возвращает:
-
Экземпляр с полной информацией Q, сформированной с помощью _transposed_ordered
- Тип возвращаемого значения:
- Вызывает исключения:
-
- RuntimeError – Если kv_indices содержит менее 2 измерений.
- AssertionError – Если передан только один из аргументов full_kv_*.
-
full_kv_indices: Tensor | None
-
full_kv_num_blocks: Tensor | None
-
full_q_indices: Tensor | None
-
full_q_num_blocks: Tensor | None
-
kv_indices: Tensor
-
kv_num_blocks: Tensor
-
mask_mod: Callable[[Tensor, Tensor, Tensor, Tensor], Tensor]
-
numel()[исходный код] -
Возвращает количество элементов маски (без учёта разреженности).
- Тип возвращаемого значения:
-
q_indices: Tensor | None
-
q_num_blocks: Tensor | None
-
seq_lengths: tuple[int, int]
-
property shape: tuple[int, ...]
-
sparsity()[исходный код] -
Вычисляет процент разреженных блоков (то есть блоков, которые не вычисляются).
- Тип возвращаемого значения:
-
to(device)[исходный код] -
Перемещает BlockMask на указанное устройство.
- Параметры:
-
device (torch.device or str) – Целевое устройство, на которое нужно переместить BlockMask. Это может быть объект torch.device или строка (например, ‘cpu’, ‘cuda:0’).
- Возвращает:
-
Новый экземпляр BlockMask, все тензорные компоненты которого перемещены на указанное устройство.
- Тип возвращаемого значения:
Примечание
Этот метод не изменяет исходный BlockMask на месте. Вместо этого он возвращает новый экземпляр BlockMask, отдельные атрибуты-тензоры которого могут быть перемещены на указанное устройство или остаться на прежнем — в зависимости от того, на каком устройстве они находятся в данный момент.
-
to_dense()[исходный код] -
Возвращает плотный блок, эквивалентный блочной маске.
- Тип возвращаемого значения:
-
to_string(grid_size=(20, 20), limit=4)[исходный код] -
Возвращает строковое представление блочной маски. Очень удобно.
Если grid_size равен -1, выводится несжатая версия. Внимание: она может быть очень большой!
- Тип возвращаемого значения:
-
-
BlockMask.as_tuple(flatten: Literal[True] = True) → tuple[int, int, Tensor, Tensor, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, bool | None, int, int, Callable[[Tensor, Tensor, Tensor, Tensor], Tensor]][исходный код] - BlockMask.as_tuple(flatten:Literal[False]) tuple[tuple[int,int],Tensor,Tensor,Tensor|None,Tensor|None,Tensor|None,Tensor|None,Tensor|None,Tensor|None,Tensor|None,Tensor|None,bool|None,tuple[int,int],Callable[[Tensor,Tensor,Tensor,Tensor],Tensor]]
-
Возвращает кортеж атрибутов BlockMask.
- Параметры:
-
flatten (bool) – Если значение равно True, кортеж (KV_BLOCK_SIZE, Q_BLOCK_SIZE) будет развёрнут.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/nn.attention.flex_attention.html