Files
WildAi 10e70d8e25 perf: fp8-resident weights + streaming safetensors load (v2.8.0)
Loading VibeVoice-*-fp8_e4m3.safetensors materialized the full state
dict and dequantized all ~380 fp8 layers to bf16 in RAM: 18->43 GB
spike, Pin-error flood, ~4.5 GB partial offload on 16 GB cards.

- FP8Linear keeps fp8 weight + scalar fp32 scale resident in VRAM
  (~1 byte/weight); forward dequantizes per-tensor via comfy-kitchen
  into the activation dtype. Missing kitchen backend falls back to
  dequant-at-load.
- Quantized safetensors load streams per-tensor (never a full
  in-memory state dict); dense files keep the batch path.
- Non-Linear fp8 targets (real exports quantize embed_tokens) demote
  to dequant-at-load; ConvRot non-Linear targets keep the hard fail.
- Diffusion head resolves the timestep-mlp input dtype via the
  module's compute_dtype, never the fp8 storage dtype — casting
  activations to weight.dtype fed fp8 into the dequant kernel and
  crashed generate() (NoCapableBackendError).

GPU-gated on 1.5B + 7B fp8: full VRAM fit, no partial offload, no
pin flood, audio generates.
2026-08-28 01:24:26 +03:00

169 lines
6.2 KiB
Python

