Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
86 changes: 86 additions & 0 deletions tests/v1/attention/test_mla_backends.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,8 @@
MLAAttention,
QueryLenSupport,
_DecodeConcatQuantFP8,
_materialize_kv_b_proj_weight,
_release_b12x_mxfp8_kv_b_proj,
)
from vllm.model_executor.layers.attention.sparse_mla_attention import (
SparseMLACommonImpl,
Expand Down Expand Up @@ -124,6 +126,7 @@ def test_mla_post_load_preserves_runtime_weight_addresses(monkeypatch):
layer.is_aiter_triton_fp8_bmm_enabled = False
layer.quant_config = None
layer.layer_name = "test"
layer.impl = SimpleNamespace(can_release_kv_b_proj_after_loading=False)

monkeypatch.setattr(
mla_attention_module, "set_default_quant_scales", lambda *_, **__: None
Expand All @@ -147,6 +150,89 @@ def test_mla_post_load_preserves_runtime_weight_addresses(monkeypatch):
torch.testing.assert_close(layer.W_UK_T, old_w_uk_t + 100)


class _PackedLinearMethod:
def __init__(self, weight: torch.Tensor):
self.weight = weight

def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
assert layer.b12x_mxfp8_packed_weight is not None
assert bias is None
return x @ self.weight.T


@pytest.mark.cpu_test
def test_b12x_mxfp8_mla_post_load_releases_absorbed_sources(monkeypatch):
weight = torch.arange(28.0, dtype=torch.float32).reshape(14, 2)
layer = MLAAttention.__new__(MLAAttention)
torch.nn.Module.__init__(layer)
layer.kv_lora_rank = 2
layer.num_heads = 2
layer.qk_nope_head_dim = 3
layer.v_head_dim = 4
layer.kv_b_proj = torch.nn.Module()
layer.kv_b_proj.register_parameter(
"weight", torch.nn.Parameter(weight, requires_grad=False)
)
layer.kv_b_proj.register_parameter(
"weight_scale", torch.nn.Parameter(torch.ones((14, 1)), requires_grad=False)
)
layer.kv_b_proj.input_size_per_partition = 2
layer.kv_b_proj.quant_method = _PackedLinearMethod(weight)
layer.kv_b_proj.b12x_mxfp8_packed_weight = object()
layer.is_aiter_triton_fp4_bmm_enabled = False
layer.is_aiter_triton_fp8_bmm_enabled = False
layer.quant_config = None
layer.layer_name = "test"
layer.impl = SimpleNamespace(can_release_kv_b_proj_after_loading=True)
layer.prefill_backend = None

monkeypatch.setattr(
mla_attention_module, "set_default_quant_scales", lambda *_, **__: None
)

with torch.no_grad():
layer.process_weights_after_loading(torch.float32)

assert not hasattr(layer.kv_b_proj, "weight")
assert not hasattr(layer.kv_b_proj, "weight_scale")
assert layer.kv_b_proj.b12x_mxfp8_packed_weight is None


@pytest.mark.cpu_test
def test_release_kv_b_proj_is_inert_without_b12x_mxfp8_owner():
weight = torch.nn.Parameter(torch.ones((4, 3)), requires_grad=False)
layer = torch.nn.Module()
layer.register_parameter("weight", weight)

released = _release_b12x_mxfp8_kv_b_proj(layer)

assert released is False
assert layer.weight is weight


@pytest.mark.cpu_test
def test_reloads_absorbed_weight_from_regenerated_packed_owner():
expected = torch.arange(12, dtype=torch.float32).view(4, 3)
layer = SimpleNamespace(
input_size_per_partition=3,
quant_method=_PackedLinearMethod(expected),
b12x_mxfp8_packed_weight=object(),
)

materialized = _materialize_kv_b_proj_weight(
layer,
out_dtype=torch.float32,
fallback_device=torch.device("cpu"),
)

torch.testing.assert_close(materialized, expected)


# Filtered per-test via validate_configuration (capability/deps/dims).
PREFILL_BACKENDS_TO_TEST = [
MLAPrefillBackendEnum.FLASH_ATTN,
Expand Down
55 changes: 53 additions & 2 deletions vllm/model_executor/layers/attention/mla_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -292,6 +292,46 @@

_FP8_DTYPE = current_platform.fp8_dtype()

_KV_B_PROJ_SOURCE_PARAMETERS = ("weight", "weight_scale")


def _materialize_kv_b_proj_weight(
layer: torch.nn.Module,
*,
out_dtype: torch.dtype,
fallback_device: torch.device | None,
) -> torch.Tensor:
"""Return kv_b_proj weights, including after source storage was released."""
source_names = ("weight", "qweight", "weight_packed")
if any(
isinstance(getattr(layer, name, None), torch.Tensor) for name in source_names
):
return get_and_maybe_dequant_weights(layer, out_dtype=out_dtype)

packed = getattr(layer, "b12x_mxfp8_packed_weight", None)
quant_method = getattr(layer, "quant_method", None)
if packed is None or quant_method is None or fallback_device is None:
return get_and_maybe_dequant_weights(layer, out_dtype=out_dtype)

identity = torch.eye(
layer.input_size_per_partition,
dtype=out_dtype,
device=fallback_device,
)
return quant_method.apply(layer, identity, bias=None).to(out_dtype).T


def _release_b12x_mxfp8_kv_b_proj(layer: torch.nn.Module) -> bool:
"""Drop source tensors after B12X has produced the absorbed MLA pair."""
if getattr(layer, "b12x_mxfp8_packed_weight", None) is None:
return False

for name in _KV_B_PROJ_SOURCE_PARAMETERS:
if hasattr(layer, name):
delattr(layer, name)
layer.b12x_mxfp8_packed_weight = None
return True


def _can_use_b12x_dcp_prefill_workspace(
*,
Expand Down Expand Up @@ -1281,8 +1321,11 @@ def process_weights_after_loading(self, act_dtype: torch.dtype):
# we currently do not have quantized bmm's which are needed for
# `W_UV` and `W_UK_T`, we just store fp16/bf16 copies and perform
# the bmm's in 16-bit, the extra memory overhead of this is fairly low
kv_b_proj_weight = get_and_maybe_dequant_weights(
self.kv_b_proj, out_dtype=act_dtype
fallback_device = self.W_UV.device if hasattr(self, "W_UV") else None
kv_b_proj_weight = _materialize_kv_b_proj_weight(
self.kv_b_proj,
out_dtype=act_dtype,
fallback_device=fallback_device,
).T

assert kv_b_proj_weight.shape == (
Expand Down Expand Up @@ -1366,6 +1409,14 @@ def process_weights_after_loading(self, act_dtype: torch.dtype):
# Convert from (L, N, P) to (N, P, L)
replace_parameter(self, "W_UK_T", W_UK.permute(1, 2, 0), prefer_copy=True)

if self.impl.can_release_kv_b_proj_after_loading:
if self.prefill_backend is not None:
raise RuntimeError(
"An MLA backend cannot release kv_b_proj while MHA prefill "
"is enabled."
)
_release_b12x_mxfp8_kv_b_proj(self.kv_b_proj)

# If we should not load quant weights, we initialize the scales to 1.0
# as the default value. See [Note: Register q/k/v/prob scales in state dict]
# for more details.
Expand Down
2 changes: 2 additions & 0 deletions vllm/v1/attention/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -1234,6 +1234,8 @@ def do_rope_and_kv_cache_update(
class MLAAttentionImpl(AttentionImplBase[T], Generic[T]):
"""MLA attention implementation with forward_mqa and forward_mha methods."""

can_release_kv_b_proj_after_loading: bool = False

@abstractmethod
def __init__(
self,
Expand Down
1 change: 1 addition & 0 deletions vllm/v1/attention/backends/mla/b12x_mla_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -905,6 +905,7 @@ class B12xMLASparseImpl(MLAAttentionImpl[B12xMLASparseMetadata]):
# B12X handles decode and extend inside its own top-k MQA kernels; the
# generic dense-MHA prefill path assumes cache layouts it never validated.
supports_mha_prefill: bool = False
can_release_kv_b_proj_after_loading: bool = True
supports_dcp_project_before_merge: bool = True
supports_dcp_gather_query_in_workspace: bool = True
supports_dcp_project_before_merge_in_workspace: bool = True
Expand Down
Loading