h3-blackwell-runtime/tests/test_fl2va_contracts.py
2026-08-20 17:39:44 +07:00

251 lines
11 KiB
Python

import math
import sys
from types import SimpleNamespace
import unittest
from unittest.mock import patch
import torch
from torch import nn
from torch.nn import functional as F
from h3_blackwell_runtime.packing import FRAME_RESCALE, H3PromptPacker, _video_t_spans
from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, run_attention
from h3_blackwell_runtime.qwen3vl_vision import (
TEXT_HEAD_DIM,
TEXT_ROPE_DIMS,
TEXT_ROPE_THETA,
VISION_HIDDEN,
Qwen3VL32BVision,
_VisionAttention,
_VisionPatchMerger,
_apply_rope_vision,
_text_run_ids,
mrope_freqs_cis,
)
from h3_blackwell_runtime.runtime import H3HotRuntime
from h3_blackwell_runtime.vae_encoder import MiniMaxH3VideoVAEEncoder, _downsample
class Fl2vaVAEContracts(unittest.TestCase):
def test_quant_conv_is_a_required_checkpoint_weight(self):
names = MiniMaxH3VideoVAEEncoder()._required_encoder_names()
self.assertIn("quant_conv.weight", names)
self.assertIn("quant_conv.bias", names)
def test_single_frame_is_encoded_without_temporal_prepad(self):
encoder = MiniMaxH3VideoVAEEncoder(tiling=False)
seen = []
def fake_encode(x):
seen.append(tuple(x.shape))
return torch.zeros((x.shape[0], 48, x.shape[2], 1, 1), device=x.device)
encoder._adaptive_encode = fake_encode
result = encoder.encode(torch.zeros(1, 3, 8, 8))
self.assertEqual(seen, [(1, 3, 1, 8, 8)])
self.assertEqual(tuple(result.shape), (1, 24, 1, 1, 1))
def test_downsample_pads_only_right_and_bottom(self):
x = torch.arange(16, dtype=torch.float32).reshape(1, 1, 1, 4, 4)
weight = torch.ones(1, 1, 3, 3, 3)
params = {"w": weight, "b": torch.zeros(1), "time": 1, "space": 2}
actual = _downsample(x, params)
padded = F.pad(x, (0, 1, 0, 1, 0, 0), mode="reflect")
expected = F.conv3d(padded, weight[:, :, -1:], params["b"], stride=(1, 2, 2))
torch.testing.assert_close(actual, expected)
def test_hot_runtime_preserves_zero_to_one_images(self):
runtime = H3HotRuntime.__new__(H3HotRuntime)
image = torch.tensor([[[[0.0, 0.5, 1.0]]]])
converted = runtime._image_to_uint8_nhwc(image)
self.assertEqual(converted.flatten().tolist(), [0, 127, 255])
class AttentionBackendContracts(unittest.TestCase):
def test_hot_backends_include_benchmark_candidates(self):
self.assertTrue({"sage2", "cudnn_sdpa", "ck_int8"}.issubset(AVAILABLE_BACKENDS))
def test_cudnn_backend_is_forced_without_fallback(self):
q = torch.randn(1, 2, 3, 4)
expected = torch.randn_like(q)
context = unittest.mock.MagicMock()
with (
patch("torch.nn.attention.sdpa_kernel", return_value=context) as kernel,
patch("h3_blackwell_runtime.attention.functional.scaled_dot_product_attention", return_value=expected) as sdpa,
):
actual = run_attention(q, q, q, backend="cudnn_sdpa", is_causal=False)
self.assertIs(actual, expected)
self.assertEqual(kernel.call_args.args[0], [torch.nn.attention.SDPBackend.CUDNN_ATTENTION])
sdpa.assert_called_once_with(q, q, q, is_causal=False)
def test_comfy_kitchen_int8_backend_dispatches_hnd_tensors(self):
q = torch.randn(1, 2, 3, 4)
expected = torch.randn_like(q)
kitchen = SimpleNamespace(int8_attention=unittest.mock.MagicMock(return_value=expected))
with patch.dict(sys.modules, {"comfy_kitchen": kitchen}):
actual = run_attention(q, q, q, backend="ck_int8", is_causal=False)
self.assertIs(actual, expected)
kitchen.int8_attention.assert_called_once_with(q, q, q)
def test_comfy_kitchen_int8_rejects_causal_attention(self):
q = torch.randn(1, 2, 3, 4)
with self.assertRaisesRegex(ValueError, "does not support causal"):
run_attention(q, q, q, backend="ck_int8", is_causal=True)
class Fl2vaVisionContracts(unittest.TestCase):
def test_visual_rotary_coordinates_are_block_major(self):
class CoordinateTable(nn.Module):
def forward(self, length):
return torch.arange(length, dtype=torch.float32).unsqueeze(1)
vision = Qwen3VL32BVision.__new__(Qwen3VL32BVision)
nn.Module.__init__(vision)
vision.spatial_merge_size = 2
vision.rotary_pos_emb = CoordinateTable()
coordinates = vision.rot_pos_emb(torch.tensor([[1, 4, 4]])).tolist()
self.assertEqual(coordinates[:8], [
[0.0, 0.0], [0.0, 1.0], [1.0, 0.0], [1.0, 1.0],
[0.0, 2.0], [0.0, 3.0], [1.0, 2.0], [1.0, 3.0],
])
def test_sdpa_output_is_restored_to_token_major_layout(self):
torch.manual_seed(7)
sequence, heads, head_dim = 3, 2, 2
hidden = heads * head_dim
qkv_weight = torch.randn(hidden * 3, hidden)
qkv_bias = torch.randn(hidden * 3)
proj_weight = torch.randn(hidden, hidden)
proj_bias = torch.randn(hidden)
module = _VisionAttention(qkv_weight, qkv_bias, proj_weight, proj_bias, num_heads=heads, head_dim=head_dim)
x = torch.randn(sequence, hidden)
cos = torch.ones(sequence, 1, head_dim)
sin = torch.zeros(sequence, 1, head_dim // 2)
actual = module(x, torch.tensor([0, sequence], dtype=torch.int32), (cos, sin, sin))
qkv = F.linear(x, qkv_weight, qkv_bias)
query, key, value = qkv.reshape(sequence, 3, heads, head_dim).permute(1, 0, 2, 3).unbind(0)
output = F.scaled_dot_product_attention(
query.transpose(0, 1).unsqueeze(0),
key.transpose(0, 1).unsqueeze(0),
value.transpose(0, 1).unsqueeze(0),
)
expected = F.linear(output.transpose(1, 2).reshape(sequence, hidden), proj_weight, proj_bias)
torch.testing.assert_close(actual, expected)
def test_vision_rope_uses_original_halves(self):
q = torch.tensor([[[1.0, 2.0, 3.0, 4.0]]])
k = q + 4
cos = torch.full_like(q, 0.5)
sin = torch.full_like(q[..., :2], 0.25)
neg_sin = -sin
actual_q, actual_k = _apply_rope_vision(q, k, (cos, sin, neg_sin))
def expected(x):
return torch.cat((x[..., :2] * 0.5 + x[..., 2:] * -0.25,
x[..., 2:] * 0.5 + x[..., :2] * 0.25), dim=-1)
torch.testing.assert_close(actual_q, expected(q))
torch.testing.assert_close(actual_k, expected(k))
def test_mrope_uses_reference_section_boundaries(self):
positions = torch.stack((torch.arange(8), torch.arange(8) + 10, torch.arange(8) + 20))
actual = mrope_freqs_cis(positions)[0]
inv_freq = 1.0 / (
TEXT_ROPE_THETA ** (torch.arange(0, TEXT_HEAD_DIM, 2, dtype=torch.float32) / TEXT_HEAD_DIM)
)
freqs = (inv_freq[None, :, None].expand(3, -1, 1) @ positions[:, None, :].float()).transpose(1, 2)
interleaved = freqs[0].clone()
for axis, offset in ((1, 1), (2, 2)):
index = slice(offset, TEXT_ROPE_DIMS[axis] * 3, 3)
interleaved[..., index] = freqs[axis, ..., index]
expected = torch.cat((interleaved, interleaved), dim=-1).cos().unsqueeze(0)
torch.testing.assert_close(actual, expected)
def test_mergers_preserve_existing_block_major_order(self):
x = torch.arange(4 * VISION_HIDDEN, dtype=torch.float32).reshape(4, VISION_HIDDEN)
passthrough = lambda value, *args, **kwargs: value
with (
patch("h3_blackwell_runtime.qwen3vl_vision.F.layer_norm", side_effect=passthrough),
patch("h3_blackwell_runtime.qwen3vl_vision.F.linear", side_effect=passthrough),
patch("h3_blackwell_runtime.qwen3vl_vision.F.gelu", side_effect=passthrough),
):
main = _VisionPatchMerger(*(torch.empty(1) for _ in range(6)), merge_size=2,
out_hidden_size=1, norm_dim=VISION_HIDDEN)
deepstack = _VisionPatchMerger(*(torch.empty(1) for _ in range(6)), merge_size=2,
out_hidden_size=1)
torch.testing.assert_close(main(x), x.reshape(1, -1))
torch.testing.assert_close(deepstack(x), x.reshape(1, -1))
def test_configured_tokenizer_preserves_empty_prompt(self):
calls = []
class RawTokenizer:
def __call__(self, text, **kwargs):
calls.append((text, kwargs))
return SimpleNamespace(input_ids=torch.empty((1, 0), dtype=torch.long))
self.assertEqual(_text_run_ids(SimpleNamespace(tokenizer=RawTokenizer()), ""), [])
self.assertEqual(calls[0][0], "")
self.assertFalse(calls[0][1]["add_special_tokens"])
class _FakeCheckpoint:
def tensor(self, name, dtype=None):
if name == "video_patch_proj.weight":
value = torch.zeros(5376, 96)
elif name == "video_patch_proj.bias":
value = torch.zeros(5376)
elif name == "audio_patch_proj.weight":
value = torch.zeros(5376, 32)
elif name == "audio_patch_proj.bias":
value = torch.zeros(5376)
else:
value = torch.empty(0)
return value.to(dtype=dtype) if dtype is not None else value
class Fl2vaPackingContracts(unittest.TestCase):
def test_each_keyframe_keeps_its_own_condition_segment(self):
packer = H3PromptPacker(_FakeCheckpoint())
text = torch.zeros(1, 3, 5376)
video = torch.zeros(1, 24, 2, 2, 2)
audio = torch.zeros(1, 32, 2, 2)
keyframes = [torch.zeros(1, 24, 1, 2, 2) for _ in range(2)]
_, _, segments, _, _, _ = packer(
text,
video,
audio,
0.5,
cond_latents=keyframes,
cond_frame_indices=[0, 21],
frame_count=22,
)
self.assertEqual(segments[1][:2], (3, 4))
self.assertEqual(segments[2][:2], (4, 5))
self.assertEqual(segments[1][2], segments[2][2])
def test_last_only_anchor_and_targets_share_reference_cursor(self):
packer = H3PromptPacker(_FakeCheckpoint())
text = torch.zeros(1, 3, 5376)
video = torch.zeros(1, 24, 2, 2, 2)
audio = torch.zeros(1, 32, 2, 2)
last = torch.zeros(1, 24, 1, 2, 2)
_, _, _, positions, _, _ = packer(
text,
video,
audio,
0.5,
cond_latents=[last],
cond_frame_indices=[21],
frame_count=22,
seed=1,
)
expected_last_t = 3.0 + sum(_video_t_spans(2)) - FRAME_RESCALE
self.assertTrue(math.isclose(float(positions[3, 0]), expected_last_t))
self.assertEqual(float(positions[4, 0]), 3.0) # target audio
self.assertEqual(float(positions[8, 0]), 3.0) # target video
if __name__ == "__main__":
unittest.main()