Skip to content

Feature request: chunk-causal mask type for Flex Flash Attention #316

Description

@xin-w8023

Feature request: chunk-causal mask type for Flex Flash Attention

Hi MagiAttention team,

We have a workload that currently needs to express many small full-attention ranges in flex_flash_attn_func, and it becomes very expensive in FA3 backward due to a large number of tiny q_ranges.

Current pattern

For a sequence layout like interleaved token/latent pairs:

prefix, token_0, latent_0, token_1, latent_1, ..., token_n, latent_n

We need each pair to be mutually visible, while each pair should also be able to attend to all previous context. Today we express this as many full-attention ranges:

q=[100,102), k=[0,102), attn_type=FULL
q=[102,104), k=[0,104), attn_type=FULL
...
q=[198,200), k=[0,200), attn_type=FULL

This preserves semantics, but creates O(num_pairs) ranges with very small q_len and growing k_len. In our profile, this causes Flex Flash / FA3 backward to dominate training time.

Desired representation

It would be useful to represent the same pattern compactly as:

q_range = [100, 200)
k_range = [0, 200)
chunk_size = 2
attn_type = CHUNK_CAUSAL

Semantically equivalent to:

for each q chunk [q_start + i*chunk_size, q_start + (i+1)*chunk_size):
    attend to k range [k_start, q_start + (i+1)*chunk_size)
    with full attention inside the current chunk

So for chunk_size=2:

q=[100,102), k=[0,102), FULL
q=[102,104), k=[0,104), FULL
...
q=[198,200), k=[0,200), FULL

Why existing mask types are not enough

Existing attn_type_map values appear to support:

0 FULL
1 CAUSAL
2 INV_CAUSAL
3 BI_CAUSAL

CAUSAL is not equivalent, because within each pair/chunk, earlier q positions cannot attend to later k positions. For example, token_i cannot attend to latent_i.

BI_CAUSAL is also not equivalent, because it represents an aligned band/intersection rather than prefix attention plus chunk-local full attention.

auto_range_merge also does not solve this case, because the current small q_ranges are different ranges, not repeated identical outer ranges.

Proposed API

One possible API extension:

flex_flash_attn_func(
    q,
    k,
    v,
    q_ranges,
    k_ranges,
    attn_type_map,
    chunk_size_map=None,  # optional int32 tensor, same length as q_ranges
)

Where:

attn_type_map[i] = CHUNK_CAUSAL
chunk_size_map[i] = 2  # or another chunk size

Alternatively, a scalar chunk_size could be accepted if all chunk-causal ranges use the same chunk size.

Expected benefit

This would reduce range count from O(num_pairs) to O(num_sequences) or O(num_large_regions), while preserving exact mask semantics. It should significantly reduce overhead and improve backward performance for interleaved autoregressive + pair-local full-attention workloads.

Thanks for considering this feature.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions