Files
WildAi 8389ab4ce6 refactor: log through ComfyUI's own root-logger idiom, re-audit every level
The package installed its own handler, Formatter and level on a
module-level logger, then logged through that logger. ComfyUI core already
installs app.logger.ColoredFormatter on the ROOT logger, so a bare
logging.<level>() call whose message starts with "[ComfyUI-VibeVoice] " is
tagged and coloured for free. This removes the duplicate machinery rather
than extending it -- no vv_logging module, Logger subclass, LoggerAdapter,
vendored ANSI table, new handler, setLevel or propagate.

Three hazards that motivated deleting the setup block rather than moving it:

- propagate=False severed pytest's caplog; the tests only passed because
  __init__.py's `if "pytest" in sys.modules` guard skipped the block.
  Bare root calls propagate by default.
- logger.setLevel(INFO) on the package logger pinned every descendant, so
  ComfyUI's --verbose never applied to this node. Deleting the block fixes it.
- The level must stay a literal at the call site; no env-var switch was
  added. modules/diagnostics.py gates diagnostic CONTENT and is untouched.

Every call site (235 across 34 files) is now a bare root call with the
prefix. The vendored src/vibevoice tree used transformers.utils.logging,
which has no module-level info/warning/debug; those files now import stdlib
logging, so a new call site there would fail loudly instead of silently.

Level audit against ERROR=failure / WARNING=degradation / INFO=user-facing
progress / DEBUG=internals. Deleted as noise: the resample notice (44.1 kHz
reference audio is the normal case) and the two "Successfully loaded
external VibeVoice" confirmations, which duplicated the patcher's load line.
Promoted to WARNING: a mixed-naming GGUF, which is resolved by heuristic
majority vote and silently aliases the rest. Demoted to DEBUG: the four
SageAttention kernel-selection lines, model discovery, shard counts,
per-retry download attempts, attention-mode confirmation and the
save_pretrained notice. Kept at INFO: generation complete, transcription
results, model downloads and load starts.

Judgement calls recorded in docs/2026-10-01-vv-logging-cleanup-design.md.
The two I am least sure of: the ungated memory-census profile line was
demoted rather than gated (adding a gate would not be presentation-only),
and patcher.py's "Loading VibeVoice models for..." was kept at INFO against
the ask's example list because it is the only load-start line the package has.

48 new tests in tests/test_logging_idiom.py plus tests/test_audio_utils.py:
byte-exact ColoredFormatter rendering, no-leftover-machinery, AST prefix and
two-direction level policy, and the 44.1 kHz resample regression (the call
still fires with (44100, 24000); nothing is logged). 37 caplog.at_level
pins that named a module logger were stripped -- they only lowered a named
logger and left root at WARNING, so the record was discarded before capture.

Gate: 1924 passed, 30 skipped, 0 failed (1876 before this change). Twelve
deliberate mutations -- re-deleting a message, re-leveling, re-adding
setLevel and propagate=False, stripping a prefix, moving the resample out
of its branch -- were each caught by at least one test.

Not run: ComfyUI was never launched and no checkpoint was loaded.
2026-10-01 23:08:04 +03:00

179 lines
6.3 KiB
Python

