Files
WildAi cfa57c8738 feat: streaming quant load, single-rounding dequant, release hygiene (v2.11.0)
Loading
- GGUF install is now two passes: a metadata pass that decides each
  tensor's disposition, then an install pass that assigns residents in
  place and streams dense tensors. The whole checkpoint is no longer
  buffered in a dict alongside the model being built.
- dequantize_reader_tensor takes a target dtype, so dequant-at-load
  writes straight into the destination parameter and the fp32
  intermediate is never allocated.
- A bundle whose heavy fields were released is rebuilt from its recorded
  source_path instead of failing the consumer.
- Host memory is released after install.

Numerics
- Q8_0 dequant computes in fp32 so the result is rounded once, at the
  final cast. The activation-dtype path was reverted: it rounded twice
  and moved stored weights.
- Removed a redundant weight-sized copy from the dequant kernel.
  Bitwise-identical, ~1.16x.
- Precision gates are bitwise rather than tolerance-based.

Docs and tests
- README condensed; changelog moved to CHANGELOG.md.
- Third-party project references removed from source comments.
- Tests no longer assert README prose; the e2e smoke contract follows
  the developer script to its new location and skips when absent.
- Version guard reads CHANGELOG.md.
2026-09-28 18:25:19 +03:00

320 lines
12 KiB
Python

