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.
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 tinyq_ranges.Current pattern
For a sequence layout like interleaved token/latent pairs:
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:
This preserves semantics, but creates
O(num_pairs)ranges with very smallq_lenand growingk_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:
Semantically equivalent to:
So for
chunk_size=2:Why existing mask types are not enough
Existing
attn_type_mapvalues appear to support:CAUSALis not equivalent, because within each pair/chunk, earlier q positions cannot attend to later k positions. For example,token_icannot attend tolatent_i.BI_CAUSALis also not equivalent, because it represents an aligned band/intersection rather than prefix attention plus chunk-local full attention.auto_range_mergealso does not solve this case, because the current smallq_rangesare different ranges, not repeated identical outer ranges.Proposed API
One possible API extension:
Where:
Alternatively, a scalar
chunk_sizecould be accepted if all chunk-causal ranges use the same chunk size.Expected benefit
This would reduce range count from
O(num_pairs)toO(num_sequences)orO(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.