Files
WildAi 08df29df25 feat: auto-detect config_name, drop VibeVoice-Large (v2.7.0)
Config auto-detection:
- config_detect: read the checkpoint embedding shape as an architecture
  fingerprint (header-only for safetensors, reuses the open GGUF reader);
  7B=[152064,3584], 1.5B=[151936,1536]; orientation-agnostic for
  shape-reversed GGUF files
- Auto-detect is the new default config_name; resolves the family before
  any heavy load, or fails fast with an actionable error (.bin/.pt and
  unknown families cannot be fingerprinted)
- an explicit config_name that contradicts the weights self-corrects to
  the detected family with one WARNING (reconcile_config)
- loader: friendly shape pre-check in _apply_state_dict names the
  offending tensors and hints at config_name instead of torch's raw
  size-mismatch RuntimeError

Dropdown dedup:
- VibeVoice-Large removed from config_name options (duplicate of 7B);
  kept as a legacy alias so saved workflows still load (normalize at node
  + loader entry; validate_inputs(**kwargs) override makes core skip its
  combo-membership check)
- node resolves Auto-detect before computing the cache identity so the
  request key matches the consumer's bundle-derived key (no per-run churn)

Console noise:
- demote ~30 internal INFO logs to DEBUG across loaders/patcher/registry
- drop two stray tie_weights prints; tied lm_head.weight no longer warned
  as missing (expected under tie_word_embeddings)

Tests: 1060 passed / 5 pre-existing failures / 4 skipped
2026-08-27 20:33:19 +03:00

422 lines
19 KiB
Python

