Files
wildminder-ComfyUI-DyPE/tests/test_hap_runtime.py
T

845 lines
34 KiB
Python

"""
Tests for the HAP runtime (``src/hap.py``).
Covers plan phases P0 (probes), P1 (ScopePlan + band math), P2 (backends +
runtime facade). See ``docs/plans/2026-08-15-hrdit-full-hap-implementation.md``.
All tests are CPU-safe; FlexAttention-specific tests are CUDA-gated and
auto-skip elsewhere.
"""
import pytest
import torch
from src import hap
@pytest.fixture
def mock_attn():
"""The conftest-provided (pristine SDPA) mock ``comfy.ldm.modules.attention`` module."""
import comfy.ldm.modules.attention as attn_mod
return attn_mod
# ---------------------------------------------------------------------------
# T0.1 — FlexAttention availability probe
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestFlexProbe:
def test_flex_probe_returns_bool(self):
"""The probe must always return a plain bool, whatever the env."""
result = hap.hap_flex_available()
assert isinstance(result, bool)
def test_flex_probe_no_raise_on_cpu(self, monkeypatch):
"""Probe must not raise when CUDA is unavailable."""
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
assert hap.hap_flex_available() is False
def test_flex_probe_false_when_import_fails(self, monkeypatch):
"""Probe returns False (never raises) if the flex_attention import blows up.
W3 note (2026-08-25): the probe now imports via
``from torch.nn.attention import flex_attention`` (the F823 fix), so
the simulated failure must trigger on that module path.
"""
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(hap, "_torch_version_at_least", lambda maj, mnr: True)
import builtins
real_import = builtins.__import__
def fake_import(name, *args, **kwargs):
if name == "torch.nn.attention":
raise ImportError("simulated missing flex_attention")
return real_import(name, *args, **kwargs)
monkeypatch.setattr(builtins, "__import__", fake_import)
assert hap.hap_flex_available() is False
def test_flex_probe_false_on_old_torch(self, monkeypatch):
"""Probe returns False when torch < 2.5 even with CUDA."""
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
monkeypatch.setattr(hap, "_torch_version_at_least", lambda maj, mnr: False)
assert hap.hap_flex_available() is False
def test_flex_probe_cuda_gate_reachable(self, monkeypatch):
"""W3 REGRESSION (ruff F823): the local
``import torch.nn.attention.flex_attention`` used to bind ``torch``
function-locally, so ``torch.cuda.is_available()`` raised
UnboundLocalError (swallowed by the except) and the probe returned
False on EVERY environment — the CUDA gate was unreachable. With a
CUDA-mocked-positive env and an old-torch stub, the version gate must
be what returns False (proving execution reached past the CUDA check).
"""
reached = {"version_gate": False}
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
def _version_gate(maj, mnr):
reached["version_gate"] = True
return False
monkeypatch.setattr(hap, "_torch_version_at_least", _version_gate)
assert hap.hap_flex_available() is False
assert reached["version_gate"], (
"execution never reached the version gate — the CUDA check raised "
"(UnboundLocalError regression)")
@pytest.mark.unit
class TestTorchVersionParse:
@pytest.mark.parametrize(
"version,expected",
[
("2.5.0", True),
("2.5.1+cu124", True),
("2.6.0.dev20241101", True),
("2.4.1", False),
("2.4", False),
("1.13.1+cpu", False),
("3.0.0", True),
],
)
def test_version_at_least(self, monkeypatch, version, expected):
monkeypatch.setattr(torch, "__version__", version)
assert hap._torch_version_at_least(2, 5) is expected
def test_version_parse_never_raises(self, monkeypatch):
monkeypatch.setattr(torch, "__version__", "garbage")
assert isinstance(hap._torch_version_at_least(2, 5), bool)
# ---------------------------------------------------------------------------
# Constants sanity (reference parity)
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestConstants:
def test_block_size(self):
assert hap.HAP_BLOCK == 64
def test_default_text_len(self):
assert hap.HAP_DEFAULT_TEXT_LEN == 512
def test_anchor_off_sentinel(self):
assert hap.HAP_ANCHOR_OFF == 1 << 30
def test_train_seq_len(self):
# FLUX training resolution: 64x64 image tokens + 512 text tokens.
assert hap.HAP_TRAIN_SEQ_LEN == 4608
# ---------------------------------------------------------------------------
# T0.2 — Synthetic multi-layer DiT fixture
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestToyDiT:
def test_toy_dit_call_order(self, mock_attn):
"""4 layers -> exactly 4 optimized_attention calls, in order 0..3."""
from _hrdit_fixtures import CallRecorder, make_toy_dit
rec = CallRecorder().install()
try:
dit = make_toy_dit(num_layers=4, heads=2, dim=16, text_len=8, img_hw=4, seed=7)
dit.forward()
finally:
rec.uninstall()
assert len(rec.calls) == 4
assert [c["layer"] for c in dit.call_log] == [0, 1, 2, 3]
def test_toy_dit_shapes_and_layout(self, mock_attn):
"""Each call sees (1, heads, text_len + img_hw^2, dim) tensors."""
from _hrdit_fixtures import CallRecorder, make_toy_dit
rec = CallRecorder().install()
try:
dit = make_toy_dit(num_layers=2, heads=3, dim=8, text_len=5, img_hw=3, seed=1)
dit.forward()
finally:
rec.uninstall()
expected_seq = 5 + 3 * 3
for q, k, v, heads in rec.calls:
assert tuple(q.shape) == (1, 3, expected_seq, 8)
assert tuple(k.shape) == tuple(q.shape)
assert tuple(v.shape) == tuple(q.shape)
assert heads == 3
assert dit.seq_len == expected_seq
def test_toy_dit_deterministic(self, mock_attn):
"""Same seed -> identical outputs across two fresh instances."""
from _hrdit_fixtures import make_toy_dit
out1 = make_toy_dit(num_layers=3, seed=42).forward()
out2 = make_toy_dit(num_layers=3, seed=42).forward()
assert torch.equal(out1, out2)
def test_toy_dit_output_finite(self, mock_attn):
from _hrdit_fixtures import make_toy_dit
out = make_toy_dit(num_layers=4, seed=3).forward()
assert torch.isfinite(out).all()
assert out.shape == (1, 8 + 16, 2 * 16)
# ---------------------------------------------------------------------------
# T1.1 — ScopePlan load / validate / round-trip
# ---------------------------------------------------------------------------
def _tiny_plan_dict():
return {
"alphas": [[2048.0, 0.0], [128.0, 64.0]],
"betas": [[0.0, 0.25], [0.5, 0.0]],
}
@pytest.mark.unit
class TestScopePlan:
def test_scopeplan_roundtrip(self):
d = _tiny_plan_dict()
plan = hap.ScopePlan.from_dict(d)
assert plan.to_dict() == d
assert plan.num_layers == 2
assert plan.num_heads == 2
def test_scopeplan_loads_reference_flux_plan(self):
"""The REAL reference FLUX plan must load unchanged (format compat).
Uses the SHIPPED plan (configs/scope_plan_flux.json) — version-
controlled and present in CI. (An identical copy lives under the
gitignored .dev/data research tree; do not reference that path.)"""
import pathlib
path = (
pathlib.Path(__file__).parent.parent
/ "configs" / "scope_plan_flux.json"
)
assert path.exists(), f"shipped reference plan missing: {path}"
plan = hap.ScopePlan.load(path)
assert plan.num_layers == 57
assert plan.num_heads == 24
assert all(a == 2048.0 for row in plan.alphas for a in row)
assert all(b == 0.0 for row in plan.betas for b in row)
def test_scopeplan_save_load_roundtrip(self, tmp_path):
plan = hap.ScopePlan.from_dict(_tiny_plan_dict())
out = tmp_path / "plan.json"
plan.save(out)
reloaded = hap.ScopePlan.load(out)
assert reloaded.to_dict() == plan.to_dict()
# -- excluded_head_counts metadata (2026-08-23 head-count warning fix) ----
def test_scopeplan_excluded_head_counts_roundtrip(self):
"""A plan WITH excluded_head_counts round-trips the field exactly."""
d = dict(_tiny_plan_dict())
d["excluded_head_counts"] = [20]
plan = hap.ScopePlan.from_dict(d)
assert plan.excluded_head_counts == [20]
assert plan.to_dict() == d
def test_scopeplan_excluded_head_counts_omitted_when_absent(self):
"""A legacy plan WITHOUT the field round-trips to the exact legacy
shape (no spurious key) — full backward compatibility."""
d = _tiny_plan_dict()
plan = hap.ScopePlan.from_dict(d)
assert plan.excluded_head_counts is None
assert plan.to_dict() == d
assert "excluded_head_counts" not in plan.to_dict()
def test_scopeplan_excluded_head_counts_empty_omitted(self):
"""An EMPTY excluded list is treated as absent (omitted from to_dict)."""
plan = hap.ScopePlan(
alphas=_tiny_plan_dict()["alphas"],
betas=_tiny_plan_dict()["betas"],
excluded_head_counts=[],
)
assert "excluded_head_counts" not in plan.to_dict()
def test_scopeplan_rejects_bad_excluded_head_counts(self):
"""excluded_head_counts must be a list of ints; else ValueError."""
d = dict(_tiny_plan_dict())
d["excluded_head_counts"] = [20, "x"]
with pytest.raises(ValueError, match="excluded_head_counts"):
hap.ScopePlan.from_dict(d)
d2 = dict(_tiny_plan_dict())
d2["excluded_head_counts"] = "not-a-list"
with pytest.raises(ValueError, match="excluded_head_counts"):
hap.ScopePlan.from_dict(d2)
def test_scopeplan_rejects_ragged(self):
d = {"alphas": [[1.0, 2.0], [3.0]], "betas": [[0.0, 0.0], [0.0]]}
with pytest.raises(ValueError, match="ragged"):
hap.ScopePlan.from_dict(d)
def test_scopeplan_rejects_negative(self):
d = {"alphas": [[-1.0]], "betas": [[0.0]]}
with pytest.raises(ValueError, match=">= 0"):
hap.ScopePlan.from_dict(d)
def test_scopeplan_rejects_nonfinite(self):
d = {"alphas": [[float("inf")]], "betas": [[0.0]]}
with pytest.raises(ValueError, match="finite"):
hap.ScopePlan.from_dict(d)
def test_scopeplan_rejects_missing_key(self):
with pytest.raises(ValueError, match="missing required key"):
hap.ScopePlan.from_dict({"alphas": [[1.0]]})
def test_scopeplan_rejects_layer_count_mismatch(self):
d = {"alphas": [[1.0]], "betas": [[0.0], [0.0]]}
with pytest.raises(ValueError, match="layers"):
hap.ScopePlan.from_dict(d)
def test_scopeplan_rejects_head_count_mismatch(self):
d = {"alphas": [[1.0, 2.0]], "betas": [[0.0]]}
with pytest.raises(ValueError, match="heads"):
hap.ScopePlan.from_dict(d)
def test_scopeplan_rejects_non_numeric(self):
d = {"alphas": [["x"]], "betas": [[0.0]]}
with pytest.raises(ValueError, match="number"):
hap.ScopePlan.from_dict(d)
def test_scopeplan_rejects_empty(self):
with pytest.raises(ValueError, match="at least one layer"):
hap.ScopePlan.from_dict({"alphas": [], "betas": []})
def test_layer_bands(self):
plan = hap.ScopePlan.from_dict(_tiny_plan_dict())
# alpha=2048, beta=0, seq=66048 -> band 63 (reference FLUX numbers).
assert plan.layer_bands(0, 66048)[0] == 63
# ---------------------------------------------------------------------------
# T1.2 — band_blocks (reference formula, exact)
# ---------------------------------------------------------------------------
def _reference_band_blocks(alphas, betas, seq_len, block=64):
"""Inline copy of hrdit/hap.py HapRuntime.band_blocks (parity oracle)."""
nbx = seq_len // block
return [max(2 * int(a // block + b * nbx) - 1, 1) for a, b in zip(alphas, betas)]
@pytest.mark.unit
class TestBandBlocks:
def test_band_blocks_reference_flux_plan(self):
# alpha=2048, beta=0, seq=66048 (4K FLUX) -> 2*int(2048/64)-1 = 63.
assert hap.band_blocks([2048.0], [0.0], 66048) == [63]
def test_band_blocks_beta_only(self):
# alpha=0, beta=0.5, seq=66048 -> 2*int(0.5*1032)-1 = 1031.
assert hap.band_blocks([0.0], [0.5], 66048) == [2 * int(0.5 * (66048 // 64)) - 1]
assert hap.band_blocks([0.0], [0.5], 66048) == [1031]
def test_band_blocks_min_one(self):
assert hap.band_blocks([0.0], [0.0], 66048) == [1]
def test_band_blocks_matches_reference_impl(self):
"""Property test vs the inline reference copy on 50 random cases."""
g = torch.Generator().manual_seed(123)
for _ in range(50):
n = int(torch.randint(1, 8, (1,), generator=g).item())
alphas = [float(torch.randint(0, 4097, (1,), generator=g).item()) for _ in range(n)]
betas = [float(torch.rand(1, generator=g).item()) for _ in range(n)]
seq = int(torch.randint(64, 70000, (1,), generator=g).item())
assert hap.band_blocks(alphas, betas, seq) == _reference_band_blocks(alphas, betas, seq)
def test_half_blocks(self):
assert hap.half_blocks([63, 1, 1031, 4]) == [31, 0, 515, 1]
# ---------------------------------------------------------------------------
# T2.1 — HapContext + contextvars
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestHapContextVars:
def setup_method(self):
hap.HapRuntime.reset()
def teardown_method(self):
from src.spa_context import set_hap_context, set_hrdit_layer_idx
set_hap_context(None)
set_hrdit_layer_idx(0)
hap.HapRuntime.reset()
def test_hap_context_default_inactive(self):
from src.spa_context import get_hap_context
assert get_hap_context() is None
def test_hap_context_set_get_clear(self):
from src.spa_context import get_hap_context, set_hap_context
plan = hap.ScopePlan.from_dict(_tiny_plan_dict())
ctx = hap.HapContext(active=True, plan=plan)
set_hap_context(ctx)
assert get_hap_context() is ctx
set_hap_context(None)
assert get_hap_context() is None
def test_layer_counter_set_get_reset(self):
from src.spa_context import get_hrdit_layer_idx, next_hrdit_layer_idx, set_hrdit_layer_idx
assert get_hrdit_layer_idx() == 0
assert next_hrdit_layer_idx() == 0
assert next_hrdit_layer_idx() == 1
assert get_hrdit_layer_idx() == 2
set_hrdit_layer_idx(0)
assert get_hrdit_layer_idx() == 0
def test_context_isolation_across_copies(self):
"""Contextvar mutations inside a copied context don't leak out."""
import contextvars
from src.spa_context import get_hrdit_layer_idx, next_hrdit_layer_idx, set_hrdit_layer_idx
set_hrdit_layer_idx(0)
ctx = contextvars.copy_context()
ctx.run(next_hrdit_layer_idx)
ctx.run(next_hrdit_layer_idx)
# The outer context is untouched.
assert get_hrdit_layer_idx() == 0
def test_hap_context_resolve_backend_auto(self, monkeypatch):
plan = hap.ScopePlan.from_dict(_tiny_plan_dict())
ctx = hap.HapContext(active=True, plan=plan, backend="auto")
monkeypatch.setattr(hap, "hap_flex_available", lambda: False)
assert ctx.resolve_backend() == "dense"
monkeypatch.setattr(hap, "hap_flex_available", lambda: True)
assert ctx.resolve_backend() == "flex"
def test_hap_context_resolve_backend_explicit(self):
plan = hap.ScopePlan.from_dict(_tiny_plan_dict())
for backend in ("flex", "dense", "off"):
ctx = hap.HapContext(active=True, plan=plan, backend=backend)
assert ctx.resolve_backend() == backend
# ---------------------------------------------------------------------------
# T2.2 — Dense backend
# ---------------------------------------------------------------------------
def _rand_qkv(B=1, H=2, S=128, D=16, seed=0, dtype=torch.float64):
g = torch.Generator().manual_seed(seed)
q = torch.randn(B, H, S, D, generator=g, dtype=dtype)
k = torch.randn(B, H, S, D, generator=g, dtype=dtype)
v = torch.randn(B, H, S, D, generator=g, dtype=dtype)
return q, k, v
@pytest.mark.unit
class TestDenseBackend:
def test_dense_backend_matches_manual_masked_softmax(self):
"""fp64: manual softmax over masked logits @ v == backend output."""
S, H, text_len = 128, 2, 32
q, k, v = _rand_qkv(H=H, S=S, seed=5)
halves = [1, 3]
mask = hap.build_band_mask(S, text_len, halves, anchor_stride=2)
out = hap.hap_attn_dense(q, k, v, mask)
scale = q.shape[-1] ** -0.5
logits = (q @ k.transpose(-1, -2)) * scale
neg_inf = torch.finfo(logits.dtype).min
amask = torch.where(mask.unsqueeze(0), torch.zeros_like(logits), neg_inf)
ref = torch.softmax(logits + amask, dim=-1) @ v
assert torch.allclose(out, ref, atol=1e-12)
def test_dense_backend_text_tokens_full_attention(self):
"""Text query rows equal plain SDPA output (mask all-True there)."""
import torch.nn.functional as F
S, H, text_len = 96, 2, 16
q, k, v = _rand_qkv(H=H, S=S, seed=6)
mask = hap.build_band_mask(S, text_len, [0], 0)
out = hap.hap_attn_dense(q, k, v, mask)
plain = F.scaled_dot_product_attention(q, k, v, scale=q.shape[-1] ** -0.5)
assert torch.allclose(out[:, :, :text_len], plain[:, :, :text_len], atol=1e-12)
def test_dense_backend_scale_passthrough(self):
"""An explicit scale is honoured (differs from the default)."""
S, H, text_len = 64, 1, 0
q, k, v = _rand_qkv(H=H, S=S, seed=7)
mask = hap.build_band_mask(S, text_len, [10], 0) # full attention
out_default = hap.hap_attn_dense(q, k, v, mask)
out_custom = hap.hap_attn_dense(q, k, v, mask, scale=2.0)
assert not torch.allclose(out_default, out_custom, atol=1e-9)
# ---------------------------------------------------------------------------
# T2.3 — Flex backend (CUDA-gated; auto-skip elsewhere)
# ---------------------------------------------------------------------------
_FLEX_SKIP = pytest.mark.skipif(
not hap.hap_flex_available(),
reason="FlexAttention requires CUDA + torch>=2.5",
)
@pytest.mark.unit
class TestFlexBackend:
@_FLEX_SKIP
def test_flex_matches_dense_backend(self):
S, H, text_len = 256, 2, 64
q, k, v = _rand_qkv(H=H, S=S, seed=11, dtype=torch.float32)
q, k, v = q.cuda(), k.cuda(), v.cuda()
halves = [1, 2]
# band = 2*int(beta*nbx)-1, half = (band-1)//2 -> beta = (half+1)/nbx.
nbx = S // 64
plan = hap.ScopePlan(
alphas=[[0.0, 0.0]],
betas=[[(halves[0] + 1) / nbx, (halves[1] + 1) / nbx]],
)
ctx = hap.HapContext(active=True, plan=plan, text_len=text_len, backend="flex")
runtime = hap.HapRuntime.get()
out_flex = runtime.attn(q, k, v, 0, ctx=ctx)
mask = hap.build_band_mask(S, text_len, halves, 0)
out_dense = hap.hap_attn_dense(q, k, v, mask.cuda())
assert torch.allclose(out_flex, out_dense, rtol=1e-2, atol=1e-3)
@_FLEX_SKIP
def test_flex_mask_cache_reuse(self):
S, H, text_len = 256, 2, 64
q, k, v = _rand_qkv(H=H, S=S, seed=12, dtype=torch.float32)
q, k, v = q.cuda(), k.cuda(), v.cuda()
plan = hap.ScopePlan(alphas=[[0.0, 0.0]], betas=[[0.1, 0.1]])
ctx = hap.HapContext(active=True, plan=plan, text_len=text_len, backend="flex")
runtime = hap.HapRuntime.get()
runtime.attn(q, k, v, 0, ctx=ctx)
n_after_first = runtime.prepare_count
runtime.attn(q, k, v, 0, ctx=ctx)
assert runtime.prepare_count == n_after_first
# ---------------------------------------------------------------------------
# T2.4 — HapRuntime facade
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestHapRuntime:
def setup_method(self):
hap.HapRuntime.reset()
def teardown_method(self):
from src.spa_context import set_hap_context
set_hap_context(None)
hap.HapRuntime.reset()
def _ctx(self, num_layers=3, backend="dense"):
# W2.7 fix (2026-08-25): DISTINCT betas per layer. The mask cache is
# keyed by the RESOLVED halves, so a plan with identical per-layer
# scopes shares one mask across all layers and a "one prepare per
# distinct scope" count can never reach ``num_layers`` (the pre-fix
# fixture used [0.5, 0.5] everywhere -> prepare_count == 1).
# Cycle [0.5, 0.75, 1.0]: at nbx=4 (seq 256) these resolve to halves
# {1, 2, 3} — three distinct masks.
cycle = [0.5, 0.75, 1.0]
plan = hap.ScopePlan(
alphas=[[0.0, 0.0]] * num_layers,
betas=[[cycle[i % len(cycle)]] * 2 for i in range(num_layers)],
)
return hap.HapContext(active=True, plan=plan, text_len=0, backend=backend)
def test_runtime_lazy_prepare_counts(self):
"""3 layers x 2 calls -> exactly 3 mask builds (one per distinct scope)."""
from src.spa_context import set_hap_context
ctx = self._ctx(num_layers=3)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
# S=256 -> nbx=4 so the cycled betas resolve to DISTINCT halves.
q, k, v = _rand_qkv(H=2, S=256, seed=21)
for _ in range(2):
for layer in range(3):
out = runtime.attn(q, k, v, layer)
assert out is not None and out.shape == q.shape
assert runtime.prepare_count == 3
def test_runtime_seq_len_change_reprepares(self):
from src.spa_context import set_hap_context
ctx = self._ctx(num_layers=1)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q1, k1, v1 = _rand_qkv(H=2, S=128, seed=22)
runtime.attn(q1, k1, v1, 0)
assert runtime.prepare_count == 1
q2, k2, v2 = _rand_qkv(H=2, S=192, seed=23)
runtime.attn(q2, k2, v2, 0)
assert runtime.prepare_count == 2
def test_runtime_inactive_returns_none(self):
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=2, S=64, seed=24)
assert runtime.attn(q, k, v, 0) is None # no context set
def test_runtime_off_backend_falls_back_with_warning(self, caplog):
import logging
from src.spa_context import set_hap_context
ctx = self._ctx(num_layers=1, backend="off")
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=2, S=64, seed=25)
with caplog.at_level(logging.WARNING, logger="src.hap"):
out = runtime.attn(q, k, v, 0)
assert out is None
assert any("off" in rec.message for rec in caplog.records)
def test_runtime_layer_overflow_returns_none_with_warning(self, caplog):
import logging
from src.spa_context import set_hap_context
ctx = self._ctx(num_layers=1)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=2, S=64, seed=26)
with caplog.at_level(logging.WARNING, logger="src.hap"):
out = runtime.attn(q, k, v, 5)
assert out is None
assert any("exceeds" in rec.message for rec in caplog.records)
def test_runtime_dense_output_matches_oracle(self):
"""End-to-end: runtime dense dispatch == manual masked softmax."""
from src.spa_context import set_hap_context
S, text_len = 128, 32
# alpha=64 tokens -> band = 2*int(64/64)-1 = 1 -> half 0.
plan = hap.ScopePlan(alphas=[[64.0, 64.0]], betas=[[0.0, 0.0]])
ctx = hap.HapContext(active=True, plan=plan, text_len=text_len, backend="dense")
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=2, S=S, seed=27)
out = runtime.attn(q, k, v, 0)
mask = hap.build_band_mask(S, text_len, [0, 0], 0)
ref = hap.hap_attn_dense(q, k, v, mask)
assert torch.allclose(out, ref, atol=1e-12)
def test_flops_ratio_bounds(self):
plan = hap.ScopePlan.from_dict(_tiny_plan_dict())
ratio = hap.flops_ratio(plan, seq_len=2048, text_len=512)
assert 0.0 < ratio <= 1.0
# Full-attention plan (huge alpha) -> ratio ~ 1.
full = hap.ScopePlan(alphas=[[10**9]], betas=[[0.0]])
assert hap.flops_ratio(full, seq_len=2048, text_len=0) == pytest.approx(1.0)
# ---------------------------------------------------------------------------
# T3.1/T3.2 — Decline guards (plan 2026-08-16 G3)
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestDeclineGuards:
"""HAP must DECLINE (return None -> plain attention) — never crash, never
silent wrong math — for calls its square, plan-shaped mask cannot serve:
* non-square attention (cross-attention, ``kv_len != q_len``), and
* head-count mismatch (scope plan heads != model heads).
Both decline with a one-time log latch (reset by ``HapRuntime.reset()``).
"""
def setup_method(self):
hap.HapRuntime.reset()
def teardown_method(self):
from src.spa_context import set_hap_context
set_hap_context(None)
hap.HapRuntime.reset()
def _ctx(self, num_layers=3, num_heads=2, backend="dense"):
plan = hap.ScopePlan(
alphas=[[64.0] * num_heads for _ in range(num_layers)],
betas=[[0.0] * num_heads for _ in range(num_layers)],
)
return hap.HapContext(active=True, plan=plan, text_len=0, backend=backend)
def test_nonsquare_cross_attention_returns_none(self):
"""Cross-attention (kv_len != q_len) -> None (plain-attention fallback)."""
from src.spa_context import set_hap_context
ctx = self._ctx(num_heads=2)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
# q has 64 tokens, k/v have 32 (cross-attention).
g = torch.Generator().manual_seed(30)
q = torch.randn(1, 2, 64, 16, generator=g)
k = torch.randn(1, 2, 32, 16, generator=g)
v = torch.randn(1, 2, 32, 16, generator=g)
assert runtime.attn(q, k, v, 0) is None
def test_nonsquare_decline_is_one_time_debug(self, caplog):
"""The non-square decline logs at most ONE debug line per runtime."""
import logging
from src.spa_context import set_hap_context
ctx = self._ctx(num_heads=2)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
g = torch.Generator().manual_seed(31)
q = torch.randn(1, 2, 64, 16, generator=g)
k = torch.randn(1, 2, 32, 16, generator=g)
v = torch.randn(1, 2, 32, 16, generator=g)
with caplog.at_level(logging.DEBUG, logger="src.hap"):
runtime.attn(q, k, v, 0)
runtime.attn(q, k, v, 1) # second call must be silent
nonsquare = [r for r in caplog.records if "non-square" in r.message]
assert len(nonsquare) == 1
def test_head_mismatch_returns_none(self):
"""Plan with 24 heads vs q with 16 heads -> None (wrong plan for model)."""
from src.spa_context import set_hap_context
ctx = self._ctx(num_heads=24) # FLUX plan shape
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=16, S=64, seed=32) # Anima runs 16 heads
assert runtime.attn(q, k, v, 0) is None
def test_head_mismatch_decline_is_one_time_warning(self, caplog):
"""The head-mismatch decline logs ONE warning naming both counts."""
import logging
from src.spa_context import set_hap_context
ctx = self._ctx(num_heads=24)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=16, S=64, seed=33)
with caplog.at_level(logging.WARNING, logger="src.hap"):
runtime.attn(q, k, v, 0)
runtime.attn(q, k, v, 1) # second call must be silent
mismatch = [r for r in caplog.records if "heads" in r.message]
assert len(mismatch) == 1
# The warning names both the plan's and the model's head counts.
assert "24" in mismatch[0].message
assert "16" in mismatch[0].message
def test_matching_heads_unaffected(self):
"""Matching head count still engages the kernel (existing behaviour)."""
from src.spa_context import set_hap_context
ctx = self._ctx(num_heads=2)
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=2, S=64, seed=34)
out = runtime.attn(q, k, v, 0)
assert out is not None and out.shape == q.shape
# -- EXPECTED auxiliary fallback vs GENUINE wrong plan (2026-08-23) -------
def _ctx_with_excluded(self, num_heads=2, excluded=(16,)):
"""A plan that DECLARED ``excluded`` head counts during calibration."""
plan = hap.ScopePlan(
alphas=[[64.0] * num_heads for _ in range(3)],
betas=[[0.0] * num_heads for _ in range(3)],
excluded_head_counts=list(excluded),
)
return hap.HapContext(active=True, plan=plan, text_len=0, backend="dense")
def test_aux_head_mismatch_returns_none(self):
"""A head count in excluded_head_counts still declines to None."""
from src.spa_context import set_hap_context
ctx = self._ctx_with_excluded(num_heads=2, excluded=(16,))
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=16, S=64, seed=40)
assert runtime.attn(q, k, v, 0) is None
def test_aux_head_mismatch_decline_is_one_time_info(self, caplog):
"""An EXPECTED aux fallback logs ONE INFO (not a WARNING) naming the
excluded head count, and stays silent on repeat calls."""
import logging
from src.spa_context import set_hap_context
ctx = self._ctx_with_excluded(num_heads=2, excluded=(16,))
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=16, S=64, seed=41)
with caplog.at_level(logging.INFO, logger="src.hap"):
runtime.attn(q, k, v, 0)
runtime.attn(q, k, v, 1) # second call must be silent
infos = [r for r in caplog.records
if r.levelno == logging.INFO and "EXCLUDED during" in r.message]
warnings = [r for r in caplog.records
if r.levelno == logging.WARNING and "does not match" in r.message]
assert len(infos) == 1
assert warnings == [] # NOT a scary wrong-plan warning
assert "16" in infos[0].message
def test_genuine_head_mismatch_still_warning(self, caplog):
"""A head count NOT in excluded_head_counts still logs the WARNING
(genuinely wrong plan) — the pre-fix behaviour is preserved."""
import logging
from src.spa_context import set_hap_context
# Plan declares excluded=(99,) but the model runs 16 -> NOT excluded.
ctx = self._ctx_with_excluded(num_heads=2, excluded=(99,))
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=16, S=64, seed=42)
with caplog.at_level(logging.WARNING, logger="src.hap"):
runtime.attn(q, k, v, 0)
warnings = [r for r in caplog.records
if r.levelno == logging.WARNING and "does not match" in r.message]
infos = [r for r in caplog.records
if r.levelno == logging.INFO and "EXCLUDED during" in r.message]
assert len(warnings) == 1
assert infos == []
assert "2" in warnings[0].message and "16" in warnings[0].message
def test_aux_and_genuine_latches_independent(self, caplog):
"""The aux INFO latch and the genuine WARNING latch are independent:
an aux call then a genuine mismatch each log exactly once."""
import logging
from src.spa_context import set_hap_context
ctx = self._ctx_with_excluded(num_heads=2, excluded=(16,))
set_hap_context(ctx)
runtime = hap.HapRuntime.get()
q_aux, k_aux, v_aux = _rand_qkv(H=16, S=64, seed=43)
q_bad, k_bad, v_bad = _rand_qkv(H=8, S=64, seed=44) # not excluded
with caplog.at_level(logging.INFO, logger="src.hap"):
runtime.attn(q_aux, k_aux, v_aux, 0) # aux -> INFO
runtime.attn(q_bad, k_bad, v_bad, 1) # genuine -> WARNING
infos = [r for r in caplog.records
if r.levelno == logging.INFO and "EXCLUDED during" in r.message]
warnings = [r for r in caplog.records
if r.levelno == logging.WARNING and "does not match" in r.message]
assert len(infos) == 1
assert len(warnings) == 1
assert runtime._noted_aux_fallback is True
assert runtime._warned_head_mismatch is True
def test_decline_latches_reset_by_runtime_reset(self):
"""``HapRuntime.reset()`` drops the singleton -> fresh latches."""
from src.spa_context import set_hap_context
ctx = self._ctx(num_heads=24)
set_hap_context(ctx)
r1 = hap.HapRuntime.get()
q, k, v = _rand_qkv(H=16, S=64, seed=35)
r1.attn(q, k, v, 0)
assert r1._warned_head_mismatch is True
hap.HapRuntime.reset()
r2 = hap.HapRuntime.get()
assert r2 is not r1
assert r2._warned_head_mismatch is False
assert r2._warned_nonsquare is False