Files
WildAi 567fb03e30 feat: align VibeVoice loading with core DynamicVRAM, fix RAM ghost (v2.12.0)
Loading filled host RAM 20->40GB while streaming from SSD and never released
it; inference then re-served the whole model per AR step. Teardown against
core's aimdo/DynamicVRAM machinery, measured on the real 9.47GB 7B fp8 and
5.16GB 1.5B bf16 checkpoints, found four loader defects plus one that was
never ours to fix.

Defects fixed (each with before/after numbers):

* Quant pass 1 pre-read each dequant-at-load layer's scale with
  safetensors.safe_open + get_tensor. On Windows one such call commits ~1x
  file size as private, untouched memory, pinned for the lifetime of the
  returned tensor: +9,050MB on the 7B fp8, invisible to both the working-set
  counter and the storage census. Now maps the file once through core's
  aimdo arm and clones only the tiny scales. Whole load: 1.51x -> 0.13x file
  size; retained after free 11.3GB -> 96MB.

* Quant families were excluded from core's dynamic patcher by an inherited
  "quant streams natively" stop condition. Because a legacy patcher has
  nowhere to page from, those routes cloned every tensor into host RAM and
  fully H2D'd it. select_patcher_class now follows core's availability rule
  for every family, with gguf_block still excluded for a measured reason
  (GGUFTensor.from_reader_tensor clones the reader's view).

* replace_linears_for_quant built its resident modules outside the meta
  context: 8.08GB of never-written host allocation per 7B fp8 load.

* Dormant weight_function/bias_function double application: core already
  applies them inside cast_bias_weight (ops.py:431-438).

Exonerated with numbers, not assumed: the dense read path is core's own
(13.5MB private for a 5.16GB file), and core's file->VRAM paging read is
cache-clean (+0.13GB per 2GB). Only host-side view reads reproduce the
reported 1:1 RAM at 0.55GB/s signature.

Inference: reproduced headlessly that when the tree does not fit in the
VRAM that is actually free, every forward re-reads every weight from the
checkpoint file (3072/3072 file reads, 21ms -> 125ms per step). Our wrappers
are protocol-correct; residency is stable at 100% with headroom. [vvpull]
now reports resident vs reread with bytes and free VRAM so one live line
settles it, behind VIBEVOICE_VBAR_OBSERVER=0.

Not changed after being tried and reverted: deriving fast_disk from the
checkpoint. It measured as a no-op on the dev host but converted RAM-speed
pinned re-reads into disk-speed reads live, which made loading dramatically
slower. Reverted in full; see the report for the probe-design lesson.

Adds host-RAM instrumentation ([vvrss] with machine-level start/end, [vvcensus]
storage census, tts-generate bracket), four standalone probes, and ~4000
lines of tests pinning the invariants above.
2026-09-30 18:24:52 +03:00

1001 lines
42 KiB
Python

"""Phase D tests: quant-resident paths through load_external_vibevoice_model.
Exercises the REAL loader with a synthetic GGUF file and a REAL stub module
tree (meta-initialized, like production), asserting:
- residents are GGUFLinear modules holding raw uint8 bytes,
- dense floats keep native dtypes (BF16 embed stays BF16),
- NO gguf.dequantize call happens on the load path (the RAM-spike proof),
- bundle records weight_family/quant_stats,
- ConvRot safetensors branch swaps linears before assignment.
"""
import json
import pytest
import torch
from unittest.mock import MagicMock, patch
import gguf
from safetensors.torch import save_file
from ComfyUI_VibeVoice.modules import external_loader as EL
from ComfyUI_VibeVoice.modules.external_loader import (
_install_gguf_weights,
load_external_vibevoice_model,
)
from ComfyUI_VibeVoice.modules.gguf_quant import GGUFLinear
from ComfyUI_VibeVoice.modules.convrot_quant import ConvRotInt8Linear
from ComfyUI_VibeVoice.modules.fp8_quant import FP8Linear
from conftest import build_stub_vv, stub_vv_gguf_spec
class _FakeStreamingCfg:
pass
@pytest.fixture
def stubbed_load_env():
"""Common mocks for load_external_vibevoice_model around a REAL model tree.
Yields (run_fn, state) where run_fn(weight_path, **kw) executes the loader.
Pass dims=(n_layers, hidden, ffn) to match the GGUF spec under test
(K-quants require in_features % 256 == 0).
"""
holder = {"dims": (1, 256, 512)}
def _instantiate(config, is_streaming, attn_implementation, final_load_dtype,
use_meta=True):
n_layers, hidden, ffn = holder["dims"]
return build_stub_vv(n_layers=n_layers, hidden=hidden, ffn=ffn)
def _run(weight_path, dims=None, **kwargs):
if dims is not None:
holder["dims"] = dims
kwargs.setdefault("config_name", "VibeVoice-1.5B")
kwargs.setdefault("attention_mode", "sdpa")
kwargs.setdefault("use_llm_4bit", False)
kwargs.setdefault("dtype_str", "auto")
with patch.object(EL.VibeVoiceLoader, "_load_config",
return_value=holder.get("config", MagicMock())), \
patch.object(EL.VibeVoiceLoader, "_load_tokenizer", return_value=MagicMock()), \
patch.object(EL.VibeVoiceLoader, "_load_processor", return_value=MagicMock()), \
patch.object(EL.VibeVoiceLoader, "_instantiate_model",
side_effect=_instantiate), \
patch.object(EL, "resolve_sidecar_config", return_value="/fake/config.json"), \
patch.object(EL, "resolve_sidecar_preprocessor", return_value=""), \
patch.object(EL, "resolve_sidecar_tokenizer_dir", return_value="/fake/dir"), \
patch.object(EL, "resolve_dtype", return_value=torch.bfloat16), \
patch.object(EL, "resolve_attention_mode", side_effect=lambda m, q: m), \
patch.object(EL, "get_attn_implementation_for_load", return_value="eager"), \
patch.object(EL, "VibeVoiceStreamingConfig", _FakeStreamingCfg), \
patch.object(EL.model_management, "get_torch_device",
return_value=torch.device("cpu")):
bundle = load_external_vibevoice_model(
weight_path=str(weight_path), **kwargs
)
holder["bundle"] = bundle
holder["model"] = bundle["model"]
return bundle
return _run, holder
class TestGGUFLazyInstall:
def test_residents_installed_and_dense_kept_native(
self, make_gguf_file, stubbed_load_env):
run, _ = stubbed_load_env
path = make_gguf_file(stub_vv_gguf_spec(n_layers=1, qtype="Q8_0"),
tag="lazy")
with patch.object(gguf, "dequantize",
side_effect=AssertionError(
"legacy full dequant ran on the lazy path")):
bundle = run(path, dims=(1, 64, 128))
assert bundle["weight_family"] == "gguf_block"
assert bundle["quant_stats"]["n_resident_layers"] == 7
model = bundle["model"]
q = model.model.language_model.layers[0].self_attn.q_proj
assert isinstance(q, GGUFLinear)
assert q.weight.dtype == torch.uint8
assert not q.weight.is_meta
assert q._quant_resident
# Dense passthrough keeps NATIVE dtype through install+cast:
# BF16 embed stays BF16 under a bf16 target dtype.
embed = model.model.language_model.embed_tokens.weight
assert embed.dtype == torch.bfloat16
assert not embed.is_meta
# Resident byte accounting equals raw Q8_0 sizes.
total = sum(
m.weight.numel() for m in model.modules() if isinstance(m, GGUFLinear)
)
assert total == bundle["quant_stats"]["raw_bytes"]
assert total > 0
def test_no_full_dequant_on_real_reader(self, make_gguf_file, stubbed_load_env):
"""Second, stricter spike-proof variant: gguf.dequantize must be called
ZERO times; even indirect numpy float expansion of quant blocks is
avoided because we only touch .data bytes. Uses Q4_K, which requires
in_features % 256 == 0 -> hidden=256, ffn=512."""
run, _ = stubbed_load_env
path = make_gguf_file(
stub_vv_gguf_spec(n_layers=1, hidden=256, ffn=512, qtype="Q4_K"),
tag="kq",
)
calls = []
real_deq = gguf.dequantize
def _spy(data, qtype):
calls.append(qtype)
return real_deq(data, qtype)
with patch.object(gguf, "dequantize", side_effect=_spy):
bundle = run(path, dims=(1, 256, 512))
assert calls == [], "gguf.dequantize must not run on the lazy path"
assert bundle["weight_family"] == "gguf_block"
down = bundle["model"].model.language_model.layers[0].mlp.down_proj
assert isinstance(down, GGUFLinear)
assert down.ggml_type == gguf.constants.GGMLQuantizationType.Q4_K
def test_forward_works_after_install(self, make_gguf_file, stubbed_load_env):
"""Smoke: a resident linear produces finite outputs post-install."""
run, _ = stubbed_load_env
path = make_gguf_file(stub_vv_gguf_spec(n_layers=1, qtype="Q8_0"),
tag="fwd")
bundle = run(path, dims=(1, 64, 128))
q = bundle["model"].model.language_model.layers[0].self_attn.q_proj
x = torch.randn(2, q.in_features)
y = q(x)
assert y.shape == (2, q.out_features)
assert torch.isfinite(y).all()
class TestConvRotLoaderBranch:
def test_convrot_checkpoint_replaces_linears(self, tmp_path, stubbed_load_env):
run, _ = stubbed_load_env
prefix = "model.language_model.layers.0.self_attn.q_proj"
tensors = {
f"{prefix}.weight": torch.randint(-100, 100, (64, 64), dtype=torch.int8),
f"{prefix}.weight_scale": torch.full((64, 1), 0.01),
}
meta = json.dumps({
"format": "int8_tensorwise", "convrot": True,
"convrot_groupsize": 16,
}).encode("utf-8")
tensors[f"{prefix}.comfy_quant"] = torch.frombuffer(
bytearray(meta), dtype=torch.uint8
)
p = tmp_path / "convrot.safetensors"
save_file(tensors, str(p))
bundle = run(p, dims=(1, 64, 128))
assert bundle["weight_family"] == "convrot_int8"
assert bundle["quant_stats"]["n_resident_layers"] == 1
q = bundle["model"].model.language_model.layers[0].self_attn.q_proj
assert isinstance(q, ConvRotInt8Linear)
assert q.weight.dtype == torch.int8
assert not q.weight.is_meta
assert q.convrot_groupsize == 16
def test_convrot_with_bnb_rejected(self, tmp_path, stubbed_load_env):
run, _ = stubbed_load_env
prefix = "model.prediction_head.cond_proj" # exists in the stub tree
tensors = {
f"{prefix}.weight": torch.randint(-100, 100, (64, 64), dtype=torch.int8),
f"{prefix}.weight_scale": torch.full((64, 1), 0.01),
}
meta = json.dumps({"format": "int8_tensorwise", "convrot": True,
"convrot_groupsize": 64}).encode("utf-8")
tensors[f"{prefix}.comfy_quant"] = torch.frombuffer(
bytearray(meta), dtype=torch.uint8
)
p = tmp_path / "convrot2.safetensors"
save_file(tensors, str(p))
from ComfyUI_VibeVoice.modules.quant_common import QuantTargetMismatch
with pytest.raises((ValueError, QuantTargetMismatch)):
run(p, dims=(1, 64, 128), use_llm_4bit=True)
class TestGGUFInstallUnitLevel:
def test_unmapped_module_raises_actionable(self, make_gguf_file):
"""A quantized tensor targeting a non-existent module fails loudly."""
spec = [("does.not.exist.weight", "Q8_0", (32, 32))]
path = make_gguf_file(spec, tag="badmap")
reader = gguf.GGUFReader(str(path))
model = build_stub_vv()
from ComfyUI_VibeVoice.modules.quant_common import QuantTargetMismatch
# Either layer of the pipeline reports the offending key: the
# key-mapper (UnmappedKeyError, a ValueError) or the module
#-resolution check (QuantTargetMismatch).
with pytest.raises((QuantTargetMismatch, ValueError), match="does.not.exist"):
_install_gguf_weights(model, reader)
def test_shape_mismatch_raises(self, make_gguf_file):
spec = [("model.language_model.norm.weight", "Q8_0", (128, 32))]
# norm is a plain parameter (not Linear); shape (32,) vs file (32,128)
path = make_gguf_file(spec, tag="badshape")
reader = gguf.GGUFReader(str(path))
model = build_stub_vv()
from ComfyUI_VibeVoice.modules.quant_common import QuantTargetMismatch
with pytest.raises(QuantTargetMismatch):
_install_gguf_weights(model, reader)
class TestASRGGUFPath:
def test_asr_gguf_lazy_branch(self, make_gguf_file):
"""The ASR branch uses the same lazy installer for GGUF files."""
from ComfyUI_VibeVoice.modules.external_loader import (
load_external_vibevoice_asr_model,
)
class _StubASR(torch.nn.Module):
"""Exposes the stub tree at top level so checkpoint keys
(model.language_model.*, lm_head.*) resolve directly."""
def __init__(self, config):
super().__init__()
vv = build_stub_vv(n_layers=1)
self.model = vv.model
self.lm_head = vv.lm_head
path = make_gguf_file(stub_vv_gguf_spec(n_layers=1, qtype="Q8_0"),
tag="asr")
with patch.object(EL, "VibeVoiceASRForConditionalGeneration", _StubASR), \
patch.object(EL, "_load_asr_config", return_value=MagicMock()), \
patch.object(EL, "_load_asr_tokenizer", return_value=MagicMock()), \
patch.object(EL, "_load_asr_processor", return_value=MagicMock()), \
patch.object(EL, "resolve_sidecar_config", return_value="/fake/c.json"), \
patch.object(EL, "resolve_sidecar_preprocessor", return_value=""), \
patch.object(EL, "resolve_sidecar_tokenizer_dir", return_value="/fake/d"), \
patch.object(EL, "resolve_dtype", return_value=torch.float32), \
patch.object(EL, "resolve_attention_mode",
side_effect=lambda m, quantize_4bit=False: m), \
patch.object(EL, "get_attn_implementation_for_load", return_value="eager"), \
patch.object(EL.model_management, "get_torch_device",
return_value=torch.device("cpu")), \
patch.object(gguf, "dequantize",
side_effect=AssertionError("full dequant on ASR path")):
bundle = load_external_vibevoice_asr_model(
weight_path=str(path), config_name="VibeVoice-ASR",
attention_mode="sdpa", dtype_str="auto",
)
assert bundle["is_asr"] is True
assert bundle["weight_family"] == "gguf_block"
# ====================================================================
# Rowwise int8 / fp8 dequant-at-load + hard-fail propagation
# ====================================================================
import json as _json
from ComfyUI_VibeVoice.modules.convrot_quant import QuantLayerInfo
from ComfyUI_VibeVoice.modules.external_loader import (
_assert_dense_loadable,
_prepare_quantized_safetensors_load,
)
def _save_qcheckpoint(path, entries):
"""entries: list of (prefix, kind, shape) where kind in
{'convrot', 'rowwise', 'rowwise_fp8'}."""
from safetensors.torch import save_file
tensors = {}
for prefix, kind, shape in entries:
if kind == "rowwise_fp8":
tensors[f"{prefix}.weight"] = torch.randint(
-100, 100, shape).to(torch.float8_e4m3fn)
# Real fp8 checkpoints carry PER-TENSOR scalar scales.
tensors[f"{prefix}.weight_scale"] = torch.tensor(0.5)
else:
tensors[f"{prefix}.weight"] = torch.randint(
-100, 100, shape, dtype=torch.int8)
if kind == "block":
gs = 16
tensors[f"{prefix}.weight_scale"] = torch.rand(
(shape[0] // gs, shape[1] // gs),
dtype=torch.float32) * 0.1
else:
tensors[f"{prefix}.weight_scale"] = torch.full(
(shape[0], 1), 0.5)
if kind == "convrot":
meta = {"convrot": True, "convrot_groupsize": 16,
"format": "int8_tensorwise",
"orig_dtype": "torch.bfloat16"}
elif kind == "rowwise":
meta = {"format": "int8_tensorwise", "per_row": True,
"orig_dtype": "torch.bfloat16"}
elif kind == "block":
meta = {"format": "int8_blockwise", "group_size": 16,
"orig_dtype": "torch.bfloat16"}
else:
meta = {"format": "float8_e4m3fn",
"orig_dtype": "torch.bfloat16"}
tensors[f"{prefix}.comfy_quant"] = torch.frombuffer(
bytearray(_json.dumps(meta).encode("utf-8")), dtype=torch.uint8
)
save_file(tensors, str(path))
return path
class TestPrepareQuantizedSafetensorsLoad:
def test_rowwise_dequant_math_and_key_stripping(self):
sd = {
"l.weight": torch.tensor([[100, -101], [50, -25]],
dtype=torch.int8),
"l.weight_scale": torch.tensor([[0.5], [2.0]]),
"l.comfy_quant": torch.zeros(4, dtype=torch.uint8),
"l.bias": torch.zeros(2),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, in_features=2, out_features=2,
convrot=False, orig_dtype="torch.bfloat16",
)}
plan, n_rowwise, n_fp8 = _prepare_quantized_safetensors_load(sd, qmap)
assert n_rowwise == 1 and plan == {}
expected = (
torch.tensor([[100., -101.], [50., -25.]])
* torch.tensor([[0.5], [2.0]])
).to(torch.bfloat16)
assert sd["l.weight"].dtype == torch.bfloat16
assert torch.equal(sd["l.weight"], expected)
assert "l.weight_scale" not in sd
assert "l.comfy_quant" not in sd
assert "l.bias" in sd # untouched
def test_convrot_entries_become_plan_and_keep_scale(self):
sd = {
"c.weight": torch.zeros(64, 64, dtype=torch.int8),
"c.weight_scale": torch.full((64, 1), 0.01),
"c.comfy_quant": torch.zeros(3, dtype=torch.uint8),
}
qmap = {"c": QuantLayerInfo(
prefix="c", group_size=16, in_features=64, out_features=64,
convrot=True,
)}
plan, n_rowwise, n_fp8 = _prepare_quantized_safetensors_load(sd, qmap)
assert n_rowwise == 0 and set(plan) == {"c"}
assert "c.weight_scale" in sd # consumed by resident assign
assert "c.comfy_quant" not in sd # metadata stripped up front
def test_missing_scale_raises(self):
sd = {"l.weight": torch.zeros(2, 2, dtype=torch.int8)}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, convrot=False,
orig_dtype="torch.bfloat16",
)}
with pytest.raises(Exception, match="missing its"):
_prepare_quantized_safetensors_load(sd, qmap)
def test_scalar_scale_accepted(self):
sd = {
"l.weight": torch.tensor([[100, -101]], dtype=torch.int8),
"l.weight_scale": torch.tensor(0.5),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, in_features=2, out_features=1,
convrot=False, orig_dtype="torch.bfloat16",
)}
plan, n_rowwise, n_fp8 = _prepare_quantized_safetensors_load(sd, qmap)
assert n_rowwise == 1
assert torch.equal(
sd["l.weight"],
torch.tensor([[50.0, -50.5]]).to(torch.bfloat16),
)
def test_bad_scale_shape_raises(self):
sd = {
"l.weight": torch.zeros(2, 2, dtype=torch.int8),
"l.weight_scale": torch.zeros(2), # neither scalar nor [out,1]
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, convrot=False,
orig_dtype="torch.bfloat16", out_features=2, in_features=2,
)}
with pytest.raises(Exception, match="neither a scalar"):
_prepare_quantized_safetensors_load(sd, qmap)
def test_weight_metadata_mismatch_raises(self):
sd = {
"l.weight": torch.zeros(2, 2, dtype=torch.int8),
"l.weight_scale": torch.tensor(0.5),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, convrot=False,
orig_dtype="torch.bfloat16", out_features=4, in_features=4,
)}
with pytest.raises(Exception, match="disagrees with metadata"):
_prepare_quantized_safetensors_load(sd, qmap)
class TestRowwiseLoaderIntegration:
def test_mixed_convrot_rowwise_end_to_end(self, tmp_path, stubbed_load_env):
"""Real-world layout: some layers rotated residents, the rest plain
rowwise dequantized at load (the user's int8_convrot file)."""
run, holder = stubbed_load_env
p = _save_qcheckpoint(tmp_path / "mixed.safetensors", [
("model.language_model.layers.0.self_attn.q_proj",
"convrot", (64, 64)),
("model.prediction_head.cond_proj", "rowwise", (64, 64)),
])
bundle = run(p, dims=(1, 64, 128))
assert bundle["weight_family"] == "convrot_int8"
assert bundle["quant_stats"] == {
"n_resident_layers": 1,
"n_rowwise_layers": 1,
"n_fp8_resident_layers": 0,
}
q = bundle["model"].model.language_model.layers[0].self_attn.q_proj
assert isinstance(q, ConvRotInt8Linear)
cond = bundle["model"].model.prediction_head.cond_proj
assert isinstance(cond, torch.nn.Linear)
assert not isinstance(cond, ConvRotInt8Linear)
# Dequanted value matches q * per-row scale exactly (bf16 cast).
from safetensors.torch import load_file
raw = load_file(str(p))
ref_w = (raw["model.prediction_head.cond_proj.weight"].to(torch.float32)
* raw["model.prediction_head.cond_proj.weight_scale"]).to(torch.bfloat16)
assert cond.weight.dtype == torch.bfloat16
assert torch.equal(cond.weight.data, ref_w)
def test_fp8_per_row_scale_dequants_at_load(self, tmp_path, stubbed_load_env):
"""[out, 1] per-row fp8 scales cannot run through the per-tensor
kitchen kernel -> legacy dequant-at-load (plan 2026-08-27, D2)."""
run, holder = stubbed_load_env
from safetensors.torch import save_file as _sf
prefix = "model.prediction_head.cond_proj"
tensors = {
f"{prefix}.weight": torch.randint(-100, 100, (64, 64)).to(
torch.float8_e4m3fn),
f"{prefix}.weight_scale": torch.full((64, 1), 0.5),
}
tensors[f"{prefix}.comfy_quant"] = torch.frombuffer(
bytearray(_json.dumps({"format": "float8_e4m3fn",
"orig_dtype": "torch.bfloat16"})
.encode("utf-8")), dtype=torch.uint8
)
p = tmp_path / "fp8_perrow.safetensors"
_sf(tensors, str(p))
bundle = run(p, dims=(1, 64, 128))
assert bundle["quant_stats"]["n_rowwise_layers"] == 1
assert bundle["quant_stats"]["n_fp8_resident_layers"] == 0
cond = bundle["model"].model.prediction_head.cond_proj
assert not isinstance(cond, FP8Linear) # stays a plain Linear
assert cond.weight.dtype == torch.bfloat16
ref_w = (tensors[f"{prefix}.weight"].float()
* tensors[f"{prefix}.weight_scale"]).to(torch.bfloat16)
assert torch.equal(cond.weight.data, ref_w)
class TestFP8ResidentLoading:
"""Scalar-scale fp8 checkpoints execute resident (plan 2026-08-27, D1)."""
def test_fp8_scalar_scale_resident_end_to_end(self, tmp_path, stubbed_load_env):
run, holder = stubbed_load_env
p = _save_qcheckpoint(tmp_path / "fp8.safetensors", [
("model.prediction_head.cond_proj", "rowwise_fp8", (64, 64)),
])
bundle = run(p, dims=(1, 64, 128))
assert bundle["weight_family"] == "fp8_resident"
assert bundle["quant_stats"] == {
"n_resident_layers": 1,
"n_rowwise_layers": 0,
"n_fp8_resident_layers": 1,
}
cond = bundle["model"].model.prediction_head.cond_proj
assert isinstance(cond, FP8Linear)
assert cond.weight.dtype == torch.float8_e4m3fn
assert not cond.weight.is_meta
assert cond.weight_scale.dtype == torch.float32
assert not cond.weight_scale.is_meta
# Forward parity vs the manual dequant reference (bit-exact eager).
from safetensors.torch import load_file
raw = load_file(str(p))
w = raw["model.prediction_head.cond_proj.weight"]
s = raw["model.prediction_head.cond_proj.weight_scale"]
x = torch.randn(2, 64, dtype=torch.bfloat16)
ref = torch.nn.functional.linear(
x, (w.float() * s).to(torch.bfloat16)
)
assert torch.equal(cond(x), ref)
def test_fp8_resident_survives_final_dtype_cast(self, tmp_path,
stubbed_load_env):
"""cast_model_to_dtype_if_needed(bf16) after load must not touch the
fp8 storage or the fp32 scale (_quant_protected_names)."""
run, holder = stubbed_load_env
p = _save_qcheckpoint(tmp_path / "fp8cast.safetensors", [
("model.prediction_head.cond_proj", "rowwise_fp8", (64, 64)),
])
bundle = run(p, dims=(1, 64, 128))
cond = bundle["model"].model.prediction_head.cond_proj
assert cond.weight.dtype == torch.float8_e4m3fn
assert cond.weight_scale.dtype == torch.float32
def test_fp8_no_backend_falls_back_to_dequant(self, tmp_path,
stubbed_load_env):
"""Without a kitchen fp8 backend the scalar-scale file dequantizes at
load (correct, just heavier) instead of failing (plan D3)."""
run, holder = stubbed_load_env
from ComfyUI_VibeVoice.modules import fp8_quant as FQ
p = _save_qcheckpoint(tmp_path / "fp8fb.safetensors", [
("model.prediction_head.cond_proj", "rowwise_fp8", (64, 64)),
])
with patch.object(FQ, "probe_fp8_backend", return_value=None):
bundle = run(p, dims=(1, 64, 128))
assert bundle["quant_stats"]["n_rowwise_layers"] == 1
assert bundle["quant_stats"]["n_fp8_resident_layers"] == 0
cond = bundle["model"].model.prediction_head.cond_proj
assert not isinstance(cond, FP8Linear)
assert cond.weight.dtype == torch.bfloat16
assert torch.isfinite(cond.weight.data.float()).all()
def test_mixed_convrot_and_fp8_resident(self, tmp_path, stubbed_load_env):
"""ConvRot int8 + scalar fp8 in one file: both resident families
install; the family label stays convrot_int8 (it dominates)."""
run, holder = stubbed_load_env
p = _save_qcheckpoint(tmp_path / "mixfp8.safetensors", [
("model.language_model.layers.0.self_attn.q_proj",
"convrot", (64, 64)),
("model.prediction_head.cond_proj", "rowwise_fp8", (64, 64)),
])
bundle = run(p, dims=(1, 64, 128))
assert bundle["weight_family"] == "convrot_int8"
assert bundle["quant_stats"] == {
"n_resident_layers": 2,
"n_rowwise_layers": 0,
"n_fp8_resident_layers": 1,
}
q = bundle["model"].model.language_model.layers[0].self_attn.q_proj
assert isinstance(q, ConvRotInt8Linear)
cond = bundle["model"].model.prediction_head.cond_proj
assert isinstance(cond, FP8Linear)
def test_prepare_routes_resident_fp8_and_keeps_tensors(self):
sd = {
"l.weight": torch.randint(-100, 100, (4, 4)).to(
torch.float8_e4m3fn),
"l.weight_scale": torch.tensor(0.25),
"l.comfy_quant": torch.zeros(3, dtype=torch.uint8),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, in_features=4, out_features=4,
convrot=False, orig_dtype="torch.bfloat16",
rowwise_dtype=torch.float8_e4m3fn, resident_fp8=True,
)}
plan, n_rowwise, n_fp8 = _prepare_quantized_safetensors_load(sd, qmap)
assert n_rowwise == 0 and n_fp8 == 1 and set(plan) == {"l"}
# Storage stays fp8 + scale stays in the dict for assign.
assert sd["l.weight"].dtype == torch.float8_e4m3fn
assert "l.weight_scale" in sd
assert "l.comfy_quant" not in sd
def test_prepare_resident_fp8_rejects_wrong_storage_dtype(self):
sd = {
"l.weight": torch.zeros(4, 4, dtype=torch.int8),
"l.weight_scale": torch.tensor(0.25),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, in_features=4, out_features=4,
convrot=False, orig_dtype="torch.bfloat16",
rowwise_dtype=torch.float8_e4m3fn, resident_fp8=True,
)}
with pytest.raises(Exception, match="expected"):
_prepare_quantized_safetensors_load(sd, qmap)
def test_prepare_resident_fp8_rejects_per_row_scale(self):
sd = {
"l.weight": torch.randint(-100, 100, (4, 4)).to(
torch.float8_e4m3fn),
"l.weight_scale": torch.full((4, 1), 0.25),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=0, in_features=4, out_features=4,
convrot=False, orig_dtype="torch.bfloat16",
rowwise_dtype=torch.float8_e4m3fn, resident_fp8=True,
)}
with pytest.raises(Exception, match="not a per-tensor scalar"):
_prepare_quantized_safetensors_load(sd, qmap)
class TestTiedLmHeadGuard:
"""A quantized lm_head contradicts tie_word_embeddings (plan 2026-08-27)."""
def _cfg(self, tied):
class _Cfg:
pass
c = _Cfg()
c.tie_word_embeddings = tied
return c
def test_tied_config_with_quantized_lm_head_raises(self):
from ComfyUI_VibeVoice.modules.external_loader import (
_assert_lm_head_not_tied,
)
from ComfyUI_VibeVoice.modules.quant_common import QuantTargetMismatch
qmap = {"lm_head": QuantLayerInfo(prefix="lm_head", group_size=0)}
with pytest.raises(QuantTargetMismatch, match="tie_word_embeddings"):
_assert_lm_head_not_tied(self._cfg(True), qmap)
def test_untied_config_passes(self):
from ComfyUI_VibeVoice.modules.external_loader import (
_assert_lm_head_not_tied,
)
qmap = {"lm_head": QuantLayerInfo(prefix="lm_head", group_size=0)}
_assert_lm_head_not_tied(self._cfg(False), qmap) # no raise
def test_tied_config_without_lm_head_quant_passes(self):
from ComfyUI_VibeVoice.modules.external_loader import (
_assert_lm_head_not_tied,
)
qmap = {"model.layers.0.mlp": QuantLayerInfo(
prefix="model.layers.0.mlp", group_size=0)}
_assert_lm_head_not_tied(self._cfg(True), qmap) # no raise
def test_magicmock_config_not_treated_as_tied(self):
"""Test doubles must not read as tied (non-bool flags ignored)."""
from ComfyUI_VibeVoice.modules.external_loader import (
_assert_lm_head_not_tied,
)
qmap = {"lm_head": QuantLayerInfo(prefix="lm_head", group_size=0)}
_assert_lm_head_not_tied(MagicMock(), qmap) # no raise
def test_loader_rejects_tied_fp8_lm_head_end_to_end(self, tmp_path,
stubbed_load_env):
run, holder = stubbed_load_env
p = _save_qcheckpoint(tmp_path / "tiedlm.safetensors", [
("lm_head", "rowwise_fp8", (96, 64)),
])
cfg = MagicMock()
cfg.tie_word_embeddings = True
cfg.decoder_config = None
holder["config"] = cfg
with pytest.raises(RuntimeError, match="tie_word_embeddings"):
run(p, dims=(1, 64, 128))
def test_unknown_format_propagates_not_swallowed(self, tmp_path,
stubbed_load_env):
"""Unsupported formats must HARD-FAIL the node, never fall back to
the dense loader (which would silently misload fp8/int8 weights)."""
run, holder = stubbed_load_env
p = _save_qcheckpoint(tmp_path / "mystery.safetensors", [])
from safetensors.torch import save_file
tensors = {
"model.prediction_head.cond_proj.weight": torch.zeros(64, 64),
"model.prediction_head.cond_proj.weight_scale": torch.ones(64, 1),
}
tensors["model.prediction_head.cond_proj.comfy_quant"] = (
torch.frombuffer(
bytearray(_json.dumps({"format": "brand_new_scheme"})
.encode("utf-8")), dtype=torch.uint8)
)
save_file(tensors, str(p))
with pytest.raises(RuntimeError, match="brand_new_scheme"):
run(p, dims=(1, 64, 128))
def test_dense_net_rejects_unplanned_int_weights(self, tmp_path,
stubbed_load_env):
"""A 'quantized' file WITHOUT usable metadata must fail loudly at the
dense gate instead of crashing inside torch or corrupting output."""
run, holder = stubbed_load_env
from safetensors.torch import save_file
p = tmp_path / "naive.safetensors"
save_file({
"model.prediction_head.cond_proj.weight": torch.zeros(
64, 64, dtype=torch.int8),
"model.prediction_head.cond_proj.bias": torch.zeros(64),
}, str(p))
# Loader wraps load errors as RuntimeError; the actionable message
# from the dense gate must survive the wrapping.
with pytest.raises(RuntimeError, match="quantized-weight tensors"):
run(p, dims=(1, 64, 128))
class TestBlockwiseLoading:
"""int8_blockwise checkpoints (unrotated, [out, in/gs] scales)."""
def test_scanner_parses_blockwise(self, tmp_path):
from ComfyUI_VibeVoice.modules.convrot_quant import (
scan_checkpoint_quantization,
)
p = _save_qcheckpoint(tmp_path / "blk.safetensors", [
("model.prediction_head.cond_proj", "block", (64, 64)),
])
qmap = scan_checkpoint_quantization(p)
info = qmap["model.prediction_head.cond_proj"]
assert info.convrot is False
assert info.group_size == 16
assert (info.out_features, info.in_features) == (64, 64)
def test_blockwise_dequant_math(self):
# gs=2 -> scale grid (out/2, in/2), one scale per 2x2 block.
sd = {
"l.weight": (torch.arange(16, dtype=torch.int8)
.reshape(4, 4) - 8),
"l.weight_scale": torch.tensor([[0.5, 0.25], [2.0, 1.5]]),
}
qmap = {"l": QuantLayerInfo(
prefix="l", group_size=2, in_features=4, out_features=4,
convrot=False, orig_dtype="torch.float32",
)}
plan, n_rowwise, n_fp8 = _prepare_quantized_safetensors_load(sd, qmap)
assert n_rowwise == 1 and plan == {}
# Block layout: rows 0-1 scaled by s[:,0] column-group 0, etc.
expected = torch.tensor([
[-8 * 0.5, -7 * 0.5, -6 * 0.25, -5 * 0.25],
[-4 * 0.5, -3 * 0.5, -2 * 0.25, -1 * 0.25],
[-0 * 2.0, 1 * 2.0, 2 * 1.5, 3 * 1.5],
[4 * 2.0, 5 * 2.0, 6 * 1.5, 7 * 1.5],
])
assert torch.equal(sd["l.weight"].float(), expected)
assert sd["l.weight"].dtype == torch.float32
assert "l.weight_scale" not in sd
def test_blockwise_end_to_end_through_loader(self, tmp_path,
stubbed_load_env):
"""File-C layout: whole checkpoint int8_blockwise incl. an
Embedding-shaped target — dequant-at-load handles any module kind."""
run, holder = stubbed_load_env
# cond_proj is Linear; use a big 'embedding-like' target too by
# pointing at the same Linear tree (module type irrelevant for the
# dequant path).
p = _save_qcheckpoint(tmp_path / "blockmodel.safetensors", [
("model.prediction_head.cond_proj", "block", (64, 64)),
("lm_head.weight-placeholder", "block", (64, 64)),
][0:1])
bundle = run(p, dims=(1, 64, 128))
assert bundle["quant_stats"] == {
"n_resident_layers": 0,
"n_rowwise_layers": 1,
"n_fp8_resident_layers": 0,
}
cond = bundle["model"].model.prediction_head.cond_proj
assert isinstance(cond, torch.nn.Linear)
assert not isinstance(cond, ConvRotInt8Linear)
assert cond.weight.dtype == torch.bfloat16
assert torch.isfinite(cond.weight.data.float()).all()
def test_rotated_embedding_fails_with_reexport_guidance(
self, tmp_path, stubbed_load_env):
"""File-A/B layout: convrot-flagged embed_tokens must fail with an
actionable message, never a silent misload."""
run, holder = stubbed_load_env
from safetensors.torch import save_file
tensors = {
"model.language_model.embed_tokens.weight": torch.randint(
-100, 100, (96, 64), dtype=torch.int8),
"model.language_model.embed_tokens.weight_scale":
torch.rand((96, 1)) * 0.01,
"model.language_model.embed_tokens.comfy_quant":
torch.frombuffer(bytearray(_json.dumps({
"convrot": True, "convrot_groupsize": 16,
"format": "int8_tensorwise",
"orig_dtype": "torch.bfloat16", "per_row": True,
}).encode("utf-8")), dtype=torch.uint8),
}
p = tmp_path / "rotembed.safetensors"
save_file(tensors, str(p))
with pytest.raises(RuntimeError, match="[Rr]e-export"):
run(p, dims=(1, 64, 128))
def test_plain_rowwise_embedding_dequants_fine(self, tmp_path,
stubbed_load_env):
"""Non-Linear targets are fine when NOT rotated (plain rowwise)."""
run, holder = stubbed_load_env
class _EmbHolder(torch.nn.Module):
def __init__(self, config):
super().__init__()
vv = build_stub_vv(n_layers=1)
self.model = vv.model
self.lm_head = vv.lm_head
self.embed_extra = torch.nn.Embedding(96, 64)
# quick sanity of the prep path on an Embedding module directly:
sd = {
"emb.weight": torch.randint(-50, 50, (96, 64), dtype=torch.int8),
"emb.weight_scale": torch.full((96, 1), 0.25),
}
qmap = {"emb": QuantLayerInfo(
prefix="emb", group_size=0, in_features=64, out_features=96,
convrot=False, orig_dtype="torch.bfloat16",
)}
plan, n_rowwise, n_fp8 = _prepare_quantized_safetensors_load(sd, qmap)
assert n_rowwise == 1 and plan == {}
assert sd["emb.weight"].dtype == torch.bfloat16
assert torch.isfinite(sd["emb.weight"].float()).all()
class TestOptionalAbsentPrefixes:
"""A released checkpoint may omit a whole subtree; that is not a failure.
VibeVoice-Realtime-0.5B ships an acoustic-tokenizer decoder only, and
vanilla ``from_pretrained`` reports the same 276 encoder keys as MISSING.
Our loader must not print an alarming warning for that, while still
reporting genuinely missing keys.
"""
def test_prefix_is_declared(self):
from ComfyUI_VibeVoice.modules.loader import OPTIONAL_ABSENT_PREFIXES
assert "acoustic_tokenizer.encoder." in OPTIONAL_ABSENT_PREFIXES
def test_wholly_omitted_prefix_is_suppressed(self):
from ComfyUI_VibeVoice.modules.loader import mark_optional_absent
missing = [
"acoustic_tokenizer.encoder.stages.0.0.weight",
"acoustic_tokenizer.encoder.head.conv.bias",
]
assigned = {"acoustic_tokenizer.decoder.head.weight", "model.language_model.x"}
known = mark_optional_absent(missing, assigned)
assert known == set(missing)
def test_partially_supplied_prefix_is_still_reported(self):
from ComfyUI_VibeVoice.modules.loader import mark_optional_absent
missing = ["acoustic_tokenizer.encoder.stages.9.9.weight"]
assigned = {"acoustic_tokenizer.encoder.stages.0.0.weight"}
known = mark_optional_absent(missing, assigned)
assert known == set()
def test_unrelated_missing_keys_are_untouched(self):
from ComfyUI_VibeVoice.modules.loader import mark_optional_absent
missing = ["model.language_model.layers.0.weight", "tts_eos_classifier.fc1.bias"]
assigned = set()
known = mark_optional_absent(missing, assigned)
assert known == set()
def test_existing_known_missing_is_preserved(self):
from ComfyUI_VibeVoice.modules.loader import mark_optional_absent
missing = ["acoustic_tokenizer.encoder.head.conv.bias"]
known = mark_optional_absent(missing, set(), {"lm_head.weight"})
assert "lm_head.weight" in known
assert "acoustic_tokenizer.encoder.head.conv.bias" in known
def test_caller_set_is_not_mutated(self):
from ComfyUI_VibeVoice.modules.loader import mark_optional_absent
original = {"lm_head.weight"}
mark_optional_absent(["acoustic_tokenizer.encoder.head.bias"], set(), original)
assert original == {"lm_head.weight"}
class _QuantGuardModel(torch.nn.Module):
"""Tiny real module for the streaming quant-storage guard tests."""
def __init__(self):
super().__init__()
self.layer1 = torch.nn.Linear(2, 2, bias=False)
self.layer2 = torch.nn.Linear(2, 2, bias=False)
self.config = MagicMock()
self.config.decoder_config.tie_word_embeddings = False
self.config.tie_word_embeddings = False
class TestStreamApplyDenseQuantGuard:
"""The streaming dense route must reject quantized storages too.
``external_loader._assert_dense_loadable`` guards the batch route; without
a matching check the per-tensor route would happily assign an int8 storage
into a float parameter. One shared dtype set keeps both routes speaking
the same message.
"""
@staticmethod
def _apply(pairs):
from ComfyUI_VibeVoice.modules.loader import VibeVoiceLoader
return VibeVoiceLoader._stream_apply_dense(_QuantGuardModel(), iter(pairs))
def test_int8_storage_rejected(self):
with pytest.raises(ValueError) as exc:
self._apply([("layer1.weight", torch.ones(2, 2, dtype=torch.int8))])
assert "carries no executable quantization metadata" in str(exc.value)
assert "layer1.weight" in str(exc.value)
assert "torch.int8" in str(exc.value)
def test_uint8_storage_rejected(self):
with pytest.raises(ValueError) as exc:
self._apply([("layer1.weight", torch.ones(2, 2, dtype=torch.uint8))])
assert "carries no executable quantization metadata" in str(exc.value)
assert "torch.uint8" in str(exc.value)
def test_float8_e4m3fn_storage_rejected(self):
with pytest.raises(ValueError) as exc:
self._apply([("layer1.weight", torch.ones(2, 2).to(torch.float8_e4m3fn))])
assert "carries no executable quantization metadata" in str(exc.value)
assert "torch.float8_e4m3fn" in str(exc.value)
def test_quant_storage_rejected_under_an_unexpected_key(self):
"""Guard sits at the top of the loop, before the params lookup.
``_assert_dense_loadable`` checks every key of the state dict, not
just the ones the model has a target for; an int8 tensor under an
unknown key must not slip through silently.
"""
with pytest.raises(ValueError) as exc:
self._apply([("no.such.key", torch.ones(2, 2, dtype=torch.int8))])
assert "no.such.key" in str(exc.value)
def test_bf16_and_fp32_dense_still_assign(self):
from ComfyUI_VibeVoice.modules.loader import VibeVoiceLoader
w1 = torch.full((2, 2), 3.0, dtype=torch.bfloat16)
w2 = torch.full((2, 2), 4.0, dtype=torch.float32)
model = _QuantGuardModel()
missing, unexpected = VibeVoiceLoader._stream_apply_dense(
model, iter([("layer1.weight", w1), ("layer2.weight", w2)])
)
assert missing == []
assert unexpected == []
assert model.layer1.weight.dtype == torch.bfloat16
assert model.layer2.weight.dtype == torch.float32
assert torch.equal(model.layer1.weight.data, w1)
assert torch.equal(model.layer2.weight.data, w2)
def test_dtype_set_is_shared_with_external_loader(self):
# Drift guard: the dtype set has exactly ONE definition object,
# module-level in ``modules.loader`` (external_loader imports it from
# there, so the dependency runs one way only).
from ComfyUI_VibeVoice.modules import loader as L
assert L.QUANT_STORAGE_DTYPES is EL._QUANT_STORAGE_DTYPES
assert L.QUANT_STORAGE_DTYPES == frozenset({
torch.int8, torch.uint8, torch.float8_e4m3fn, torch.float8_e5m2,
})