Files
WildAi 4d8686f9ad test: clean the suite to match the current architecture, and de-localise it
Suite goes from 1894 passed / 66 failed to 1876 passed / 0 failed.

The 66 failures were not 66 problems. Two mechanisms caused all of them:

  * comfy_stream._STREAMING_CONVERSION_ENABLED is False and nothing in
    production sets it, so convert_tree_for_streaming is a no-op. Every
    test asserting "this module was converted" was asserting on a sweep
    that never ran.
  * select_patcher_class ignores both its arguments and returns
    legacy_cls. Every test parametrized over a "dynamic" arm was
    exercising a configuration production cannot produce.

Deleted (assert wiring that no longer exists): test_dynamic_patcher_
selection.py, test_dynamic_vram_mechanism.py, test_ram_residual_
attribution.py, test_loader_streaming_integration.py (three of its five
tests were passing VACUOUSLY once conversion went), and
TestCensusOnRealDynamicRoute.

Corrected to the real contract, keeping the coverage: the conversion
tests now enable the gate explicitly and test the capability, which is
what they were written for; the TestP1 load tests now assert that the
LOADER places each tensor (it took a target_device) instead of asserting
the superseded CPU staging; the both-patcher-classes tests were
de-parametrized onto the selector's real answer.

Two production error messages had lost the hints their tests pinned, and
the tests were right: UnmappedKeyError no longer told a user with a
non-standard converter to place a sidecar config, and the GGUF block-size
error no longer said the layout could not be recovered.

Also fixes a latent bug: test_comfy_stream.py imported `modules.comfy_stream`
while everything else imported `ComfyUI_VibeVoice.modules.comfy_stream`.
Under pytest those are two distinct module objects with independent
globals, so a gate set through one was invisible to the other. All bare
`modules.*` imports in collected tests now use the package name.

Portability, verified by simulation rather than inspection:

  * conftest no longer falls back to a hardcoded C:\_Dev ComfyUI path. It
    used to put a nonexistent directory on sys.path off-box and surface
    as an opaque import error across 14 test files; now it names
    COMFYUI_ROOT and stops.
  * The 5.4 GB and 3.2 GB checkpoint fixtures are opt-in via env vars with
    no machine-specific default (VIBEVOICE_TEST_DENSE_CHECKPOINT,
    VIBEVOICE_TEST_GGUF).
  * comfy_kitchen guards added. It is genuinely optional; a driver
    script that blocks the import went from 9 failed to 37 passed /
    11 skipped / 0 errors.
  * comfy_aimdo import guarded, though it is in practice a hard
    dependency of the ComfyUI fork (core imports it unconditionally).
  * test_audio_backend writes to tmp_path instead of into the repo tree.
  * test_asr_loader no longer registers the real folder_paths.models_dir
    in an autouse fixture without restoring it.

The large-checkpoint and .dev-symlink tests skip cleanly off-box; that is
the intended shape for files that cannot ship with the repo.
2026-10-01 15:49:15 +03:00

380 lines
15 KiB
Python