"""FP8 rowwise checkpoint runtime via comfy-kitchen.
Executes safetensors checkpoints carrying ``*.comfy_quant`` metadata with a
rowwise fp8 format and a SCALAR per-tensor scale (produced by
comfy-model-tools):
weight : torch.float8_e4m3fn / float8_e5m2 [out_features, in_features]
weight_scale: float32 scalar
comfy_quant : uint8 JSON {"format": "float8_e4m3fn",
"orig_dtype": "torch.bfloat16"}
Weights stay FP8 resident in VRAM (1 byte/element); each forward dequantizes
the weight per-tensor through :func:`comfy_kitchen.dequantize_per_tensor_fp8`
and runs a plain ``F.linear``. LOAD-ONLY: this module executes existing
checkpoints, it does not requantize.
"""
from __future__ import annotations
import logging
import torch
from torch import nn
from torch.nn import functional as F
from .convrot_quant import (
QUANT_META_SUFFIX,
ROWWISE_FLOAT_FORMATS,
UnsupportedQuantFormat,
resolve_orig_dtype,
)
FP8_DEQUANT_CAPABILITY = "dequantize_per_tensor_fp8"
# Reverse lookup: torch fp8 dtype -> comfy_quant format name.
_FP8_FORMAT_NAMES = {v: k for k, v in ROWWISE_FLOAT_FORMATS.items()}
def probe_fp8_backend():
"""Return the comfy-kitchen backend providing per-tensor fp8 dequant.
Soft probe: returns None instead of raising if unavailable, falling back
to dequant-at-load when needed.
"""
try:
import comfy_kitchen
except ImportError:
return None
if not callable(getattr(comfy_kitchen, FP8_DEQUANT_CAPABILITY, None)):
return None
backends = comfy_kitchen.list_backends()
for name in ("triton", "cuda", "eager"):
info = backends.get(name)
if not info or not info.get("available"):
continue
if FP8_DEQUANT_CAPABILITY in set(info.get("capabilities") or ()):
return name
return None
class FP8Linear(nn.Module):
"""Drop-in nn.Linear replacement executing fp8 storage with a scalar scale.
The fp8 weight and the fp32 scalar scale stay resident at their storage
dtypes; dequantization happens per forward call into the activation
dtype, so a 7B fp8 model occupies ~1 byte/weight in VRAM instead of the
2 bytes/weight a dequant-at-load bf16 model needs.
"""
_quant_resident = True
comfy_cast_weights = True
weight_function = []
bias_function = []
def __init__(self, in_features: int, out_features: int, bias: bool,
fp8_dtype: torch.dtype, compute_dtype: torch.dtype = None):
super().__init__()
if fp8_dtype not in ROWWISE_FLOAT_FORMATS.values():
raise ValueError(
f"FP8Linear requires an fp8 dtype from "
f"{sorted(_FP8_FORMAT_NAMES.values())}, got {fp8_dtype}"
)
self.in_features = in_features
self.out_features = out_features
self.fp8_dtype = fp8_dtype
self.compute_dtype = compute_dtype
self.weight = nn.Parameter(
torch.empty(out_features, in_features, dtype=fp8_dtype),
requires_grad=False,
)
self.weight_scale = nn.Parameter(
torch.empty((), dtype=torch.float32), requires_grad=False
)
self.weight_comfy_model_dtype = fp8_dtype
self.bias = (
nn.Parameter(torch.empty(out_features), requires_grad=False)
if bias else None
)
self.quant_format = _FP8_FORMAT_NAMES[fp8_dtype]
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict,
missing_keys, unexpected_keys, error_msgs):
state_dict.pop(f"{prefix}{QUANT_META_SUFFIX}", None)
super()._load_from_state_dict(
state_dict, prefix, local_metadata, strict,
missing_keys, unexpected_keys, error_msgs,
)
def forward(self, x):
import comfy_kitchen
if not x.dtype.is_floating_point:
raise TypeError(
f"FP8Linear expects a floating-point activation, got {x.dtype}."
)
if x.dtype in _FP8_FORMAT_NAMES:
raise TypeError(
f"FP8Linear received fp8 activations ({x.dtype}). Upstream "
f"code derived an activation dtype from the quantized weight "
f"storage; use the module's compute_dtype "
f"({self.compute_dtype}) instead of weight.dtype."
)
wf = getattr(self, "weight_function", None)
if (self.weight.device != x.device) or (wf and len(wf) > 0):
return self._forward_streamed(x)
w = comfy_kitchen.dequantize_per_tensor_fp8(
self.weight, self.weight_scale, x.dtype
)
bias = self.bias
if bias is not None and bias.dtype != x.dtype:
bias = bias.to(x.dtype)
return F.linear(x, w, bias)
def _forward_streamed(self, x):
"""Paged path: pull the raw fp8 weight through core's cast machinery."""
import comfy.ops
import comfy_kitchen
weight, bias, offload_stream = comfy.ops.cast_bias_weight(
self,
x,
device=x.device,
dtype=self.weight_comfy_model_dtype,
bias_dtype=x.dtype,
offloadable=True,
)
try:
scale = (self.weight_scale if self.weight_scale.device == x.device
else self.weight_scale.to(x.device))
if bias is not None and bias.dtype != x.dtype:
bias = bias.to(x.dtype)
w = comfy_kitchen.dequantize_per_tensor_fp8(weight, scale, x.dtype)
return F.linear(x, w, bias)
finally:
comfy.ops.uncast_bias_weight(self, weight, bias, offload_stream)
def extra_repr(self):
return (
f"in_features={self.in_features}, out_features={self.out_features}, "
f"bias={self.bias is not None}, fp8_dtype={self.fp8_dtype}"
)
def make_fp8_linear(info):
"""Factory adapter for :func:`replace_linears_for_quant` plans."""
fp8_dtype = info.rowwise_dtype
try:
compute_dtype = resolve_orig_dtype(info.orig_dtype)
except UnsupportedQuantFormat:
compute_dtype = None
def _factory(in_features: int, out_features: int, has_bias: bool):
return FP8Linear(in_features, out_features, has_bias, fp8_dtype,
compute_dtype)
return _factory