"""Phase 2 tests: assign-based weight loading (plan 2026-08-18, D2/D3).
Covers ``VibeVoiceLoader._apply_state_dict``:
- assign semantics replace parameter objects (no copy, data_ptr identity)
- tied weights are re-tied after assign (un-defers DF-005)
- meta stragglers are zero-materialized
- missing/unexpected keys are reported
- tie is skipped when the config is not tied
- end-to-end: meta-instantiated tiny model + assign load yields checkpoint values
"""
import torch
import pytest
from unittest.mock import MagicMock
from ComfyUI_VibeVoice.modules.loader import VibeVoiceLoader
# ---------------------------------------------------------------------------
# Tiny test doubles
# ---------------------------------------------------------------------------
class _TiedTinyModel(torch.nn.Module):
"""Minimal model with a tied lm_head (mirrors VibeVoice-1.5B)."""
def __init__(self, vocab=4, dim=4):
super().__init__()
self.embed = torch.nn.Embedding(vocab, dim)
self.lm_head = torch.nn.Linear(dim, vocab, bias=False)
self.lm_head.weight = self.embed.weight # tie
# Config gate used by _apply_state_dict
self.config = MagicMock()
self.config.decoder_config.tie_word_embeddings = True
self.config.tie_word_embeddings = False
def tie_weights(self):
self.lm_head.weight = self.embed.weight
class _UntiedTinyModel(_TiedTinyModel):
"""Same model but with tying disabled in the config."""
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
self.config.decoder_config.tie_word_embeddings = False
self.tie_called = False
def tie_weights(self):
self.tie_called = True
class _MetaStragglerModel(torch.nn.Module):
"""Model with one parameter absent from the checkpoint (stays meta)."""
def __init__(self):
super().__init__()
self.present = torch.nn.Linear(2, 2, bias=False)
self.absent = torch.nn.Linear(2, 2, bias=False)
# Put 'absent' on meta to simulate a checkpoint that omits the key.
with torch.device("meta"):
self.absent = torch.nn.Linear(2, 2, bias=False)
self.config = MagicMock()
self.config.decoder_config.tie_word_embeddings = False
self.config.tie_word_embeddings = False
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestApplyStateDict:
def test_assign_replaces_parameters_no_copy(self):
model = _UntiedTinyModel()
w = torch.ones(4, 4)
sd = {"embed.weight": w}
VibeVoiceLoader._apply_state_dict(model, sd)
# assign=True: the parameter object IS the checkpoint tensor.
assert model.embed.weight.data_ptr() == w.data_ptr()
assert torch.equal(model.embed.weight, torch.ones(4, 4))
def test_tie_restored_after_assign(self):
model = _TiedTinyModel()
sd = {"embed.weight": torch.ones(4, 4)} # no lm_head.weight in sd
VibeVoiceLoader._apply_state_dict(model, sd)
# After re-tie, lm_head shares the (new) embed parameter.
assert model.lm_head.weight.data_ptr() == model.embed.weight.data_ptr()
assert torch.equal(model.lm_head.weight, torch.ones(4, 4))
def test_missing_meta_params_materialized_zero(self):
model = _MetaStragglerModel()
assert model.absent.weight.is_meta
sd = {"present.weight": torch.full((2, 2), 3.0)}
VibeVoiceLoader._apply_state_dict(model, sd)
# 'present' got the checkpoint value; 'absent' was zero-materialized.
assert torch.equal(model.present.weight, torch.full((2, 2), 3.0))
assert not model.absent.weight.is_meta
assert model.absent.weight.device.type == "cpu"
assert torch.equal(model.absent.weight, torch.zeros(2, 2))
assert model.absent.weight.shape == (2, 2)
def test_missing_unexpected_keys_reported(self):
model = _UntiedTinyModel()
sd = {"embed.weight": torch.ones(4, 4), "bogus.key": torch.zeros(1)}
missing, unexpected = VibeVoiceLoader._apply_state_dict(model, sd)
assert "bogus.key" in unexpected
# lm_head.weight is tied; with assign it is reported missing before re-tie.
assert isinstance(missing, list)
def test_tie_skipped_when_not_tied(self):
model = _UntiedTinyModel()
sd = {"embed.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
assert model.tie_called is False
def test_meta_instantiate_then_assign_yields_checkpoint_values(self):
"""End-to-end fast path: meta init + assign load (Phase 3 preview)."""
with torch.device("meta"):
model = _UntiedTinyModel()
assert all(p.is_meta for p in model.parameters())
sd = {"embed.weight": torch.full((4, 4), 7.0),
"lm_head.weight": torch.full((4, 4), 9.0)}
VibeVoiceLoader._apply_state_dict(model, sd)
assert not any(p.is_meta for p in model.parameters())
assert torch.equal(model.embed.weight, torch.full((4, 4), 7.0))
assert torch.equal(model.lm_head.weight, torch.full((4, 4), 9.0))
def test_meta_buffers_materialized_after_assign(self):
"""Meta buffers (not in checkpoint) must be materialized to CPU.
Without this, ComfyUI's unpatch_model → model.to(device) raises
NotImplementedError on meta buffers.
"""
class _ModelWithBuffer(_UntiedTinyModel):
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
self.register_buffer("position_ids", torch.arange(vocab))
with torch.device("meta"):
model = _ModelWithBuffer()
# Buffer is on meta.
assert model.position_ids.is_meta
# Checkpoint has no buffer key.
sd = {"embed.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
# Buffer is now on CPU (materialized).
assert not model.position_ids.is_meta
assert model.position_ids.device.type == "cpu"
def test_no_meta_tensors_remain_after_assign(self):
"""After _apply_state_dict, NO tensor (param or buffer) is on meta."""
class _ModelWithBuffer(_UntiedTinyModel):
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
self.register_buffer("mask", torch.ones(vocab))
with torch.device("meta"):
model = _ModelWithBuffer()
sd = {"embed.weight": torch.ones(4, 4), "lm_head.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
# All parameters AND buffers must be off meta.
assert not any(p.is_meta for p in model.parameters())
assert not any(b.is_meta for b in model.buffers())
class TestSentinelBufferRegression:
"""v2.3.1: meta-init must not zero sentinel buffers (silent-output bug).
The VibeVoice model registers ``speech_scaling_factor`` /
``speech_bias_factor`` as ``float('nan')`` and computes them at inference
time. The diffusion-inversion gate (modeling_vibevoice.py:623) is
``if not torch.isnan(sf) and not torch.isnan(bf)``. Zeroing these buffers
makes the gate TRUE and applies ``speech / 0 - 0`` → silent output. The
hotfix in 6b99be9 zero-materialized ALL meta buffers, which regressed this.
"""
def test_sentinel_buffers_restored_to_nan(self):
"""speech_scaling_factor / speech_bias_factor stay nan after meta load."""
class _SentinelModel(_UntiedTinyModel):
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
self.register_buffer("speech_scaling_factor", torch.tensor(float("nan")))
self.register_buffer("speech_bias_factor", torch.tensor(float("nan")))
with torch.device("meta"):
model = _SentinelModel()
# Buffers are on meta after meta-init.
assert model.speech_scaling_factor.is_meta
assert model.speech_bias_factor.is_meta
sd = {"embed.weight": torch.ones(4, 4), "lm_head.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
# They must be restored to nan (NOT zero) so the isnan gate works.
assert not model.speech_scaling_factor.is_meta
assert not model.speech_bias_factor.is_meta
assert torch.isnan(model.speech_scaling_factor).item()
assert torch.isnan(model.speech_bias_factor).item()
def test_fix_std_sentinel_restored_to_config_value(self):
"""fix_std (persistent=False) is restored to its config value (0.5)."""
class _FixStdModel(_UntiedTinyModel):
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
self.register_buffer("fix_std", torch.tensor(0.5), persistent=False)
with torch.device("meta"):
model = _FixStdModel()
assert model.fix_std.is_meta
sd = {"embed.weight": torch.ones(4, 4), "lm_head.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
assert not model.fix_std.is_meta
assert model.fix_std.item() == 0.5
def test_non_sentinel_buffer_still_zeroed(self):
"""Ordinary meta buffers (e.g. position_ids) are still zeroed."""
class _PlainBufferModel(_UntiedTinyModel):
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
self.register_buffer("position_ids", torch.arange(vocab))
with torch.device("meta"):
model = _PlainBufferModel()
sd = {"embed.weight": torch.ones(4, 4), "lm_head.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
assert not model.position_ids.is_meta
assert torch.equal(model.position_ids, torch.zeros(4))
class TestRopeRecomputeRegression:
"""v2.3.2: meta-init must not leave RoPE inv_freq zeroed (gibberish bug).
Rotary embeddings compute ``inv_freq`` from the config in ``__init__``;
the buffer is non-persistent and never in the checkpoint. Meta-init makes
it a meta tensor and the generic zero-materialization destroys it. With
``inv_freq == 0`` RoPE yields cos(0)=1 / sin(0)=0 → no positional
encoding → the LM emits gibberish and stops early (~2 s clip). A
deterministic eager-vs-meta diff on the real VibeVoice-1.5B checkpoint
showed these buffers were the ONLY tensors that differed.
"""
def _make_rope_model(self):
"""Build a tiny model with a real Qwen2RotaryEmbedding submodule."""
from transformers.models.qwen2.modeling_qwen2 import Qwen2RotaryEmbedding
from transformers.models.qwen2.configuration_qwen2 import Qwen2Config
class _RopeModel(_UntiedTinyModel):
def __init__(self, vocab=4, dim=4):
super().__init__(vocab, dim)
rope_cfg = Qwen2Config(
hidden_size=16,
num_attention_heads=2,
max_position_embeddings=64,
rope_theta=10000.0,
)
self.rotary_emb = Qwen2RotaryEmbedding(rope_cfg)
return _RopeModel
def test_rope_inv_freq_recomputed_after_meta_load(self):
"""inv_freq must be recomputed (non-zero) after meta-init + assign."""
_RopeModel = self._make_rope_model()
with torch.device("meta"):
model = _RopeModel()
# Under meta-init the rotary buffers are on meta.
assert model.rotary_emb.inv_freq.is_meta
sd = {"embed.weight": torch.ones(4, 4), "lm_head.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
# After the fix: inv_freq is real, on CPU, and NON-zero.
assert not model.rotary_emb.inv_freq.is_meta
assert model.rotary_emb.inv_freq.device.type == "cpu"
assert (model.rotary_emb.inv_freq != 0).any()
assert not model.rotary_emb.original_inv_freq.is_meta
assert (model.rotary_emb.original_inv_freq != 0).any()
def test_rope_inv_freq_matches_eager_init(self):
"""Recomputed inv_freq must equal the eager-init value exactly."""
_RopeModel = self._make_rope_model()
# Eager reference.
eager = _RopeModel()
expected = eager.rotary_emb.inv_freq.detach().clone()
# Meta path.
with torch.device("meta"):
model = _RopeModel()
sd = {"embed.weight": torch.ones(4, 4), "lm_head.weight": torch.ones(4, 4)}
VibeVoiceLoader._apply_state_dict(model, sd)
assert torch.equal(model.rotary_emb.inv_freq, expected)
assert torch.equal(model.rotary_emb.original_inv_freq, expected)
def test_recompute_helper_returns_count(self):
"""_recompute_rope_buffers returns the number of rotary modules fixed."""
from ComfyUI_VibeVoice.modules.loader import _recompute_rope_buffers
_RopeModel = self._make_rope_model()
with torch.device("meta"):
model = _RopeModel()
# Zero-materialize the meta buffers first (as _apply_state_dict does).
for name, buf in list(model.named_buffers()):
if buf.is_meta:
parts = name.split(".")
mod = model
for pt in parts[:-1]:
mod = getattr(mod, pt)
setattr(mod, parts[-1], torch.zeros(buf.shape, dtype=buf.dtype))
n = _recompute_rope_buffers(model)
assert n == 1
assert (model.rotary_emb.inv_freq != 0).any()
def test_recompute_noop_without_rotary_modules(self):
"""Models without rotary embeddings are unaffected."""
from ComfyUI_VibeVoice.modules.loader import _recompute_rope_buffers
model = _UntiedTinyModel()
n = _recompute_rope_buffers(model)
assert n == 0
# ---------------------------------------------------------------------------
# Phase 5 (plan 2026-08-27, D8): friendly shape pre-check
# ---------------------------------------------------------------------------
class TestShapePreCheck:
"""A config/weights family mismatch raises a clear ValueError naming the
offending keys + a config hint BEFORE torch's raw 'size mismatch' error."""
def test_shape_mismatch_raises_friendly_valueerror(self):
model = _UntiedTinyModel() # embed.weight is (4, 4)
sd = {"embed.weight": torch.ones(6, 6)} # wrong family shape
with pytest.raises(ValueError) as excinfo:
VibeVoiceLoader._apply_state_dict(model, sd)
msg = str(excinfo.value)
# Names the offending key...
assert "embed.weight" in msg
# ...shows both shapes...
assert "(6, 6)" in msg and "(4, 4)" in msg
# ...and hints at the likely cause (config mismatch).
assert "config" in msg.lower()
def test_shape_mismatch_mentions_auto_detect(self):
model = _UntiedTinyModel()
sd = {"embed.weight": torch.ones(6, 6)}
with pytest.raises(ValueError) as excinfo:
VibeVoiceLoader._apply_state_dict(model, sd)
assert "Auto-detect" in str(excinfo.value)
def test_matching_shapes_load_normally(self):
model = _UntiedTinyModel()
sd = {"embed.weight": torch.full((4, 4), 5.0)}
# No ValueError; the value is assigned.
VibeVoiceLoader._apply_state_dict(model, sd)
assert torch.equal(model.embed.weight, torch.full((4, 4), 5.0))
def test_missing_and_unexpected_keys_not_flagged(self):
"""Only SHARED keys are shape-checked; missing/unexpected are left to
strict=False (they must NOT trigger the shape pre-check)."""
model = _TiedTinyModel()
# lm_head.weight is tied (absent from sd -> missing); bogus -> unexpected.
sd = {"embed.weight": torch.ones(4, 4), "bogus.key": torch.zeros(9, 9)}
missing, unexpected = VibeVoiceLoader._apply_state_dict(model, sd)
assert "bogus.key" in unexpected
assert isinstance(missing, list)
def test_meta_param_model_matching_shapes_load(self):
"""Meta-initialized models still pass the pre-check (meta has shape)."""
with torch.device("meta"):
model = _UntiedTinyModel()
assert all(p.is_meta for p in model.parameters())
sd = {"embed.weight": torch.full((4, 4), 2.0),
"lm_head.weight": torch.full((4, 4), 3.0)}
VibeVoiceLoader._apply_state_dict(model, sd)
assert torch.equal(model.embed.weight, torch.full((4, 4), 2.0))
def test_meta_param_model_shape_mismatch_caught(self):
"""The pre-check works on meta models (shapes exist before load)."""
with torch.device("meta"):
model = _UntiedTinyModel()
sd = {"embed.weight": torch.ones(5, 5)}
with pytest.raises(ValueError, match="config"):
VibeVoiceLoader._apply_state_dict(model, sd)
def test_buffer_shape_mismatch_caught(self):
class _BufModel(_UntiedTinyModel):
def __init__(self):
super().__init__()
self.register_buffer("position_ids", torch.arange(4))
model = _BufModel()
sd = {"embed.weight": torch.ones(4, 4), "position_ids": torch.arange(9)}
with pytest.raises(ValueError) as excinfo:
VibeVoiceLoader._apply_state_dict(model, sd)
assert "position_ids" in str(excinfo.value)
def test_many_mismatches_truncated_with_more(self):
class _WideModel(torch.nn.Module):
def __init__(self):
super().__init__()
for i in range(8):
setattr(self, f"w{i}", torch.nn.Linear(2, 2, bias=False))
self.config = MagicMock()
self.config.decoder_config.tie_word_embeddings = False
self.config.tie_word_embeddings = False
model = _WideModel()
sd = {f"w{i}.weight": torch.ones(7, 7) for i in range(8)}
with pytest.raises(ValueError) as excinfo:
VibeVoiceLoader._apply_state_dict(model, sd)
msg = str(excinfo.value)
assert "8 shared" in msg
assert "... and" in msg # truncated beyond max_reported
class TestAssertShapesCompatibleHelper:
"""Direct unit tests for the _assert_shapes_compatible helper."""
def test_no_mismatch_returns_none(self):
from ComfyUI_VibeVoice.modules.loader import _assert_shapes_compatible
model = _UntiedTinyModel()
assert _assert_shapes_compatible(model, {"embed.weight": torch.ones(4, 4)}) is None
def test_empty_state_dict_ok(self):
from ComfyUI_VibeVoice.modules.loader import _assert_shapes_compatible
model = _UntiedTinyModel()
assert _assert_shapes_compatible(model, {}) is None
def test_non_tensor_value_skipped(self):
from ComfyUI_VibeVoice.modules.loader import _assert_shapes_compatible
model = _UntiedTinyModel()
# A non-tensor entry (no .shape) must not crash the pre-check.
assert _assert_shapes_compatible(model, {"embed.weight": "not-a-tensor"}) is None