"""Unit tests for the FP8-resident linear runtime (plan 2026-08-27, Phase 1).
Parity references are the MANUAL dequant formula the kitchen eager backend
implements — ``(w.float() * scale).to(out_dtype)`` — verified bit-exact
against ``comfy_kitchen.dequantize_per_tensor_fp8`` on this box. All tests
run on CPU (eager backend); no GPU required.
"""
import importlib.util
import json
import sys
import types
from unittest.mock import patch
import pytest
import torch
from torch import nn
from torch.nn import functional as F
from ComfyUI_VibeVoice.modules.fp8_quant import (
FP8Linear,
make_fp8_linear,
probe_fp8_backend,
)
from ComfyUI_VibeVoice.modules.convrot_quant import QuantLayerInfo
from ComfyUI_VibeVoice.modules.dtype_utils import (
cast_model_to_dtype,
cast_model_to_dtype_if_needed,
)
# comfy-kitchen provides the fp8 dequant kernels FP8Linear.forward calls. It is
# in neither requirements.txt nor pyproject.toml, so a machine without it must
# SKIP these rather than error: the fp8 path is simply unavailable there.
#
# find_spec is wrapped because it can raise on a broken or partially-installed
# package, and this runs at COLLECTION time -- an exception here would error the
# whole file instead of skipping the fp8 tests.
def _has_comfy_kitchen() -> bool:
try:
return importlib.util.find_spec("comfy_kitchen") is not None
except (ImportError, ValueError):
return False
requires_kitchen = pytest.mark.skipif(
not _has_comfy_kitchen(),
reason="comfy_kitchen (optional fp8 dequant backend) is not installed",
)
def _fill(m, seed=0, scale=2.0, bias_val=None):
"""Install deterministic fp8 storage + scalar scale into a fresh module."""
g = torch.Generator().manual_seed(seed)
w = (torch.randn(m.out_features, m.in_features, generator=g) * 0.3).to(m.fp8_dtype)
m.weight.data.copy_(w)
m.weight_scale.data.copy_(torch.tensor(scale, dtype=torch.float32))
if m.bias is not None:
b = torch.randn(m.out_features, generator=g) * 0.1
if bias_val is not None:
b = torch.full((m.out_features,), bias_val)
m.bias.data.copy_(b)
return m
def _reference(m, x):
ref_w = (m.weight.float() * m.weight_scale).to(x.dtype)
bias = m.bias.to(x.dtype) if m.bias is not None else None
return F.linear(x, ref_w, bias)
@requires_kitchen
class TestProbeFP8Backend:
def test_returns_backend_on_this_box(self):
assert probe_fp8_backend() in ("triton", "cuda", "eager")
def test_none_when_kitchen_missing(self):
with patch.dict(sys.modules, {"comfy_kitchen": None}):
assert probe_fp8_backend() is None
def test_none_when_callable_missing_despite_capability(self):
# D-3 trap guard: capability strings have shipped without callables.
fake = types.ModuleType("comfy_kitchen")
fake.list_backends = lambda: {
"eager": {"available": True,
"capabilities": ["dequantize_per_tensor_fp8"]}
}
with patch.dict(sys.modules, {"comfy_kitchen": fake}):
assert probe_fp8_backend() is None
def test_none_when_no_backend_available(self):
# comfy_kitchen is in neither requirements.txt nor pyproject.toml, so a
# bare import here ERRORS on a machine without it instead of skipping.
comfy_kitchen = pytest.importorskip("comfy_kitchen")
fake = types.ModuleType("comfy_kitchen")
fake.dequantize_per_tensor_fp8 = comfy_kitchen.dequantize_per_tensor_fp8
fake.list_backends = lambda: {
"eager": {"available": False,
"capabilities": ["dequantize_per_tensor_fp8"]}
}
with patch.dict(sys.modules, {"comfy_kitchen": fake}):
assert probe_fp8_backend() is None
class TestFP8LinearConstruction:
def test_rejects_non_fp8_dtype(self):
with pytest.raises(ValueError):
FP8Linear(4, 4, False, torch.float16)
@pytest.mark.parametrize("fp8_dtype", [torch.float8_e4m3fn, torch.float8_e5m2])
def test_storage_dtypes_and_markers(self, fp8_dtype):
m = FP8Linear(6, 8, True, fp8_dtype)
assert m.weight.dtype == fp8_dtype
assert m.weight_scale.dtype == torch.float32
assert m.weight_scale.shape == ()
assert m.bias is not None and m.bias.dtype == torch.float32
assert m._quant_resident is True
assert m.comfy_cast_weights is True
assert m.weight_comfy_model_dtype == fp8_dtype
def test_no_bias(self):
m = FP8Linear(6, 8, False, torch.float8_e4m3fn)
assert m.bias is None
def test_meta_init_compatible(self):
with torch.device("meta"):
m = FP8Linear(6, 8, False, torch.float8_e4m3fn)
assert m.weight.is_meta
assert m.weight_scale.is_meta
assert tuple(m.weight.shape) == (8, 6)
@requires_kitchen
class TestFP8LinearForward:
@pytest.mark.parametrize("act_dtype", [torch.bfloat16, torch.float32])
def test_parity_bit_exact(self, act_dtype):
m = _fill(FP8Linear(6, 8, False, torch.float8_e4m3fn), seed=1)
x = torch.randn(3, 6, dtype=act_dtype)
out = m(x)
assert out.dtype == act_dtype
assert torch.equal(out, _reference(m, x))
def test_parity_e5m2(self):
m = _fill(FP8Linear(6, 8, False, torch.float8_e5m2), seed=2)
x = torch.randn(3, 6, dtype=torch.bfloat16)
assert torch.equal(m(x), _reference(m, x))
def test_parity_with_bias(self):
m = _fill(FP8Linear(6, 8, True, torch.float8_e4m3fn), seed=3)
x = torch.randn(2, 5, 6, dtype=torch.bfloat16) # batched 3-D input
assert torch.equal(m(x), _reference(m, x))
def test_non_float_activation_raises(self):
m = _fill(FP8Linear(6, 8, False, torch.float8_e4m3fn))
with pytest.raises(TypeError):
m(torch.randint(0, 4, (2, 6)))
def test_streamed_path_parity(self):
"""A non-empty weight_function forces _forward_streamed; the pulled
weight must stay fp8 (dtype-preserving) and the result must match."""
m = _fill(FP8Linear(6, 8, True, torch.float8_e4m3fn), seed=4)
seen = []
def _identity(w):
seen.append(w.dtype)
return w
m.weight_function = [_identity]
x = torch.randn(3, 6, dtype=torch.bfloat16)
assert torch.equal(m(x), _reference(m, x))
assert seen == [torch.float8_e4m3fn], "streamed pull must preserve fp8 dtype"
m.weight_function = []
class TestResidentConstructionIsMetaOnly:
"""2026-09-30: swapping in quant residents must not allocate host RAM.
The model tree is meta (VibeVoiceLoader._instantiate_model use_meta=True),
but the replacement factories used to run OUTSIDE that context, so every
FP8Linear allocated a real torch.empty of its full weight — 8.08 GB for
the 7B fp8 checkpoint, all of it replaced by the checkpoint assign. The
live signature was peak_ws 5.6 GB next to peak_private 27.94 GB:
committed, never touched.
"""
def test_swapped_residents_start_on_meta(self):
from ComfyUI_VibeVoice.modules.convrot_quant import QuantLayerInfo
from ComfyUI_VibeVoice.modules.fp8_quant import make_fp8_linear
from ComfyUI_VibeVoice.modules.quant_common import replace_linears_for_quant
class Tree(torch.nn.Module):
def __init__(self):
super().__init__()
self.proj = torch.nn.Linear(32, 16)
info = QuantLayerInfo(
prefix="proj",
group_size=32,
in_features=32,
out_features=16,
has_bias=True,
rowwise_dtype=torch.float8_e4m3fn,
resident_fp8=True,
)
with torch.device("meta"):
tree = Tree()
replaced = replace_linears_for_quant(tree, {"proj": make_fp8_linear(info)})
assert replaced == ["proj"]
weight = dict(tree.named_parameters())["proj.weight"]
assert weight.is_meta, (
"the replacement resident allocated real host memory — this is "
"the 8GB-per-model waste the meta-only construction removes"
)
class TestFP8LinearStateDict:
def test_load_consumes_comfy_quant_meta(self):
m = FP8Linear(6, 8, False, torch.float8_e4m3fn)
g = torch.Generator().manual_seed(5)
w = (torch.randn(8, 6, generator=g) * 0.3).to(torch.float8_e4m3fn)
meta = json.dumps(
{"format": "float8_e4m3fn", "orig_dtype": "torch.bfloat16"}
).encode()
sd = {
"weight": w,
"weight_scale": torch.tensor(1.5),
"comfy_quant": torch.frombuffer(bytearray(meta), dtype=torch.uint8),
}
missing, unexpected = m.load_state_dict(sd, strict=False)
assert list(missing) == []
assert list(unexpected) == []
assert torch.equal(m.weight, w)
assert m.weight_scale.item() == pytest.approx(1.5)
def test_state_dict_roundtrip_keeps_storage_dtypes(self):
m = _fill(FP8Linear(6, 8, True, torch.float8_e4m3fn), seed=6)
sd = m.state_dict()
assert set(sd.keys()) == {"weight", "weight_scale", "bias"}
assert sd["weight"].dtype == torch.float8_e4m3fn
assert sd["weight_scale"].dtype == torch.float32
class TestMakeFP8Linear:
@pytest.mark.parametrize("fp8_dtype", [torch.float8_e4m3fn, torch.float8_e5m2])
def test_factory_builds_matching_module(self, fp8_dtype):
info = QuantLayerInfo(
prefix="model.language_model.layers.0.self_attn.q_proj",
group_size=0,
in_features=6,
out_features=8,
convrot=False,
rowwise_dtype=fp8_dtype,
)
mod = make_fp8_linear(info)(6, 8, True)
assert isinstance(mod, FP8Linear)
assert mod.fp8_dtype == fp8_dtype
assert mod.bias is not None
# ====================================================================
# Ranked hypothesis (task fp8-peak, 2026-09-29): "_quant_protected_names
# fails to protect FP8Linear, so cast_model_to_dtype_if_needed dequantises
# the resident fp8 weights to bf16 — 9.47GB of fp8 becoming ~17GB, the
# measured 1.8x on the 7B file."
#
# MEASURED, 2026-09-29: REFUTED. The ranked hypothesis was the whole defect
# story and it does not hold. ``_quant_protected_names`` protects the fp8
# weight and the fp32 scale of every ``_quant_resident`` module
# (modules/dtype_utils.py:138-152), and FP8Linear sets that marker
# (modules/fp8_quant.py:93), so the cast skips the residents. These tests
# exist to keep it refuted: the protection is load-bearing for a 1.8x host-RAM
# claim and a future edit that drops the marker (or the weight entry) would
# silently reintroduce it, with no other test noticing.
# ====================================================================
class _WitnessModel(nn.Module):
"""fp8 resident + one mismatched fp32 linear, so the cast is NOT a no-op.
The plain fp32 linear is the control: it proves the cast actually ran.
Without a mismatched castable parameter, ``cast_model_to_dtype_if_needed``
takes its fast path and returns untouched — which would make every
"storage survived" assertion below pass vacuously.
"""
def __init__(self, in_features=8, out_features=4):
super().__init__()
self.proj = FP8Linear(in_features, out_features, True,
torch.float8_e4m3fn, torch.bfloat16)
self.head = nn.Linear(out_features, out_features, bias=False)
self.head.weight.data = torch.zeros(out_features, out_features,
dtype=torch.float32)
def _fp8_resident_tree():
model = _WitnessModel()
_fill(model.proj, seed=3, scale=0.75)
return model
class TestFp8ResidentSurvivesDtypeCast:
"""KB-scale: a resident fp8 param keeps float8 STORAGE after the cast."""
def test_cast_runs_and_leaves_fp8_storage_untouched(self):
model = _fp8_resident_tree()
before_bytes = model.proj.weight.untyped_storage().nbytes()
assert before_bytes == 4 * 8 # 1 byte per fp8 element, no rounding
cast_model_to_dtype_if_needed(model, torch.bfloat16)
# Control: the cast really happened (else the rest proves nothing).
assert model.head.weight.dtype == torch.bfloat16
# The claim under test.
assert model.proj.weight.dtype == torch.float8_e4m3fn
assert model.proj.weight_scale.dtype == torch.float32
assert model.proj.weight.untyped_storage().nbytes() == before_bytes
assert model.proj.weight_scale.item() == pytest.approx(0.75)
@pytest.mark.parametrize("fp8_dtype", [torch.float8_e4m3fn, torch.float8_e5m2])
def test_both_fp8_storage_dtypes_survive(self, fp8_dtype):
model = _WitnessModel()
model.proj = FP8Linear(8, 4, True, fp8_dtype, torch.bfloat16)
_fill(model.proj, seed=4)
cast_model_to_dtype_if_needed(model, torch.bfloat16)
assert model.proj.weight.dtype == fp8_dtype
assert model.proj.weight_scale.dtype == torch.float32
def test_unconditional_cast_model_to_dtype_also_protects(self):
"""``cast_model_to_dtype`` shares the filtered walk, so it is safe too."""
model = _fp8_resident_tree()
cast_model_to_dtype(model, torch.bfloat16)
assert model.proj.weight.dtype == torch.float8_e4m3fn
def test_fp8_bytes_stay_half_the_bf16_footprint(self):
"""The 1.8x hypothesis, in bytes: fp8 storage is HALF a bf16 recast.
9.47GB of fp8 is ~17GB if the residents are cast — the exact shape of
the reported 25->42GB spike. With the protection in place the model
keeps the fp8 footprint, so the ceiling of the quant route is ~1x
file (owned copies), not ~2x.
"""
model = _fp8_resident_tree()
fp8_bytes = model.proj.weight.numel()
cast_model_to_dtype_if_needed(model, torch.bfloat16)
assert model.proj.weight.untyped_storage().nbytes() == fp8_bytes
assert fp8_bytes * 2 == model.proj.weight.numel() * 2 # bf16 would be this
def test_census_reads_the_resident_as_private_fp8(self):
"""The census half of the same claim, at the same scale.
A quant-assigned parameter is a PRIVATE host allocation BY DESIGN
(the stream assigns owned clones — under the aimdo arm they are
cloned from zero-copy file views, see modules/base_loader.py) —
so the census must count it as private
fp8 bytes at 1 byte/element, never as an aimdo file view. That 1x is
the ceiling of the quant route: with views the family would read as
page cache, and with a bf16 recast it would read as twice these bytes.
"""
from ComfyUI_VibeVoice.modules.memory_census import census
model = _fp8_resident_tree()
cast_model_to_dtype_if_needed(model, torch.bfloat16)
report = census(model)
# No view, no mmap: the whole tree is private host memory.
assert report["param_view_bytes"] == 0
assert report["param_mmap_bytes"] == 0
assert report["param_private_bytes"] == sum(
p.untyped_storage().nbytes() for p in model.parameters())
# The resident weight is counted at fp8 width (1 byte/element) plus
# its fp32 scalar scale and its (castable, so cast) bf16 bias.
assert report["families"]["FP8Linear/params"] == (
model.proj.weight.numel()
+ model.proj.weight_scale.untyped_storage().nbytes()
+ model.proj.bias.untyped_storage().nbytes()
)
assert report["families"]["Linear/params"] == model.head.weight.numel() * 2