Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 52 additions & 0 deletions tests/distributed/test_dcp_a2a.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ class _FakeCPGroup:
def __init__(self, world_size: int, device_group: dist.ProcessGroup):
self.world_size = world_size
self.device_group = device_group
self.rank_in_group = dist.get_rank(device_group)


def _dtype_from_name(dtype_name: str) -> torch.dtype:
Expand Down Expand Up @@ -375,6 +376,57 @@ def test_pack_unpack_combine_matches_reference(
else:
_assert_packed_a2a_close(actual, expected_out, dtype)

@pytest.mark.skipif(
torch.accelerator.device_count() < 1, reason="CUDA is required."
)
def test_pack_send_zeroes_empty_local_rows(self):
from vllm.v1.attention.ops.dcp_alltoall import (
_dcp_a2a_lse_pack_dim,
_dcp_a2a_pack_send,
)

device = torch.device("cuda")
world_size, B, h_per_rank, D = 4, 5, 2, 32
H = world_size * h_per_rank
cp_attn_out = torch.randn(B, H, D, device=device)
cp_attn_lse = torch.randn(B, H, device=device)
valid_counts = torch.tensor([3, 0, 1, 0, 2], device=device)
lse_pack_dim = _dcp_a2a_lse_pack_dim(cp_attn_out.dtype)
send_buffer = torch.empty(
(world_size, B, h_per_rank, D + lse_pack_dim),
device=device,
)

_dcp_a2a_pack_send(
cp_attn_out,
cp_attn_lse,
send_buffer,
world_size,
h_per_rank,
D,
lse_pack_dim,
valid_counts=valid_counts,
)
torch.accelerator.synchronize()

empty_rows = valid_counts == 0
non_empty_rows = ~empty_rows
expected_out = (
cp_attn_out.view(B, world_size, h_per_rank, D)
.permute(1, 0, 2, 3)
.contiguous()
)
empty_payload = send_buffer[:, empty_rows, :, :D]
torch.testing.assert_close(empty_payload, torch.zeros_like(empty_payload))
torch.testing.assert_close(
send_buffer[:, non_empty_rows, :, :D],
expected_out[:, non_empty_rows],
)
torch.testing.assert_close(
send_buffer[:, empty_rows, :, D],
torch.full_like(send_buffer[:, empty_rows, :, D], float("-inf")),
)


def _distributed_packed_a2a_worker(env: dict[str, str]) -> None:
update_environment_variables(env)
Expand Down
88 changes: 69 additions & 19 deletions vllm/model_executor/layers/attention/mla_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -513,6 +513,10 @@ def __init__(
)

# Initialize q/k/v range constants.
# Project attention output through W_UV before the DCP merge: shrinks
# the merge payload from kv_lora_rank to v_head_dim per head.
self.W_UV_dcp: torch.Tensor | None = None

self.q_range = torch.tensor(envs.Q_SCALE_CONSTANT, dtype=torch.float32)
self.k_range = torch.tensor(envs.K_SCALE_CONSTANT, dtype=torch.float32)
self.v_range = torch.tensor(envs.V_SCALE_CONSTANT, dtype=torch.float32)
Expand Down Expand Up @@ -811,25 +815,7 @@ def forward_impl(
assert attn_metadata.decode is not None
attn_out, lse = self.impl.forward_mqa(mqa_q, kv_cache, attn_metadata, self) # type: ignore[attr-defined]

# correct dcp attn_out with lse.
if self.impl.dcp_world_size > 1:
if self.dcp_a2a:
attn_out = dcp_a2a_lse_reduce(
attn_out,
lse,
get_dcp_group(),
is_lse_base_on_e=self.impl.lse_base_on_e,
)
else:
attn_out = cp_lse_ag_out_rs(
attn_out,
lse,
get_dcp_group(),
is_lse_base_on_e=self.impl.lse_base_on_e,
)

# v_up projection
self._v_up_proj(attn_out, out=mqa_output_slice)
self._dcp_merge_and_v_up_proj(attn_out, lse, mqa_output_slice)

if quant_key is not None:
quant_idx = num_mqa_tokens if mha_use_quant_output else num_actual_toks
Expand Down Expand Up @@ -958,6 +944,11 @@ def process_weights_after_loading(self, act_dtype: torch.dtype):
else:
# Convert from (L, N, V) to (N, L, V)
self.W_UV = W_UV.transpose(0, 1)
if getattr(self.impl, "dcp_world_size", 1) > 1:
# all_gather_into_tensor requires a contiguous input
self.W_UV_dcp = get_dcp_group().all_gather(
self.W_UV.contiguous(), dim=0
)
# Convert from (L, N, P) to (N, P, L)
self.W_UK_T = W_UK.permute(1, 2, 0)

Expand Down Expand Up @@ -1009,6 +1000,57 @@ def get_kv_cache_spec(self, vllm_config: VllmConfig) -> KVCacheSpec:
cache_dtype_str=vllm_config.cache_config.cache_dtype,
)

def _dcp_lse_merge(
self,
attn_out: torch.Tensor,
lse: torch.Tensor,
out: torch.Tensor | None,
) -> torch.Tensor:
"""LSE-weighted combine of the per-rank attention outputs across the
DCP group. The a2a transport can write directly into ``out``."""
if self.dcp_a2a:
return dcp_a2a_lse_reduce(
attn_out,
lse,
get_dcp_group(),
is_lse_base_on_e=self.impl.lse_base_on_e,
out=out,
valid_counts=getattr(self.impl, "_last_dcp_valid_counts", None),
)
return cp_lse_ag_out_rs(
attn_out,
lse,
get_dcp_group(),
is_lse_base_on_e=self.impl.lse_base_on_e,
)

def _dcp_merge_and_v_up_proj(
self,
attn_out: torch.Tensor,
lse: torch.Tensor,
out: torch.Tensor,
) -> None:
"""Combine the decode attention output across DCP ranks (if any) and
apply the W_UV up-projection into ``out`` (flattened v_head_dim)."""
if self.impl.dcp_world_size > 1 and self.W_UV_dcp is not None:
# Project kv_lora_rank -> v_head_dim BEFORE the merge to halve the
# DCP exchange payload; the LSE-weighted merge commutes with the
# linear W_UV projection, so this is exact.
projected = attn_out.new_empty(
attn_out.shape[0], attn_out.shape[1], self.v_head_dim
)
self._v_up_proj_bmm(attn_out, projected, self.W_UV_dcp)
out_view = out.view(-1, self.num_heads, self.v_head_dim)
merged = self._dcp_lse_merge(projected, lse, out=out_view)
if merged is not out_view:
out.copy_(merged.reshape(out.shape))
return

# Project after the merge (dcp=1, or backends without a gathered W_UV).
if self.impl.dcp_world_size > 1:
attn_out = self._dcp_lse_merge(attn_out, lse, out=None)
self._v_up_proj(attn_out, out=out)

def _v_up_proj(self, x: torch.Tensor, out: torch.Tensor):
# Convert from (B, N, L) to (N, B, L)
x = x.view(-1, self.num_heads, self.kv_lora_rank).transpose(0, 1)
Expand All @@ -1033,6 +1075,14 @@ def _v_up_proj(self, x: torch.Tensor, out: torch.Tensor):
# Multiply + Transpose (N, B, L) x (N, L, V)->(N, B, V)->(B, N, V)
torch.bmm(x, self.W_UV, out=out.transpose(0, 1))

def _v_up_proj_bmm(
self, x: torch.Tensor, out: torch.Tensor, w_uv: torch.Tensor
) -> None:
num_heads = w_uv.shape[0]
x = x.view(-1, num_heads, self.kv_lora_rank).transpose(0, 1)
out = out.view(-1, num_heads, self.v_head_dim)
torch.bmm(x, w_uv, out=out.transpose(0, 1))


def unified_mla_kv_cache_update(
kv_c_normed: torch.Tensor,
Expand Down
Loading
Loading