@@ -194,6 +194,7 @@ def _apply_mask_and_soft_cap(
194194 mask_value : float ,
195195 mask_ref ,
196196 q_sequence_ref ,
197+ kv_sequence_ref ,
197198 q_segment_ids_ref ,
198199 kv_segment_ids_ref ,
199200 * ,
@@ -207,30 +208,38 @@ def _apply_mask_and_soft_cap(
207208) -> tuple [jax .Array , jax .Array | None ]:
208209 assert mask_ref is None or q_sequence_ref is None
209210 assert (q_sequence_ref is None ) == (mask_function is None )
211+ assert kv_sequence_ref is None or mask_function is not None
210212
211213 masks = []
212214 if has_partial_mask :
213215 if mask_ref is not None :
214216 mask = mask_ref [:, k_slice ] if k_in_lanes else mask_ref [k_slice , :]
215217 masks .append (mask )
216218 elif mask_function is not None :
217- # Compute the mask using the given q_sequence indices.
218- # KV indices are computed on the fly. This works because we only support Q
219- # sequence sharding. If we wanted to compute Q indices too, then we would
220- # need to keep into account the current shard along Q sequence.
219+ # Compute the mask using original Q positions. K/V positions are computed
220+ # on the fly unless explicit original K/V positions are provided.
221221
222222 if k_in_lanes :
223223 assert q_sequence_ref .shape == (bq , NUM_LANES )
224224
225- k_sequence = k_offset + jax .lax .broadcasted_iota (jnp .int32 , (bq , k_slice .size ), 1 )
225+ if kv_sequence_ref is None :
226+ k_sequence = k_offset + jax .lax .broadcasted_iota (jnp .int32 , (bq , k_slice .size ), 1 )
227+ else :
228+ k_sequence = jnp .broadcast_to (kv_sequence_ref [:1 , k_slice ], (bq , k_slice .size ))
226229
227230 repeats , rem = divmod (k_slice .size , NUM_LANES )
228231 assert rem == 0
229232 q_sequence = jnp .tile (q_sequence_ref [...], (1 , repeats )) # [bq, k_slice.size]
230233 else :
231234 assert q_sequence_ref .shape == (NUM_SUBLANES , bq )
232235
233- k_sequence = k_offset + jax .lax .broadcasted_iota (jnp .int32 , (k_slice .size , bq ), 0 )
236+ if kv_sequence_ref is None :
237+ k_sequence = k_offset + jax .lax .broadcasted_iota (jnp .int32 , (k_slice .size , bq ), 0 )
238+ else :
239+ repeats , rem = divmod (bq , NUM_LANES )
240+ if rem :
241+ raise NotImplementedError (f"block_q must be a multiple of { NUM_LANES } " )
242+ k_sequence = jnp .tile (kv_sequence_ref [k_slice , :], (1 , repeats ))
234243 q_sequence = q_sequence_ref [:1 , :] # [1, bq]
235244 q_sequence = jnp .broadcast_to (q_sequence , (k_slice .size , bq ))
236245
@@ -290,6 +299,7 @@ def flash_attention_kernel(
290299 sinks_ref ,
291300 mask_ref ,
292301 q_sequence_ref ,
302+ kv_sequence_ref ,
293303 max_logit_value_ref ,
294304 # Outputs
295305 o_ref ,
@@ -394,6 +404,7 @@ def body(kv_compute_index, _, has_partial_mask=False):
394404 mask_value ,
395405 mask_ref ,
396406 q_sequence_ref ,
407+ kv_sequence_ref ,
397408 q_segment_ids_ref ,
398409 kv_segment_ids_ref ,
399410 attn_logits_soft_cap = attn_logits_soft_cap ,
@@ -669,6 +680,13 @@ def mask_index_map(h, grid_idx, rows_ref, cols_ref, mask_next_ref=None, *_):
669680 q_sequence = None
670681 in_specs .append (None )
671682
683+ if mask_info .kv_sequence is not None :
684+ kv_sequence = jax .lax .broadcast_in_dim (mask_info .kv_sequence , (NUM_SUBLANES , kv_seq_len ), (1 ,))
685+ in_specs .append (pl .BlockSpec ((NUM_SUBLANES , bkv ), kv_segment_ids_index_map ))
686+ else :
687+ kv_sequence = None
688+ in_specs .append (None )
689+
672690 if max_logit_value is not None :
673691 # reshape to allow sublane selection for vmap-ping and shard_map-ping
674692 max_logit_value = jnp .broadcast_to (
@@ -819,6 +837,7 @@ def _fwd_cost_estimate(
819837 sinks ,
820838 mask_info .partial_mask_blocks ,
821839 q_sequence ,
840+ kv_sequence ,
822841 max_logit_value ,
823842 )
824843 out , logsumexp , l_linear , max_logits = all_out
@@ -996,6 +1015,7 @@ def _flash_attention_dq_kernel(
9961015 di_ref ,
9971016 mask_ref ,
9981017 q_sequence_ref ,
1018+ kv_sequence_ref ,
9991019 # Outputs
10001020 dq_scratch_ref ,
10011021 dq_ref ,
@@ -1049,6 +1069,7 @@ def body(has_partial_mask: bool = False):
10491069 mask_value ,
10501070 mask_ref ,
10511071 q_sequence_ref ,
1072+ kv_sequence_ref ,
10521073 q_segment_ids_ref ,
10531074 kv_segment_ids_ref ,
10541075 attn_logits_soft_cap = attn_logits_soft_cap ,
@@ -1115,6 +1136,7 @@ def _flash_attention_dkv_kernel(
11151136 di_ref ,
11161137 mask_ref ,
11171138 q_sequence_ref ,
1139+ kv_sequence_ref ,
11181140 # aliases
11191141 dq_alias ,
11201142 dk_alias ,
@@ -1212,6 +1234,7 @@ def _load_kv(ref, layout):
12121234 mask_value ,
12131235 mask_ref ,
12141236 q_sequence_ref ,
1237+ kv_sequence_ref ,
12151238 q_segment_ids_ref ,
12161239 kv_segment_ids_ref ,
12171240 attn_logits_soft_cap = attn_logits_soft_cap ,
@@ -1436,9 +1459,8 @@ def create_dkv_index_map(h, i, j, *_):
14361459 mask_spec = pl .BlockSpec ((None , bkv , bq ), mask_index_map )
14371460
14381461 q_segment_ids_index_map = unravel (lambda h , i , j : (0 , i ))
1462+ kv_segment_ids_index_map = unravel (lambda h , i , j : (j , 0 ))
14391463 if segment_ids is not None :
1440- kv_segment_ids_index_map = unravel (lambda h , i , j : (j , 0 ))
1441-
14421464 q_segment_spec = pl .BlockSpec ((NUM_SUBLANES , bq ), q_segment_ids_index_map )
14431465 kv_segment_spec = pl .BlockSpec ((bkv , NUM_LANES ), kv_segment_ids_index_map )
14441466 q_segment_ids = jax .lax .broadcast_in_dim (segment_ids .q , (NUM_SUBLANES , q_seq_len ), (1 ,))
@@ -1485,6 +1507,13 @@ def create_dkv_index_map(h, i, j, *_):
14851507 q_sequence = None
14861508 in_specs .append (None )
14871509
1510+ if mask_info .kv_sequence is not None :
1511+ in_specs .append (pl .BlockSpec ((bkv , NUM_LANES ), kv_segment_ids_index_map ))
1512+ kv_sequence = jax .lax .broadcast_in_dim (mask_info .kv_sequence , (kv_seq_len , NUM_LANES ), (0 ,))
1513+ else :
1514+ kv_sequence = None
1515+ in_specs .append (None )
1516+
14881517 dq_reduction_steps = config .dq_reduction_steps
14891518 if not dynamic_grid and kv_steps <= 3 and dq_reduction_steps == 3 :
14901519 dq_reduction_steps = None
@@ -1581,6 +1610,7 @@ def create_dkv_index_map(h, i, j, *_):
15811610 di ,
15821611 mask_info .partial_mask_blocks ,
15831612 q_sequence ,
1613+ kv_sequence ,
15841614 ]
15851615 num_args = sum (1 for x in args if x is not None )
15861616 input_output_aliases = {}
@@ -1609,6 +1639,7 @@ def _bwd_cost_estimate(
16091639 di : jax .Array ,
16101640 partial_mask_blocks : jax .Array | None ,
16111641 q_sequence : jax .Array | None ,
1642+ kv_sequence : jax .Array | None ,
16121643 out_shapes : list [jax .ShapeDtypeStruct ],
16131644 mask_sparsity_factor : float ,
16141645 ) -> pl .CostEstimate :
@@ -1643,6 +1674,7 @@ def _bwd_cost_estimate(
16431674 di ,
16441675 partial_mask_blocks ,
16451676 q_sequence ,
1677+ kv_sequence ,
16461678 ]
16471679 input_bytes = sum (map (_bytes , inputs_ ))
16481680 output_bytes = sum (map (_bytes , out_shapes ))
@@ -1666,6 +1698,7 @@ def _bwd_cost_estimate(
16661698 di ,
16671699 mask_info .partial_mask_blocks ,
16681700 q_sequence ,
1701+ kv_sequence ,
16691702 out_shapes ,
16701703 dkv_mask_sparsity ,
16711704 )
@@ -1880,6 +1913,7 @@ def mask_info_spec(mask_info):
18801913 if mask_info .partial_mask_blocks is not None
18811914 else None ,
18821915 q_sequence = _resolve_spec (mask_info .q_sequence ),
1916+ kv_sequence = (jax .sharding .PartitionSpec () if mask_info .kv_sequence is not None else None ),
18831917 )
18841918
18851919 return SplashAttentionKernel (
@@ -2029,6 +2063,7 @@ def process_mask_shard(mask):
20292063 block_mask = mask_spec ,
20302064 partial_mask_blocks = mask_spec ,
20312065 q_sequence = None ,
2066+ kv_sequence = None ,
20322067 )
20332068 out_specs = (
20342069 mask_info_specs ,
0 commit comments