Enable H3 VAE attention backend selection.
This commit is contained in:
@@ -15,7 +15,9 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from fastvideo.attention import get_attn_backend
|
||||
from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.profiler import nvtx_range
|
||||
|
||||
|
||||
@@ -293,6 +295,7 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
|
||||
class MiniMaxH3VideoAttention(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
|
||||
"""Build projections and the selected dense FastVideo attention implementation."""
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
@@ -304,12 +307,34 @@ class MiniMaxH3VideoAttention(nn.Module):
|
||||
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
|
||||
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
|
||||
self.attn_impl = None
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
if current_platform.is_cuda_alike():
|
||||
attention_backend = get_attn_backend(
|
||||
dim_head,
|
||||
# FlashAttention executes the FP32 VAE activations in BF16 and
|
||||
# restores FP32 output, so resolve against the kernel dtype.
|
||||
torch.bfloat16,
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
),
|
||||
)
|
||||
self.attn_impl = attention_backend.get_impl_cls()(
|
||||
num_heads=heads,
|
||||
head_size=dim_head,
|
||||
softmax_scale=dim_head**-0.5,
|
||||
num_kv_heads=heads,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply dense self-attention to one spatial VAE token sequence."""
|
||||
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
|
||||
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
|
||||
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
|
||||
@@ -330,9 +355,17 @@ class MiniMaxH3VideoAttention(nn.Module):
|
||||
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
|
||||
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
|
||||
|
||||
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
|
||||
if self.attn_impl is not None and query.device.type != "cpu":
|
||||
# VAE decoding has no diffusion-step metadata, so call the selected
|
||||
# backend implementation directly with dense BSHD tensors.
|
||||
hidden_states = self.attn_impl.forward(query, key, value, None)
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
else:
|
||||
# Keep CPU construction and execution available without requiring
|
||||
# an accelerator attention backend.
|
||||
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
|
||||
return self.to_out[0](hidden_states)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,329 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU contract tests for MiniMax H3 VAE compilation and profiling ranges."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.vaes.minimax_h3_audio import (
|
||||
MiniMaxH3AudioBigVGANDecoder,
|
||||
MiniMaxH3AudioVAE,
|
||||
)
|
||||
from fastvideo.models.vaes.minimax_h3_video import (
|
||||
AutoencoderKLMiniMaxH3,
|
||||
MiniMaxH3VideoAttention,
|
||||
MiniMaxH3VideoViTDecoder3d,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
|
||||
|
||||
def _empty_typed_module(module_type: type[nn.Module]) -> nn.Module:
|
||||
"""Create a weightless instance that retains its production module type."""
|
||||
module = object.__new__(module_type)
|
||||
nn.Module.__init__(module)
|
||||
return module
|
||||
|
||||
|
||||
def _assert_dynamic_compile_selects_decoder(
|
||||
vae_type: type[nn.Module],
|
||||
decoder_type: type[nn.Module],
|
||||
) -> None:
|
||||
"""Verify one H3 VAE compiles only its top-level decoder in place."""
|
||||
vae = _empty_typed_module(vae_type)
|
||||
decoder = _empty_typed_module(decoder_type)
|
||||
same_type_under_another_name = _empty_typed_module(decoder_type)
|
||||
unrelated_submodule = nn.Identity()
|
||||
vae.decoder = decoder
|
||||
vae.same_type_under_another_name = same_type_under_another_name
|
||||
vae.unrelated_submodule = unrelated_submodule
|
||||
|
||||
compiled_forward = Mock(name="compiled_forward")
|
||||
compile_kwargs = {"backend": "inductor", "dynamic": False}
|
||||
with patch(
|
||||
"fastvideo.pipelines.composed_pipeline_base.torch.compile",
|
||||
return_value=compiled_forward,
|
||||
) as compile_mock:
|
||||
compiled_count = ComposedPipelineBase._compile_with_conditions(vae, compile_kwargs)
|
||||
|
||||
assert compiled_count == 1
|
||||
compile_mock.assert_called_once()
|
||||
selected_forward = compile_mock.call_args.args[0]
|
||||
assert selected_forward.__self__ is decoder
|
||||
assert selected_forward.__func__ is decoder_type.forward
|
||||
assert compile_mock.call_args.kwargs == compile_kwargs
|
||||
assert decoder.forward is compiled_forward
|
||||
assert "forward" not in same_type_under_another_name.__dict__
|
||||
assert "forward" not in unrelated_submodule.__dict__
|
||||
|
||||
wrong_type_vae = _empty_typed_module(vae_type)
|
||||
wrong_type_vae.decoder = nn.Identity()
|
||||
with patch("fastvideo.pipelines.composed_pipeline_base.torch.compile") as wrong_type_compile:
|
||||
wrong_type_count = ComposedPipelineBase._compile_with_conditions(wrong_type_vae, compile_kwargs)
|
||||
assert wrong_type_count == 0
|
||||
wrong_type_compile.assert_not_called()
|
||||
|
||||
|
||||
def _assert_reduce_overhead_compile(compiled_function: Any) -> None:
|
||||
"""Verify a class-owned compile boundary enables CUDA Graph replay."""
|
||||
assert hasattr(compiled_function, "get_compiler_config")
|
||||
assert compiled_function.get_compiler_config()["triton.cudagraphs"] is True
|
||||
|
||||
|
||||
def test_video_attention_uses_selected_fastvideo_backend() -> None:
|
||||
"""Pass BSHD tensors to the selected dense backend without forward metadata."""
|
||||
backend_call: dict[str, Any] = {}
|
||||
|
||||
class RecordingAttentionImpl:
|
||||
"""Record the backend construction and forward contracts."""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
backend_call["init"] = kwargs
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attention_metadata: Any,
|
||||
) -> torch.Tensor:
|
||||
"""Return values unchanged after recording the backend inputs."""
|
||||
backend_call["shapes"] = (query.shape, key.shape, value.shape)
|
||||
backend_call["metadata"] = attention_metadata
|
||||
return value
|
||||
|
||||
class RecordingAttentionBackend:
|
||||
"""Supply the recording implementation through the backend API."""
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type[RecordingAttentionImpl]:
|
||||
return RecordingAttentionImpl
|
||||
|
||||
with (
|
||||
patch("fastvideo.platforms.current_platform") as current_platform,
|
||||
patch(
|
||||
"fastvideo.models.vaes.minimax_h3_video.get_attn_backend",
|
||||
return_value=RecordingAttentionBackend,
|
||||
) as get_backend,
|
||||
):
|
||||
current_platform.is_cuda_alike.return_value = True
|
||||
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
|
||||
|
||||
attention.to_q = nn.Identity()
|
||||
attention.to_k = nn.Identity()
|
||||
attention.to_v = nn.Identity()
|
||||
attention.norm_q = nn.Identity()
|
||||
attention.norm_k = nn.Identity()
|
||||
attention.to_out[0] = nn.Identity()
|
||||
output = attention(torch.empty((1, 3, 8), device="meta"))
|
||||
|
||||
get_backend.assert_called_once_with(
|
||||
4,
|
||||
torch.bfloat16,
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
),
|
||||
)
|
||||
assert backend_call["init"] == {
|
||||
"num_heads": 2,
|
||||
"head_size": 4,
|
||||
"softmax_scale": 0.5,
|
||||
"num_kv_heads": 2,
|
||||
"causal": False,
|
||||
}
|
||||
assert backend_call["shapes"] == ((1, 3, 2, 4), ) * 3
|
||||
assert backend_call["metadata"] is None
|
||||
assert output.shape == (1, 3, 8)
|
||||
|
||||
|
||||
def test_video_attention_cpu_uses_torch_sdpa() -> None:
|
||||
"""Use PyTorch SDPA when H3 VAE attention receives CPU tensors."""
|
||||
with (
|
||||
patch("fastvideo.platforms.current_platform") as current_platform,
|
||||
patch("fastvideo.models.vaes.minimax_h3_video.get_attn_backend") as get_backend,
|
||||
):
|
||||
current_platform.is_cuda_alike.return_value = False
|
||||
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
|
||||
|
||||
get_backend.assert_not_called()
|
||||
assert attention.attn_impl is None
|
||||
|
||||
attention.to_q = nn.Identity()
|
||||
attention.to_k = nn.Identity()
|
||||
attention.to_v = nn.Identity()
|
||||
attention.norm_q = nn.Identity()
|
||||
attention.norm_k = nn.Identity()
|
||||
attention.to_out[0] = nn.Identity()
|
||||
hidden_states = torch.randn(1, 3, 8)
|
||||
query = hidden_states.unflatten(2, (2, 4)).permute(0, 2, 1, 3)
|
||||
expected = F.scaled_dot_product_attention(query, query, query).permute(0, 2, 1, 3).flatten(2, 3)
|
||||
|
||||
torch.testing.assert_close(attention(hidden_states), expected)
|
||||
|
||||
|
||||
def test_compile_with_conditions_selects_minimax_h3_video_decoder() -> None:
|
||||
"""Compile the registered video decoder with the VAE runtime kwargs."""
|
||||
assert not hasattr(MiniMaxH3VideoViTDecoder3d.forward, "get_compiler_config")
|
||||
_assert_dynamic_compile_selects_decoder(AutoencoderKLMiniMaxH3, MiniMaxH3VideoViTDecoder3d)
|
||||
|
||||
|
||||
def test_project_decoder_tile_uses_reduce_overhead_compile() -> None:
|
||||
"""Compile the per-tile decoder-input projection with CUDA Graph replay."""
|
||||
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._project_decoder_tile)
|
||||
|
||||
|
||||
def test_stitch_tiles_uses_reduce_overhead_compile() -> None:
|
||||
"""Compile spatial tile blending and concatenation with CUDA Graph replay."""
|
||||
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._stitch_tiles)
|
||||
|
||||
|
||||
def test_compile_with_conditions_selects_minimax_h3_audio_decoder() -> None:
|
||||
"""Compile the audio VAE decoder that the H3 waveform decode path calls."""
|
||||
_assert_dynamic_compile_selects_decoder(MiniMaxH3AudioVAE, MiniMaxH3AudioBigVGANDecoder)
|
||||
|
||||
|
||||
def test_decode_emits_indexed_temporal_chunk_ranges() -> None:
|
||||
"""Nest frame-segment ranges under each temporal decoder chunk range."""
|
||||
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
|
||||
vae.tokens_chunk_size = 1
|
||||
vae.token_overlap = 1
|
||||
vae.temporal_compression_ratio = 1
|
||||
vae.frame_pre_padding = 0
|
||||
vae.frame_overlap = 1
|
||||
vae.config = SimpleNamespace(token_drop=1)
|
||||
vae._decode_clip = Mock(return_value=torch.zeros((1, 1, 2, 1, 1)))
|
||||
range_events = []
|
||||
|
||||
@contextmanager
|
||||
def record_range(name: str):
|
||||
range_events.append(("enter", name))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_events.append(("exit", name))
|
||||
|
||||
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
|
||||
decoded = vae._decode(torch.zeros((1, 1, 2, 1, 1)))
|
||||
|
||||
assert decoded.shape == (1, 1, 3, 1, 1)
|
||||
assert vae._decode_clip.call_count == 2
|
||||
assert range_events == [
|
||||
("enter", "minimax_h3.vae.temporal_chunk.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.1"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.1"),
|
||||
]
|
||||
|
||||
|
||||
def test_decode_clip_no_spatial_tiling_stage_ranges() -> None:
|
||||
"""Separate untiled latent projection and decoder ranges."""
|
||||
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
|
||||
vae.use_tiling = False
|
||||
range_events = []
|
||||
vae.post_quant_conv = nn.Identity()
|
||||
vae.post_quant_conv.register_forward_hook(
|
||||
lambda _module, _args, _output: range_events.append(("call", "post_quant_conv")))
|
||||
vae.decoder = nn.Identity()
|
||||
vae.decoder.register_forward_hook(
|
||||
lambda _module, _args, _output: range_events.append(("call", "decoder_forward")))
|
||||
latent_clip = torch.zeros((1, 1, 1, 2, 2))
|
||||
|
||||
@contextmanager
|
||||
def record_range(name: str):
|
||||
range_events.append(("enter", name))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_events.append(("exit", name))
|
||||
|
||||
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
|
||||
decoded_clip = vae._decode_clip(latent_clip)
|
||||
|
||||
assert decoded_clip is latent_clip
|
||||
assert range_events == [
|
||||
("enter", "minimax_h3.vae.decode_clip"),
|
||||
("enter", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
|
||||
("call", "post_quant_conv"),
|
||||
("exit", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
|
||||
("enter", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
|
||||
("call", "decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip"),
|
||||
]
|
||||
|
||||
|
||||
def test_decode_clip_emits_tiled_stage_ranges() -> None:
|
||||
"""Nest indexed decoder tiles between tile-splitting and stitching ranges."""
|
||||
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
|
||||
vae.use_tiling = True
|
||||
vae.spatial_compression_ratio = 1
|
||||
vae.tile_sample_min_height = 1
|
||||
vae.tile_sample_min_width = 1
|
||||
vae.tile_sample_min_overlap_height = 0
|
||||
vae.tile_sample_min_overlap_width = 0
|
||||
vae._split_tiles = Mock(side_effect=[
|
||||
([0, 1], [1, 1], [0]),
|
||||
([0, 1], [1, 1], [0]),
|
||||
])
|
||||
vae.post_quant_conv = nn.Identity()
|
||||
vae._project_decoder_tile = Mock(side_effect=vae.post_quant_conv)
|
||||
vae.decoder = nn.Identity()
|
||||
stitched_clip = torch.zeros((1, 1, 1, 2, 2))
|
||||
vae._stitch_tiles = Mock(return_value=stitched_clip)
|
||||
range_events = []
|
||||
|
||||
@contextmanager
|
||||
def record_range(name: str):
|
||||
range_events.append(("enter", name))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_events.append(("exit", name))
|
||||
|
||||
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
|
||||
decoded_clip = vae._decode_clip(torch.zeros((1, 1, 1, 2, 2)))
|
||||
|
||||
assert decoded_clip is stitched_clip
|
||||
assert vae._split_tiles.call_count == 2
|
||||
assert vae._project_decoder_tile.call_count == 4
|
||||
assert vae._stitch_tiles.call_count == 1
|
||||
assert range_events == [
|
||||
("enter", "minimax_h3.vae.decode_clip"),
|
||||
("enter", "minimax_h3.vae.decode_clip.split_tiles"),
|
||||
("exit", "minimax_h3.vae.decode_clip.split_tiles"),
|
||||
("enter", "minimax_h3.vae.decode_clip.decode_tiles"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.0.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.0.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.0.1"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.0.1"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.1.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.1.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.1.1"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.1.1"),
|
||||
("exit", "minimax_h3.vae.decode_clip.decode_tiles"),
|
||||
("enter", "minimax_h3.vae.decode_clip.stitch_tiles"),
|
||||
("exit", "minimax_h3.vae.decode_clip.stitch_tiles"),
|
||||
("exit", "minimax_h3.vae.decode_clip"),
|
||||
]
|
||||
@@ -1,11 +1,40 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines import ForwardBatch
|
||||
from fastvideo.worker.gpu_worker import Worker
|
||||
from fastvideo.worker.gpu_worker import Worker, _log_cuda_device_uuid
|
||||
|
||||
|
||||
def test_cuda_device_uuid_receipt_is_disabled_without_nvtx_profiling(monkeypatch) -> None:
|
||||
"""Avoid NVIDIA property access during ordinary worker initialization."""
|
||||
get_device_properties = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
|
||||
monkeypatch.setattr(torch.cuda, "get_device_properties", get_device_properties)
|
||||
|
||||
_log_cuda_device_uuid(0, torch.device("cuda:0"))
|
||||
|
||||
get_device_properties.assert_not_called()
|
||||
|
||||
|
||||
def test_cuda_device_uuid_receipt_identifies_profiled_worker(monkeypatch) -> None:
|
||||
"""Bind one profiled worker rank to its NVIDIA device UUID in logs."""
|
||||
log_info = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "get_device_properties", lambda device: SimpleNamespace(uuid="device-uuid"))
|
||||
monkeypatch.setattr("fastvideo.worker.gpu_worker.logger.info", log_info)
|
||||
|
||||
_log_cuda_device_uuid(2, torch.device("cuda:0"))
|
||||
|
||||
log_info.assert_called_once_with(
|
||||
"Worker %d CUDA device UUID: GPU-%s",
|
||||
2,
|
||||
"device-uuid",
|
||||
local_main_process_only=False,
|
||||
)
|
||||
|
||||
|
||||
def _worker_returning(output_batch: ForwardBatch) -> Worker:
|
||||
|
||||
Reference in New Issue
Block a user