Harden DeepSeek V4 MHC warmup guards

- Rename local warmup destructured comb_mix to res_mix (same value)

- Strengthen CPU-device no-op test: spy mhc_pre_broadcast_tilelang, assert zero calls

- In _warmup_broadcast_mhc, fail closed/no-op unless hc_attn_fn_broadcast is a tensor on the same CUDA device as hc_attn_fn

- Add authentic tests for missing (non-tensor) and mismatched-device broadcast tensor

Signed-off-by: Colton Ottley <coltonottley@gmail.com>
This commit is contained in:
Colton Ottley
2026-07-24 05:23:11 -06:00
parent 32389a6b3c
commit 21ef130e16
2 changed files with 77 additions and 18 deletions
@@ -351,6 +351,28 @@ class TestFindDeepseekV4Model:
class TestWarmupBroadcastNoOpConditions:
@staticmethod
def _install_broadcast_spy(monkeypatch) -> list[str]:
calls: list[str] = []
module_path = "vllm.model_executor.kernels.mhc.tilelang"
stub = types.ModuleType(module_path)
stub.mhc_pre_broadcast_tilelang = lambda *args, **kwargs: calls.append(
"mhc_pre_broadcast_tilelang"
)
monkeypatch.setitem(sys.modules, module_path, stub)
return calls
@staticmethod
def _model_with_layer(layer):
layer.__name__ = layer.__qualname__ = "DeepseekV4DecoderLayer"
class MockModel:
def modules(self):
yield self
yield layer()
return MockModel()
def test_noop_when_no_broadcast_layer(self) -> None:
from vllm.model_executor.warmup.deepseek_v4_mhc_warmup import (
_warmup_broadcast_mhc,
@@ -359,34 +381,59 @@ class TestWarmupBroadcastNoOpConditions:
class NoBroadcastLayer:
hc_attn_fn_broadcast = None
NoBroadcastLayer.__name__ = NoBroadcastLayer.__qualname__ = (
"DeepseekV4DecoderLayer"
_warmup_broadcast_mhc(
self._model_with_layer(NoBroadcastLayer), token_sizes=[1, 2, 4]
)
class MockModel:
def modules(self):
yield self
yield NoBroadcastLayer()
_warmup_broadcast_mhc(MockModel(), token_sizes=[1, 2, 4])
def test_noop_when_device_not_cuda(self) -> None:
def test_noop_when_device_not_cuda(self, monkeypatch) -> None:
from vllm.model_executor.warmup.deepseek_v4_mhc_warmup import (
_warmup_broadcast_mhc,
)
calls = self._install_broadcast_spy(monkeypatch)
class CpuLayer:
hc_attn_fn_broadcast = torch.empty(0)
hc_attn_fn = torch.empty(0, device="cpu")
CpuLayer.__name__ = CpuLayer.__qualname__ = "DeepseekV4DecoderLayer"
_warmup_broadcast_mhc(self._model_with_layer(CpuLayer), token_sizes=[1, 2, 4])
assert not calls
class MockModel:
def modules(self):
yield self
yield CpuLayer()
def test_noop_when_broadcast_not_a_tensor(self, monkeypatch) -> None:
from vllm.model_executor.warmup.deepseek_v4_mhc_warmup import (
_warmup_broadcast_mhc,
)
_warmup_broadcast_mhc(MockModel(), token_sizes=[1, 2, 4])
calls = self._install_broadcast_spy(monkeypatch)
class BoolBroadcastLayer:
hc_attn_fn_broadcast = True
hc_attn_fn = SimpleNamespace(device=torch.device("cuda"))
_warmup_broadcast_mhc(
self._model_with_layer(BoolBroadcastLayer), token_sizes=[1, 2, 4]
)
assert not calls
def test_noop_when_broadcast_device_mismatch(self, monkeypatch) -> None:
from unittest import mock
from vllm.model_executor.warmup.deepseek_v4_mhc_warmup import (
_warmup_broadcast_mhc,
)
calls = self._install_broadcast_spy(monkeypatch)
broadcast = mock.MagicMock(spec=torch.Tensor)
broadcast.device = torch.device("cuda:1")
class MismatchedLayer:
hc_attn_fn_broadcast = broadcast
hc_attn_fn = SimpleNamespace(device=torch.device("cuda:0"))
_warmup_broadcast_mhc(
self._model_with_layer(MismatchedLayer), token_sizes=[1, 2, 4]
)
assert not calls
# ── _warmup_layer_mhc union behavior ─────────────────────────────────────
@@ -160,13 +160,13 @@ def _warmup_layer_mhc(
(layer.hc_attn_fn, layer.hc_attn_scale, layer.hc_attn_base),
(layer.hc_ffn_fn, layer.hc_ffn_scale, layer.hc_ffn_base),
):
layer_input, post_mix, comb_mix = layer.hc_pre(
layer_input, post_mix, res_mix = layer.hc_pre(
residual_slice,
fn,
scale,
base,
)
layer.hc_post(layer_input, residual_slice, post_mix, comb_mix)
layer.hc_post(layer_input, residual_slice, post_mix, res_mix)
def _warmup_hc_head(
@@ -241,6 +241,18 @@ def _warmup_broadcast_mhc(
if device.type != "cuda":
return
# Fail closed unless hc_attn_fn_broadcast is a tensor on the same CUDA
# device as hc_attn_fn. A non-tensor, a CPU tensor, or a tensor on a
# different device index indicates a partially initialized or mismatched
# broadcast-weight setup that must not be warmed up.
fn_broadcast = first_broadcast_layer.hc_attn_fn_broadcast
if not isinstance(fn_broadcast, torch.Tensor):
return
if fn_broadcast.device.type != "cuda":
return
if fn_broadcast.device != device:
return
hidden_size = first_broadcast_layer.hidden_size
broadcast_token_sizes = _select_mhc_split_key_token_sizes(
max_tokens=max(token_sizes),