f"length, got {local_q_len}; pad the global sequence to a "
f"multiple of 2 * cp_size ({2 * cp_mesh.size()})"
)
peer_rank = _head_tail_peer_rank(cp_mesh)
query_peer = _collectives.p2p_exchange(
query.narrow(2, local_q_len // 2, local_q_len // 2), peer_rank)
global_key, global_value = flex_cp_allgather(key, value, 2, cp_mesh)
keep_output = _run_head_tail_half(
attention_fn, query.narrow(2, 0, local_q_len // 2),
global_key, global_value, attention_kwargs)
peer_output = _run_head_tail_half(
attention_fn, query_peer, global_key, global_value,
attention_kwargs if peer_attention_kwargs is None else peer_attention_kwargs,
)
return torch.cat(
[keep_output, _collectives.p2p_exchange(peer_output, peer_rank)], dim=2)
def _cp_offset_causal_mask(q_len: int, kv_len: int, lo: int,