"""Regression test for the SageAttention RoPE call.
Runtime BUG-001: `sage_attention_forward` called
`apply_rotary_pos_emb(query_states, key_states, cos, sin, position_ids=None)`.
transformers >= 5 removed the `position_ids` kwarg from that function, so the
call raised `TypeError: apply_rotary_pos_emb() got an unexpected keyword
argument 'position_ids'` at generation time when the sage attention mode was
selected. The fix drops the kwarg and calls positionally.
This test loads the REAL module by file path (bypassing conftest's
`src.vibevoice` mock), stubs the CUDA sage kernel so it runs without a GPU,
and asserts the forward executes the RoPE call without TypeError and produces
the correct output shape. If the `position_ids` kwarg is ever reintroduced the
real `apply_rotary_pos_emb` raises and this test fails.
"""
import importlib.util
import os
import types
import pytest
torch = pytest.importorskip("torch")
pytest.importorskip("transformers")
_HERE = os.path.dirname(os.path.abspath(__file__))
_PATCH_PATH = os.path.abspath(
os.path.join(_HERE, "..", "src", "vibevoice", "modular", "sage_attention_patch.py")
)
spec = importlib.util.spec_from_file_location("sage_attention_patch_rt", _PATCH_PATH)
sage_attention_patch = importlib.util.module_from_spec(spec)
spec.loader.exec_module(sage_attention_patch)
from torch import nn # noqa: E402
def _make_fake_attention():
head_dim = 16
num_heads = 4
num_kv = 4
hidden = num_heads * head_dim # 64
cfg = types.SimpleNamespace(
num_attention_heads=num_heads,
num_key_value_heads=num_kv,
head_dim=head_dim,
hidden_size=hidden,
)
attn = types.SimpleNamespace(
config=cfg,
layer_idx=0,
head_dim=head_dim,
q_proj=nn.Linear(hidden, num_heads * head_dim),
k_proj=nn.Linear(hidden, num_kv * head_dim),
v_proj=nn.Linear(hidden, num_kv * head_dim),
o_proj=nn.Linear(hidden, hidden),
)
return attn, head_dim, num_heads, hidden
def test_sage_attention_forward_no_position_ids_kwarg():
fake, head_dim, num_heads, hidden = _make_fake_attention()
called = {}
def stub_sage(q, k, v, tensor_layout="HND", is_causal=False,
qk_quant_gran="per_warp", pv_accum_dtype="fp32"):
called["ok"] = True
# Mirror the expected output shape: (bsz, heads, q_len, head_dim)
return q
prev = sage_attention_patch.SAGE_ATTENTION_FUNCTION
sage_attention_patch.SAGE_ATTENTION_FUNCTION = stub_sage
try:
bsz, q_len = 1, 8
hidden_states = torch.randn(bsz, q_len, hidden)
cos = torch.randn(bsz, q_len, head_dim)
sin = torch.randn(bsz, q_len, head_dim)
out, attn_weights = sage_attention_patch.sage_attention_forward(
fake,
hidden_states,
position_embeddings=(cos, sin),
attention_mask=None,
past_key_values=None,
cache_position=None,
)
finally:
sage_attention_patch.SAGE_ATTENTION_FUNCTION = prev
assert called.get("ok") is True
assert tuple(out.shape) == (bsz, q_len, hidden)
assert attn_weights is None
# ---------------------------------------------------------------------------
# Arch dispatch: exact match, not a threshold.
# ---------------------------------------------------------------------------
# The previous chain tested ``arch_code >= 120``, then ``>= 90``, then
# ``== 89``, then ``>= 80``. sm100 / sm103 (Blackwell datacenter, CC 10.x)
# therefore fell into the SM90 branch. This module calls the kernels directly
# and never routes through sage's ``sageattn()``, so the pre-fix failure was
# not sage's ``ValueError: Unsupported CUDA architecture`` but SILENT
# wrong-kernel selection: the sm90 kernel asserts only that it COMPILED, not
# that it is running on the right silicon, and nothing downstream caught it.
_SAGE_KERNELS = (
"sageattn_qk_int8_pv_fp16_cuda",
"sageattn_qk_int8_pv_fp8_cuda",
"sageattn_qk_int8_pv_fp8_cuda_sm90",
)
def _make_kernel_stub(name):
# The dispatcher SELECTS a kernel, it never calls it, so identity has to be
# carried by the function object itself. One stable object per name, built
# once, so `is` identity can be asserted in the tests below.
def _fn(*a, **kw):
raise AssertionError("the dispatcher must not call the kernel")
_fn.kernel_name = name
return _fn
_KERNEL_STUBS = {name: _make_kernel_stub(name) for name in _SAGE_KERNELS}
def _dispatch(major, minor):
"""Run the real dispatcher against a faked device.
Returns (selected_kernel, pv_accum_dtype, arch_code). The selected kernel is
the stub object itself, so callers can assert identity.
"""
prev_funcs = {
name: getattr(sage_attention_patch, name, None) for name in _SAGE_KERNELS
}
prev_available = sage_attention_patch.SAGE_ATTENTION_AVAILABLE
prev_cuda = torch.cuda.is_available
prev_cap = torch.cuda.get_device_capability
for name in _SAGE_KERNELS:
setattr(sage_attention_patch, name, _KERNEL_STUBS[name])
sage_attention_patch.SAGE_ATTENTION_AVAILABLE = True
torch.cuda.is_available = lambda: True
torch.cuda.get_device_capability = lambda *a, **kw: (major, minor)
try:
attn_func, _, pv_accum_dtype = (
sage_attention_patch.get_sage_attention_function_and_params()
)
finally:
for name, fn in prev_funcs.items():
setattr(sage_attention_patch, name, fn)
sage_attention_patch.SAGE_ATTENTION_AVAILABLE = prev_available
torch.cuda.is_available = prev_cuda
torch.cuda.get_device_capability = prev_cap
return attn_func, pv_accum_dtype, major * 10 + minor
@pytest.mark.parametrize(
"capability,expected,pv_accum",
[
# int8 QK / fp16 PV on Ampere.
((8, 0), "sageattn_qk_int8_pv_fp16_cuda", "fp32"),
((8, 6), "sageattn_qk_int8_pv_fp16_cuda", "fp32"),
# int8 QK / fp8 PV everywhere else, accumulating in fp32+fp32.
((8, 9), "sageattn_qk_int8_pv_fp8_cuda", "fp32+fp32"),
((9, 0), "sageattn_qk_int8_pv_fp8_cuda_sm90", "fp32+fp32"),
((12, 0), "sageattn_qk_int8_pv_fp8_cuda", "fp32+fp32"),
],
)
def test_supported_arch_selects_the_right_kernel(capability, expected, pv_accum):
kernel, pv_accum_dtype, arch_code = _dispatch(*capability)
# Object identity, not a name: a threshold rule would hand back the SM90
# kernel for sm100/sm103 and a name-based assertion could be fooled by a
# renamed or aliased symbol.
assert kernel is _KERNEL_STUBS[expected], (
f"sm{arch_code} should select {expected}"
)
assert pv_accum_dtype == pv_accum
@pytest.mark.parametrize("capability", [(7, 5), (8, 7), (8, 8), (10, 0), (10, 3), (11, 0)])
def test_unsupported_arch_selects_no_kernel(capability):
"""No kernel at all — explicitly not the SM90 one."""
kernel, pv_accum_dtype, arch_code = _dispatch(*capability)
assert kernel is None, (
f"sm{arch_code} is not in sage's dispatch table; selecting a kernel for "
f"it is how a Hopper kernel ends up on silicon it was never built for"
)
assert pv_accum_dtype is None
def test_dispatch_agrees_with_the_project_capability_check():
"""The dispatcher and check_sage_attention_compatible must accept exactly
the same set, or the node offers sage on a machine the loader refuses.
Both directions are asserted: the check must not reject an arch the
dispatcher can serve, and must not accept one the dispatcher refuses. That
keeps the two literals — the dispatcher's exact-code chain and
SAGE_SUPPORTED_ARCHS — from drifting apart.
"""
from ComfyUI_VibeVoice.modules.attention_utils import (
SAGE_SUPPORTED_ARCHS,
check_sage_attention_compatible,
)
from unittest.mock import patch
# Every real arch the check knows about, accepted or not. Codes 75-120 step
# by 1; sm75 is sage's triton-only arch and the project excludes it.
all_codes = [c for c in range(75, 121)]
for code in all_codes:
major, minor = code // 10, code % 10
kernel, _, _ = _dispatch(major, minor)
arch = f"sm{code}"
with patch("ComfyUI_VibeVoice.modules.attention_utils.SAGE_ATTENTION_AVAILABLE", True), \
patch("torch.cuda.is_available", return_value=True), \
patch("torch.cuda.get_device_capability", return_value=(major, minor)):
check_says_yes = check_sage_attention_compatible()
assert (kernel is not None) == check_says_yes, (
f"{arch}: dispatcher selects {getattr(kernel, 'kernel_name', None)!r} "
f"but check_sage_attention_compatible() says {check_says_yes}. "
f"Project-supported archs are {sorted(SAGE_SUPPORTED_ARCHS)}; sage's "
f"own table is sm75/sm80/sm86/sm89/sm90/sm120."
)
assert check_says_yes == (arch in SAGE_SUPPORTED_ARCHS), (
f"{arch}: check_sage_attention_compatible() disagrees with the "
f"SAGE_SUPPORTED_ARCHS literal {sorted(SAGE_SUPPORTED_ARCHS)}"
)
def test_non_cuda_returns_no_kernel():
prev = torch.cuda.is_available
torch.cuda.is_available = lambda: False
try:
attn_func, gran, accum = sage_attention_patch.get_sage_attention_function_and_params()
finally:
torch.cuda.is_available = prev
assert (attn_func, gran, accum) == (None, None, None)
# ---------------------------------------------------------------------------
# Mask refusal
# ---------------------------------------------------------------------------
# sageattn 2.2.0 has NO attn_mask parameter. The old code used the mask only
# to pick ``is_causal`` and then dropped it, so a padded prefill (the ASR
# processor left-pads every batch) ran non-causally over the pad columns —
# silently wrong output. It now refuses.
def test_sage_attention_forward_refuses_a_non_none_mask():
fake, head_dim, num_heads, hidden = _make_fake_attention()
called = {"n": 0}
def stub_sage(q, k, v, tensor_layout="HND", is_causal=False,
qk_quant_gran="per_warp", pv_accum_dtype="fp32"):
called["n"] += 1
return q
prev = sage_attention_patch.SAGE_ATTENTION_FUNCTION
sage_attention_patch.SAGE_ATTENTION_FUNCTION = stub_sage
try:
bsz, q_len = 1, 8
hidden_states = torch.randn(bsz, q_len, hidden)
cos = torch.randn(bsz, q_len, head_dim)
sin = torch.randn(bsz, q_len, head_dim)
additive_mask = torch.zeros(bsz, 1, q_len, q_len)
with pytest.raises(ValueError, match="attention_mask"):
sage_attention_patch.sage_attention_forward(
fake,
hidden_states,
position_embeddings=(cos, sin),
attention_mask=additive_mask,
past_key_values=None,
cache_position=None,
)
finally:
sage_attention_patch.SAGE_ATTENTION_FUNCTION = prev
assert called["n"] == 0, (
"the kernel must never run with a mask it cannot honour — the mask is "
"the whole reason for the refusal"
)
def test_sage_attention_forward_still_runs_unmasked():
"""The guard must not disturb the working TTS path."""
fake, head_dim, num_heads, hidden = _make_fake_attention()
seen = {}
def stub_sage(q, k, v, tensor_layout="HND", is_causal=False,
qk_quant_gran="per_warp", pv_accum_dtype="fp32"):
seen["is_causal"] = is_causal
return q
prev = sage_attention_patch.SAGE_ATTENTION_FUNCTION
sage_attention_patch.SAGE_ATTENTION_FUNCTION = stub_sage
try:
bsz, q_len = 1, 8
hidden_states = torch.randn(bsz, q_len, hidden)
cos = torch.randn(bsz, q_len, head_dim)
sin = torch.randn(bsz, q_len, head_dim)
out, _ = sage_attention_patch.sage_attention_forward(
fake,
hidden_states,
position_embeddings=(cos, sin),
attention_mask=None,
past_key_values=None,
cache_position=None,
)
finally:
sage_attention_patch.SAGE_ATTENTION_FUNCTION = prev
assert tuple(out.shape) == (bsz, q_len, hidden)
# Causality is now decided purely by the query length, since the mask is
# guaranteed None by the guard above.
assert seen["is_causal"] is True