"""Regression test for the TimestepEmbedder dtype handling.
Runtime BUG-002: `TimestepEmbedder.timestep_embedding` returned
`embedding.to(t.dtype)`. When the timesteps tensor is integer (`int64` — which is
what `noise_scheduler.timesteps` yields), that (a) truncated the cos/sin embedding
to int64, destroying it, and (b) fed an int64 tensor into `self.mlp`, whose weights
are `c10::BFloat16` when the model is loaded in bf16 (e.g. under SageAttention) —
raising `RuntimeError: expected mat1 and mat2 to have the same dtype`.
The fix keeps the embedding floating and casts it to the MLP's compute dtype in
`TimestepEmbedder.forward`. This test loads the REAL module by file path (bypassing
the heavy `vibevoice` package __init__ chain, which would otherwise pull in
`diffusers`) and asserts the forward runs for int64/float32/bf16 timesteps and
produces a correct bf16 output that is not destroyed.
"""
import importlib.util
import os
import sys
import types
import pytest
torch = pytest.importorskip("torch")
_HERE = os.path.dirname(os.path.abspath(__file__))
_PATCH_PATH = os.path.abspath(
os.path.join(
_HERE, "..", "src", "vibevoice", "modular",
"modular_vibevoice_diffusion_head.py",
)
)
def _load_module():
# Stub the relative-import target so we don't pull in the heavy vibevoice
# package (which transitively imports diffusers and trips a huggingface_hub
# version clash in some venvs). We only need TimestepEmbedder here.
config_name = "src.vibevoice.modular.configuration_vibevoice"
stub = types.ModuleType(config_name)
stub.VibeVoiceDiffusionHeadConfig = type("VibeVoiceDiffusionHeadConfig", (), {})
sys.modules[config_name] = stub
spec = importlib.util.spec_from_file_location(
"src.vibevoice.modular.modular_vibevoice_diffusion_head", _PATCH_PATH
)
mod = importlib.util.module_from_spec(spec)
mod.__package__ = "src.vibevoice.modular"
spec.loader.exec_module(mod)
return mod
def test_timestep_embedder_handles_int64_timesteps_in_bf16():
mod = _load_module()
TimestepEmbedder = mod.TimestepEmbedder
embedder = TimestepEmbedder(hidden_size=64, frequency_embedding_size=256).to(
torch.bfloat16
)
# The exact failing case: integer scheduler timesteps + bf16 model weights.
t_int = torch.arange(0, 50, dtype=torch.int64)
out = embedder(t_int)
assert out.dtype == torch.bfloat16, f"output dtype was {out.dtype}"
assert out.shape == (50, 64)
# The embedding must not have been destroyed by an int downcast.
assert out.abs().max().item() > 0, "embedding was destroyed (all zeros)"
# Sanity: float inputs also work and stay bf16.
for t in (
torch.linspace(0, 1000, 50, dtype=torch.float32),
torch.linspace(0, 1000, 50, dtype=torch.bfloat16),
):
o = embedder(t)
assert o.dtype == torch.bfloat16
assert o.abs().max().item() > 0
# ====================================================================
# N-13 regression (2026-08-28): fp8 checkpoints replace the diffusion
# head's mlp linears with FP8Linear. `TimestepEmbedder.forward` used to
# cast `t_freq.to(self.mlp[0].weight.dtype)` — with fp8-resident weights
# that cast the activation to torch.float8_e4m3fn, and FP8Linear passed
# it straight into comfy_kitchen.dequantize_per_tensor_fp8 as the output
# dtype -> NoCapableBackendError. The forward must derive the compute
# dtype via _mlp_compute_dtype (declared compute_dtype first), never the
# storage dtype.
# ====================================================================
def test_mlp_compute_dtype_selection():
mod = _load_module()
# Plain float mlp: first layer's weight dtype (legacy behavior).
emb = mod.TimestepEmbedder(hidden_size=32).to(torch.bfloat16)
assert mod._mlp_compute_dtype(emb.mlp) == torch.bfloat16
# fp8 storage without a declared compute dtype: falls back to bf16,
# never the fp8 storage dtype.
class _FakeFp8Lin(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(
torch.empty(8, 8, dtype=torch.float8_e4m3fn)
)
self.compute_dtype = None
def forward(self, x):
return x
mlp = torch.nn.Sequential(_FakeFp8Lin(), torch.nn.SiLU())
assert mod._mlp_compute_dtype(mlp) == torch.bfloat16
# A declared compute dtype wins.
mlp[0].compute_dtype = torch.float16
assert mod._mlp_compute_dtype(mlp) == torch.float16
def test_timestep_embedder_with_fp8_mlp_stays_bf16():
from ComfyUI_VibeVoice.modules.convrot_quant import QuantLayerInfo
from ComfyUI_VibeVoice.modules.fp8_quant import (
make_fp8_linear,
probe_fp8_backend,
)
if probe_fp8_backend() is None:
pytest.skip("no comfy_kitchen fp8 backend on this box")
mod = _load_module()
embedder = mod.TimestepEmbedder(hidden_size=64, frequency_embedding_size=256)
def _fp8_linear(in_f, out_f):
info = QuantLayerInfo(
prefix="t_embedder.mlp",
group_size=0,
in_features=in_f,
out_features=out_f,
has_bias=False,
convrot=False,
orig_dtype="torch.bfloat16",
rowwise_dtype=torch.float8_e4m3fn,
resident_fp8=True,
)
lin = make_fp8_linear(info)(in_f, out_f, False)
g = torch.Generator().manual_seed(0)
w = torch.randint(
-100, 100, (out_f, in_f), generator=g
).to(torch.float8_e4m3fn)
lin.weight.data.copy_(w)
lin.weight_scale.data.fill_(0.01)
return lin
# Mirror the real fp8 checkpoint: both mlp linears replaced
# (bias=False, so the tree has NO bf16 parameter to fall back on).
embedder.mlp[0] = _fp8_linear(256, 64)
embedder.mlp[2] = _fp8_linear(64, 64)
assert mod._mlp_compute_dtype(embedder.mlp) == torch.bfloat16
# The exact failing case from the GPU gate: integer scheduler
# timesteps into the fp8-resident mlp.
t_int = torch.arange(0, 50, dtype=torch.int64)
out = embedder(t_int)
assert out.dtype == torch.bfloat16, f"output dtype was {out.dtype}"
assert out.shape == (50, 64)
assert torch.isfinite(out).all()
assert out.abs().max().item() > 0, "embedding was destroyed (all zeros)"