"""Pack unequal-width QKV into one sequence-to-head exchange."""
key, value = _key_value(module, state.kv_latent, key_rope, self.plan.local_heads, None)
if self.plan.cp_degree == 1:
return query, key, value
widths = (query.shape[-1], key.shape[-1], value.shape[-1])
# Keep projection outputs in BSND while packing to reduce layout copies.
# Source-rank/sequence reconstruction is also a view when B=1.
payload = torch.cat(tuple(tensor.transpose(1, 2) for tensor in (query, key, value)), dim=-1)
payload = ulysses_seq_to_head(payload, 1, 2, self.cp_mesh)
return tuple(tensor.transpose(1, 2) for tensor in payload.split(widths, dim=-1))
def _latent(self, module, state, query, key_rope):
"""Exchange Q heads, gather shared KV latents, then expand only owned heads."""
dimensions = self.plan.dimensions
heads = self.plan.compute_heads
rank = 0 if self.cp_mesh is None else self.cp_mesh.get_local_rank()
start = rank * heads
if self.plan.cp_degree > 1:
query = ulysses_seq_to_head(query, 2, 1, self.cp_mesh)
payload = mla_all_gather(torch.cat((state.kv_latent, key_rope), dim=-1), 1, self.cp_mesh)
kv_latent, key_rope = payload.split((dimensions.kv_rank, dimensions.rope_dim), dim=-1)
# A contiguous latent preserves the native Linear bias-fusion path after payload splitting.
key, value = _key_value(module, kv_latent.contiguous(), key_rope, heads, (start, start + heads))
return query, key, value