From a21676815b568921a45f9bf19104cd8125c97b6f Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 09:38:20 -0700 Subject: [PATCH 01/24] [wip]: exact-size pinned arenas for H3 offload and 4090 benchmarks --- fastvideo/hooks/layerwise_offload.py | 33 ++-- fastvideo/hooks/pinned_memory.py | 86 ++++++++++ .../basic/minimax_h3/minimax_h3_pipeline.py | 68 +++----- fastvideo/tests/hooks/test_pinned_memory.py | 107 ++++++++++++ .../benchmarks/minimax_h3_4090/bench_pod.py | 154 ++++++++++++++++++ .../minimax_h3_4090/download_model.py | 14 ++ .../minimax_h3_4090/prompts_1k.json | 4 + 7 files changed, 409 insertions(+), 57 deletions(-) create mode 100644 fastvideo/hooks/pinned_memory.py create mode 100644 fastvideo/tests/hooks/test_pinned_memory.py create mode 100644 scripts/benchmarks/minimax_h3_4090/bench_pod.py create mode 100644 scripts/benchmarks/minimax_h3_4090/download_model.py create mode 100644 scripts/benchmarks/minimax_h3_4090/prompts_1k.json diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 1356219b2f..bc9ac0903e 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -4,6 +4,7 @@ import torch from torch import nn from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager +from fastvideo.hooks.pinned_memory import PinnedTensorArena from fastvideo.logger import init_logger logger = init_logger(__name__) @@ -53,6 +54,7 @@ def __init__( self.cpu_named_parameters: dict[str, torch.Tensor] = {} self.module_ref: nn.Module = None # type: ignore self.device: torch.device = device + self.cpu_arena: PinnedTensorArena | None = None def _will_offload(self, name: str) -> bool: return True @@ -60,11 +62,22 @@ def _will_offload(self, name: str) -> bool: @torch.compiler.disable def on_init(self, module: nn.Module): self.module_ref = module + self.clear_cpu_storage() + self.cpu_arena = PinnedTensorArena( + (name, param) for name, param in _offload_tensors(module) if self._will_offload(name)) for name, param in _offload_tensors(self.module_ref): if self._will_offload(name): - self.cpu_named_parameters[name] = (param.data.detach().to("cpu").pin_memory()) + host = self.cpu_arena.empty_like(name, param) + host.copy_(param.data.detach()) + self.cpu_named_parameters[name] = host param.data = _tensor_placeholder(param.data, self.device) + def clear_cpu_storage(self) -> None: + self.cpu_named_parameters.clear() + if self.cpu_arena is not None: + self.cpu_arena.close() + self.cpu_arena = None + @torch.compiler.disable def wait_and_replace_params(self): torch.cuda.current_stream().wait_stream(self.async_copy_stream) @@ -108,16 +121,17 @@ def on_attach(self, module: nn.Module): self.state.on_init(module) # pyright: ignore def on_detach(self, module: nn.Module): + self.state.async_copy_stream.synchronize() named_parameters = dict(_offload_tensors(module, self.state.cpu_named_parameters)) for name, cpu_tensor in self.state.cpu_named_parameters.items(): - if name not in self.state.gpu_named_parameters: - if name in named_parameters: - named_parameters[name].data = cpu_tensor.to(device=self.state.device) - else: - logger.warning( - "Parameter {} not found in module during detachment.", - name, - ) + if name in named_parameters: + gpu_tensor = self.state.gpu_named_parameters.get(name) + named_parameters[name].data = gpu_tensor if gpu_tensor is not None else cpu_tensor.to(self.state.device) + else: + logger.warning("Parameter %s not found in module during detachment.", name) + self.state.gpu_named_parameters.clear() + self.state.clear_cpu_storage() + self.state.next_state = None @classmethod def name(cls) -> str: @@ -154,7 +168,6 @@ def mutate_params_scope(self): yield finally: # instead of releasing, we should overwrite the original params since they have been modified - self.state.cpu_named_parameters.clear() self.state.gpu_named_parameters.clear() self.state.on_init(self.state.module_ref) # pyright: ignore diff --git a/fastvideo/hooks/pinned_memory.py b/fastvideo/hooks/pinned_memory.py new file mode 100644 index 0000000000..81e7ff8af3 --- /dev/null +++ b/fastvideo/hooks/pinned_memory.py @@ -0,0 +1,86 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exact-size pinned host storage for inference offload.""" + +import weakref +from collections.abc import Iterable + +import torch + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) +_ALIGNMENT = 256 +_PAGE_ALIGNMENT = 4096 + + +def _unregister(buffer: torch.Tensor, device: int) -> None: + # The allocation must outlive any outstanding nonblocking H2D copies. + try: + with torch.cuda.device(device): + torch.cuda.synchronize() + error = torch.cuda.cudart().cudaHostUnregister(buffer.data_ptr()) + if error != 0: + logger.warning("cudaHostUnregister failed: %s", error) + except Exception as exc: + # CUDA may already be unavailable during interpreter shutdown. + logger.warning("Could not unregister pinned host arena: %s", exc) + + +class PinnedTensorArena: + """Pack tensors into one CUDA-registered CPU allocation, aligned to 256 bytes. + + PyTorch's pinned allocator rounds large allocations to powers of two. Registering + ordinary host storage avoids that overhead. Typed views retain this owner, so + registration survives even if the module or offload state is dropped first. + If registration is unavailable, allocate conventional pinned tensors instead. + Call ``close`` only after all views and pending copies have been released. + """ + + def __init__(self, tensors: Iterable[tuple[str, torch.Tensor]]) -> None: + self.offsets: dict[str, tuple[int, int]] = {} + size = 0 + for name, tensor in tensors: + size = (size + _ALIGNMENT - 1) // _ALIGNMENT * _ALIGNMENT + length = tensor.numel() * tensor.element_size() + self.offsets[name] = (size, length) + size += length + self.nbytes = size + self.buffer: torch.Tensor | None = None + self._finalizer: weakref.finalize | None = None + if not size: + return + # Register dedicated pages: small malloc allocations can otherwise share a + # registered page with another arena. The extra space is bounded by 8 KiB. + span = (size + _PAGE_ALIGNMENT - 1) // _PAGE_ALIGNMENT * _PAGE_ALIGNMENT + allocation = torch.empty(span + _PAGE_ALIGNMENT - 1, dtype=torch.uint8, device="cpu") + start = (-allocation.data_ptr()) % _PAGE_ALIGNMENT + # Give the aligned region its own storage base. Tensor.is_pinned() queries + # the storage pointer, which would precede the registered region for a + # plain narrow() view. The memoryview retains the original allocation. + buffer = torch.frombuffer(memoryview(allocation.numpy())[start:start + span], dtype=torch.uint8) + device = torch.cuda.current_device() + try: + error = torch.cuda.cudart().cudaHostRegister(buffer.data_ptr(), span, 0) + if error != 0: + raise RuntimeError(f"cudaHostRegister returned {error}") + except Exception as exc: + logger.warning("Exact-size host registration failed; using the pinned allocator: %s", exc) + return + self.buffer = buffer + self._finalizer = weakref.finalize(self, _unregister, buffer, device) + + def empty_like(self, name: str, tensor: torch.Tensor) -> torch.Tensor: + """Return a contiguous host view with the source's dtype and shape.""" + if self.buffer is None: + return torch.empty(tensor.shape, dtype=tensor.dtype, device="cpu", pin_memory=True) + offset, length = self.offsets[name] + host = self.buffer.narrow(0, offset, length).view(tensor.dtype).reshape(tensor.shape) + host._pinned_arena = self + return host + + def close(self) -> None: + """Unregister before releasing storage; safe to call more than once.""" + if self._finalizer is not None: + self._finalizer() + self._finalizer = None + self.buffer = None diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index c09e39b0e9..005f981a9c 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -20,6 +20,7 @@ from fastvideo.configs.pipelines.minimax_h3 import (FASTH3_INFERENCE_FILE, FASTH3_INFERENCE_SCHEMA, MiniMaxH3PipelineConfig) from fastvideo.fastvideo_args import FastVideoArgs +from fastvideo.hooks.pinned_memory import PinnedTensorArena from fastvideo.logger import init_logger from fastvideo.models.hf_transformer_utils import get_diffusers_config from fastvideo.pipelines.basic.minimax_h3.stages import ( @@ -76,64 +77,37 @@ def _checkpoint_has_vsa_gates(transformer_dir: Path) -> bool: return False -def _exact_pinned_views(tensors: list[torch.Tensor]) -> tuple[list[torch.Tensor], torch.Tensor | None]: - """Page-locked host copies of ``tensors`` backed by one exact-size allocation. - - ``Tensor.pin_memory()`` goes through torch's caching host allocator, which rounds every block up to a power - of two (1.76x for H3's packed FFN weights), so pinning a 20 GB DiT plus a 15 GB encoder overruns a 60 GB - container. Registering one plain allocation with ``cudaHostRegister`` pins exactly what is needed; the views - stay pinned and keep the arena alive. - """ - sizes = [-(-t.numel() * t.element_size() // 256) * 256 for t in tensors] - arena = torch.empty(max(sum(sizes), 1), dtype=torch.uint8) - cudart = torch.cuda.cudart() - if cudart.cudaHostRegister(arena.data_ptr(), arena.numel(), 0) != cudart.cudaError.success: - logger.warning("cudaHostRegister failed; falling back to torch pinned allocations") - return [t.detach().to("cpu").pin_memory() for t in tensors], None - views, offset = [], 0 - for tensor, size in zip(tensors, sizes, strict=True): - nbytes = tensor.numel() * tensor.element_size() - view = arena[offset:offset + nbytes].view(tensor.dtype).view(tensor.shape) - view.copy_(tensor) - views.append(view) - offset += size - return views, arena - - def _pinned_swap(module: Any, device: torch.device) -> None: - """Move a module's tensors between the GPU and persistent, exactly sized pinned host copies. + """Move a module's tensors between the GPU and a persistent pinned host copy. Inference weights never change, so a parameter's pinned copy is made once and parking just repoints the - parameter at it (no transfer); restoring is one pinned host-to-device copy. Buffers keep a persistent host - copy too and are copied back into it on every park, so mutable buffers stay correct without new allocations. + parameter at it (no transfer); restoring is one pinned host-to-device copy. Buffers are copied every time. """ store = module.__dict__.setdefault("_pinned_host_tensors", {}) - arenas = module.__dict__.setdefault("_pinned_host_arenas", []) params = dict(module.named_parameters()) - named = [(name, tensor) for name, tensor in list(params.items()) + list(module.named_buffers()) - if tensor is not None] + tensors = list(params.items()) + list(module.named_buffers()) if device.type == "cpu": - missing = [(name, tensor) for name, tensor in named if tensor.device.type != "cpu" and ( - name not in store or store[name].shape != tensor.shape or store[name].dtype != tensor.dtype)] - if missing: - views, arena = _exact_pinned_views([tensor.detach() for _, tensor in missing]) - store.update({name: view for (name, _), view in zip(missing, views, strict=True)}) - if arena is not None: - arenas.append(arena) - fresh = {name for name, _ in missing} - else: - fresh = set() - for name, tensor in named: + missing = [(name, tensor) for name, tensor in tensors + if tensor is not None and tensor.device.type != "cpu" and ( + name not in store or store[name].shape != tensor.shape or store[name].dtype != tensor.dtype)] + arena = PinnedTensorArena(missing) if missing else None + for name, tensor in tensors: + if tensor is None: + continue + if device.type == "cpu": if tensor.device.type == "cpu": continue - host = store[name] - if name not in params and name not in fresh: + host = store.get(name) + if host is None or host.shape != tensor.shape or host.dtype != tensor.dtype: + assert arena is not None + host = arena.empty_like(name, tensor) + host.copy_(tensor) + store[name] = host + elif name not in params: host.copy_(tensor) tensor.data = host - else: - for _, tensor in named: - if tensor.device != device: - tensor.data = tensor.data.to(device, non_blocking=True) + elif tensor.device != device: + tensor.data = tensor.data.to(device, non_blocking=True) if device.type != "cpu" and torch.cuda.is_available(): torch.cuda.current_stream(device).synchronize() diff --git a/fastvideo/tests/hooks/test_pinned_memory.py b/fastvideo/tests/hooks/test_pinned_memory.py new file mode 100644 index 0000000000..7c7e76d2a7 --- /dev/null +++ b/fastvideo/tests/hooks/test_pinned_memory.py @@ -0,0 +1,107 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CUDA host-registration lifetime and offload regressions (one GPU required).""" + +import gc +import weakref + +import pytest +import torch +from torch import nn + +from fastvideo.hooks.hooks import ModuleHookManager +from fastvideo.hooks.layerwise_offload import LayerwiseOffloadHook, LayerwiseOffloadState +from fastvideo.hooks.pinned_memory import PinnedTensorArena + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA host registration requires one GPU") + + +def test_arena_mixed_dtype_exact_copy_and_lifetime(): + sources = { + "weight": torch.randn(17, 33, device="cuda", dtype=torch.bfloat16), + "scale": torch.randn(17, 1, device="cuda", dtype=torch.float32), + "packed": torch.arange(513, device="cuda").to(torch.uint8), + "scalar": torch.tensor(3.0, device="cuda"), + "empty": torch.empty(0, 4, device="cuda"), + } + arena = PinnedTensorArena(sources.items()) + assert arena.buffer is not None, "This GPU must exercise registration, not fallback" + assert arena.buffer.numel() < sum(t.numel() * t.element_size() for t in sources.values()) + 4096 + 256 * len(sources) + hosts = {} + for name, source in sources.items(): + host = arena.empty_like(name, source) + host.copy_(source) + if host.numel(): + assert host.is_pinned() + assert (host.data_ptr() - arena.buffer.data_ptr()) % 256 == 0 or host.numel() == 0 + torch.testing.assert_close(host.to("cuda", non_blocking=True), source, rtol=0, atol=0) + hosts[name] = host + owner = weakref.ref(arena) + del arena + gc.collect() + assert owner() is not None, "Live typed views must retain the registration" + del hosts, host + gc.collect() + assert owner() is None + + +def test_registration_failure_uses_pinned_allocator(monkeypatch): + class RefusingRuntime: + + def cudaHostRegister(self, *_args): + return 1 + + monkeypatch.setattr(torch.cuda, "cudart", lambda: RefusingRuntime()) + source = torch.arange(27, dtype=torch.float32) + arena = PinnedTensorArena([("weight", source)]) + assert arena.buffer is None + host = arena.empty_like("weight", source) + host.copy_(source) + assert host.is_pinned() + torch.testing.assert_close(host, source, rtol=0, atol=0) + arena.close() + arena.close() + + +def test_offload_mutation_and_prefetched_detach(monkeypatch): + monkeypatch.setenv("FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS", "1") + module = nn.Linear(16, 16, device="cuda", dtype=torch.bfloat16) + module.register_buffer("packed", torch.arange(1 << 20, device="cuda").to(torch.uint8)) + expected = {name: tensor.clone() for name, tensor in list(module.named_parameters()) + list(module.named_buffers())} + state = LayerwiseOffloadState(torch.cuda.Stream(), torch.device("cuda")) + hook = LayerwiseOffloadHook(state) + manager = ModuleHookManager.get_from_or_default(module) + manager.append_forward_hook(hook) + old_buffer = state.cpu_arena.buffer + assert old_buffer.is_pinned() + with hook.mutate_params_scope(), torch.no_grad(): + module.weight.add_(1) + module.packed.add_(1) + assert not old_buffer.is_pinned(), "Reinitialization must unregister old storage" + expected["weight"].add_(1) + expected["packed"].add_(1) + state.prefetch_params() + new_buffer = state.cpu_arena.buffer + manager.remove_forward_hook(hook.name()) + assert not new_buffer.is_pinned(), "Detachment must unregister storage" + assert state.cpu_arena is None + assert not state.cpu_named_parameters and not state.gpu_named_parameters + for name, tensor in list(module.named_parameters()) + list(module.named_buffers()): + torch.testing.assert_close(tensor, expected[name], rtol=0, atol=0) + + +def test_h3_swap_reuses_host_storage_and_updates_buffers(): + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import _pinned_swap + + module = nn.Linear(16, 16, device="cuda", dtype=torch.bfloat16) + module.register_buffer("cache", torch.ones(11, device="cuda")) + expected_weight = module.weight.detach().clone() + _pinned_swap(module, torch.device("cpu")) + hosts = module._pinned_host_tensors + pointers = {name: tensor.data_ptr() for name, tensor in hosts.items()} + assert all(tensor.is_pinned() for tensor in hosts.values()) + _pinned_swap(module, torch.device("cuda", torch.cuda.current_device())) + module.cache.add_(3) + _pinned_swap(module, torch.device("cpu")) + assert pointers == {name: tensor.data_ptr() for name, tensor in hosts.items()} + torch.testing.assert_close(module.cache, torch.full((11,), 4.0), rtol=0, atol=0) + torch.testing.assert_close(module.weight, expected_weight.cpu(), rtol=0, atol=0) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py new file mode 100644 index 0000000000..417dcdc594 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -0,0 +1,154 @@ +"""Pruned FastH3 8-step benchmark on a single RTX 4090. + +usage: python bench_pod.py [--timed N] [--offload-buffers] [--no-layerwise] +Stage logs (FASTVIDEO_STAGE_LOGGING=1) carry per-stage time and memory peaks; results.json lands in outputs//. +""" +import argparse +import json +import os +import pathlib +import shlex +import statistics +import sys +import threading +import time + + +class HostMemoryPeak: + """Sample pod-wide cgroup usage; anon excludes cached checkpoint file pages.""" + + def __init__(self): + self.stop = threading.Event() + self.peak_bytes = 0 + self.peak_anon_bytes = 0 + self.thread = threading.Thread(target=self._sample, daemon=True) + + def _sample(self): + while not self.stop.is_set(): + try: + root = pathlib.Path("/sys/fs/cgroup") + self.peak_bytes = max(self.peak_bytes, int((root / "memory.current").read_text())) + stats = dict(line.split() for line in (root / "memory.stat").read_text().splitlines()) + self.peak_anon_bytes = max(self.peak_anon_bytes, int(stats["anon"])) + except (OSError, KeyError, ValueError): + return + self.stop.wait(0.1) + + def __enter__(self): + self.thread.start() + return self + + def __exit__(self, *_args): + self.stop.set() + self.thread.join() + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("name"); ap.add_argument("model"); ap.add_argument("quant", choices=("nvfp4", "fp8", "bf16")) + ap.add_argument("--timed", type=int, default=2) + ap.add_argument("--offload-buffers", action="store_true") + ap.add_argument("--no-layerwise", action="store_true") + ap.add_argument("--resident-encoder", action="store_true") + ap.add_argument("--tile-batch", default=None) + ap.add_argument("--frames", type=int, default=243) + ap.add_argument("--height", type=int, default=768) + ap.add_argument("--width", type=int, default=1344) + ap.add_argument("--warmup", type=int, default=1) + ap.add_argument("--prompt-file", type=pathlib.Path, default=pathlib.Path(__file__).with_name("prompts_1k.json")) + ap.add_argument("--output-root", type=pathlib.Path, default=pathlib.Path("/workspace/outputs")) + ap.add_argument("--no-vae-compile", action="store_true") + ap.add_argument("--adaln-cache", action="store_true") + ap.add_argument("--profile", action="store_true") + ap.add_argument("--sparsity", type=float, default=0.8) + ap.add_argument("--decode", default="h3-vae") + ap.add_argument("--lazy", action="store_true", help="lazy_module_load: reload released modules per request") + ap.add_argument("--prompts", default=None, help="comma-separated prompt ids (default: both)") + a = ap.parse_args() + if a.timed < 2 or a.warmup < 1: + ap.error("Use at least one warmup and two timed runs") + if not pathlib.Path(a.model, "fastvideo_inference.json").is_file(): + ap.error("The model directory must contain fastvideo_inference.json for the 8-step DMD contract") + + os.environ.setdefault("FASTVIDEO_STAGE_LOGGING", "1") + # NVFP4 is retained for other GPUs; the RTX 4090 DiT uses FP8. + if a.quant == "nvfp4": + os.environ.setdefault("FASTVIDEO_H3_VSA_FP4", "1") + os.environ.setdefault("FASTVIDEO_NVFP4_MM_BACKEND", "cutlass") + os.environ.setdefault("FASTVIDEO_MINIMAX_H3_FUSIONS", "all") + os.environ.setdefault("FASTVIDEO_H3_VAE_TILE_BATCH", "28") + os.environ.setdefault("FASTVIDEO_VSA_TRITON", "1") + os.environ.setdefault("FASTVIDEO_VSA_SM100A", "0") + os.environ.setdefault("FASTVIDEO_FA4", "0") + os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") + ap_table = os.environ.get("FASTVIDEO_H3_ADALN_TABLE") + if a.profile: + os.environ["FASTVIDEO_H3_SP_PROFILE"] = "1" + if a.adaln_cache and not ap_table: + os.environ["FASTVIDEO_H3_ADALN_CACHE"] = "1" + # Export the exact modulation tables (and inputs) for table-only loads and low-rank experiments. + os.environ.setdefault("FASTVIDEO_H3_ADALN_DUMP", f"/workspace/adaln_tables_{a.name}.pt") + if a.offload_buffers: + os.environ["FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS"] = "1" + if a.tile_batch: + os.environ["FASTVIDEO_H3_VAE_TILE_BATCH"] = a.tile_batch + import torch + from fastvideo import VideoGenerator + + texts = json.loads(a.prompt_file.read_text()) + layerwise = not a.no_layerwise + engine = {"num_gpus": 1, "use_fsdp_inference": False, + "parallelism": {"tp_size": 1, "sp_size": 1}, + "offload": {"dit": False, "dit_layerwise": layerwise, "text_encoder": not a.resident_encoder, + "vae": layerwise, "pin_cpu_memory": True, "lazy_module_load": a.lazy}, + "compile": {"enabled": False, "vae_enabled": not a.no_vae_compile}} + if a.quant == "nvfp4": + engine["quantization"] = {"transformer_quant": "NVFP4", "layer_profile": "h3_dit_vsa"} + elif a.quant == "fp8": + engine["quantization"] = {"transformer_quant": "FP8"} + experimental = {"attention_backend": "VIDEO_SPARSE_ATTN_H3", "VSA_sparsity": a.sparsity, "VSA_tile_size": 64, + "h3_sequential_load": not a.resident_encoder, "inference_torch_compile": False, + "vae_parallel_decode": False, "video_decode_backend": a.decode} + config = {"model_path": a.model, "engine": engine, "pipeline": {"experimental": experimental}} + out_dir = a.output_root / a.name + out_dir.mkdir(parents=True, exist_ok=True) + sampling = {"seed": 20260929, "height": a.height, "width": a.width, "num_frames": a.frames, "fps": 24, + "num_inference_steps": 9, "guidance_scale": 1.0, "batch_cfg": False} + results = {"name": a.name, "quant": a.quant, "command": shlex.join([sys.executable, "-P", *sys.argv]), + "env": {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_")) + or k in ("CUDA_VISIBLE_DEVICES", "MAX_JOBS")}, + "torch": torch.__version__, "cuda": torch.version.cuda, + "source_commit": os.environ.get("FASTVIDEO_SOURCE_COMMIT"), + "gpu": torch.cuda.get_device_name(0), "config": config, "sampling": sampling, "runs": []} + (out_dir / "results.json").write_text(json.dumps(results, indent=2)) + t0 = time.perf_counter() + generator = VideoGenerator.from_config(config) + results["load_s"] = round(time.perf_counter() - t0, 1) + ids = a.prompts.split(",") if a.prompts else list(texts) + order = [ids[i % len(ids)] for i in range(a.warmup + a.timed)] + try: + for i, pid in enumerate(order): + request = {"prompt": texts[pid], "negative_prompt": "", + "sampling": sampling, + "output": {"output_path": str(out_dir / f"{i:02d}_{pid}.mp4"), "save_video": True, + "return_frames": False}} + t = time.perf_counter() + with HostMemoryPeak() as host_peak: + generator.generate(request) + wall = round(time.perf_counter() - t, 2) + results["runs"].append({"prompt": pid, "warmup": i < a.warmup, "wall_s": wall, + "clip": request["output"]["output_path"], + "peak_host_cgroup_gib": round(host_peak.peak_bytes / 2**30, 3), + "peak_host_anon_gib": round(host_peak.peak_anon_bytes / 2**30, 3)}) + timed = [run["wall_s"] for run in results["runs"] if not run["warmup"]] + if timed: + results["median_e2e_s"] = statistics.median(timed) + print("RUN", json.dumps(results["runs"][-1]), flush=True) + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + finally: + generator.shutdown() + print("DONE", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_4090/download_model.py b/scripts/benchmarks/minimax_h3_4090/download_model.py new file mode 100644 index 0000000000..d76718a028 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/download_model.py @@ -0,0 +1,14 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Download the private FP8 model without exposing its credential.""" + +import os +from pathlib import Path + +from huggingface_hub import snapshot_download + +os.environ.setdefault("HF_XET_HIGH_PERFORMANCE", "1") +snapshot_download( + "FastVideo/FastH3-Pruned-8Step-FP8-ckpt300", + local_dir="/workspace/vol/pruned_fp8_300", + token=Path("/root/.hf-fastvideo/token").read_text().strip(), +) diff --git a/scripts/benchmarks/minimax_h3_4090/prompts_1k.json b/scripts/benchmarks/minimax_h3_4090/prompts_1k.json new file mode 100644 index 0000000000..553a5a7d67 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/prompts_1k.json @@ -0,0 +1,4 @@ +{ + "latency-ceramics-005": "In a quiet pottery studio an adult potter steadies a small spinning bowl while an adult apprentice watches. The apprentice asks, \"Is the rim ready?\" The potter says, \"One more gentle pass,\" and smooths the lip with a damp sponge. Begin with a close view of the hands, then make one restrained cut to a shoulder-level view showing both faces. The wheel hum, damp clay, a light splash and breathing form the soundscape. The movement is careful and unhurried, with no background music. The entire event is one finishing pass on the same bowl, not a demonstration of the whole pottery process.\nThe apprentice's apron pocket is indigo, with a plain surface. This detail belongs to the existing object, stays at its established location and remains subordinate to the main action. Preserve its material and appearance through the camera movement.\nThe wheel sits near a tall north-facing window on the left wall of the studio.\nIts circular splash pan is at the seated potter's waist.\nThe apprentice stands beyond the right side of the wheel, with enough room between their clothes and the wet rim to avoid accidental contact.\nBehind the potter, a shallow shelf holds a few dry cups.\nA sink is farther back on the right, beneath a small rack of towels.\nThese positions remain fixed through the cut.\nThe wider view reveals the same work area that surrounds the hands in the opening close view.\nThe bowl is modest in size, comfortably held between two hands.\nIts side rises from a low foot in a continuous shallow curve and ends in a slightly thickened lip.\nThe clay is warm gray with a brown undertone, darker where it is wet.\nA narrow spiral line below the rim records the potter's earlier touch.\nThat line rotates with the bowl, while a stationary highlight from the window travels over the passing surface.\nKeep these two motions distinguishable: the material marks belong to the clay, and the reflected light belongs to the relationship between the surface and the window.\nOpen with the camera low enough to see the bowl's interior without looking directly down onto the wheel.\nThe potter's left fingertips support the inside edge.\nThe right hand holds a small natural sponge outside the lip.\nNeither hand blocks the whole form.\nThe near edge of the splash pan occupies the bottom of the composition as a soft curved boundary.\nBeyond the hands, the apprentice's apron is initially out of focus.\nThis arrangement gives the close shot depth and prepares the later view of the two people without requiring another establishing shot.\nThe left hand is already stable when the clip begins.\nIts fingers form a loose supportive curve rather than squeezing the wall.\nThe sponge approaches the outside edge with only a small adjustment of the wrist.\nAs it contacts the clay, it compresses slightly and darkens where moisture gathers.\nThe bowl continues to rotate at a steady moderate speed.\nShow the finishing pass as a change in the surface's smoothness and the evenness of the lip, not as a large change in the bowl's overall shape.\nThe work is nearly finished before this moment begins.\nThe apprentice asks the question while looking at the rim, then briefly lifts their eyes toward the potter.\nTheir hands rest loosely together in front of the apron, safely away from the rotating work.\nThe question is curious and quiet, with the natural upward inflection of someone checking a detail.\nThe potter answers without stopping the wheel or turning their whole body.\nA small glance toward the apprentice is sufficient before attention returns to the clay.\nKeep the spoken words exactly as given, with no narrator explaining the technique and no extra exchange after the answer.\nCut once after the question has made the apprentice's presence clear.\nThe shoulder-level view places the potter to the left and the apprentice to the right, preserving the established relation to the wheel.\nThe bowl remains visible between them in the lower part of the frame.\nThe potter's right hand still holds the same sponge at the same point on the rim.\nContinue the wheel sound across the cut without a restart.\nThe change of view should feel like a closer understanding of the same instant, not a jump forward to a later stage of the work.\nThe potter wears a practical cotton work shirt with the sleeves rolled above the wrists.\nThe folds gather at the elbows and remain dry there.\nSmall clay marks on the forearms and apron are concentrated near the work area.\nThey do not spread or migrate during the pass.\nThe apprentice's apron is cleaner but shows a few dry pale smudges near one pocket.\nBoth garments have weight and ordinary creases.\nAvoid pristine costumes or exaggerated distressing; this is a used studio where people work carefully and clean their tools regularly.\nGive the potter a focused, patient expression.\nTheir mouth moves only for the brief reply, then settles while they feel the rim through the sponge.\nAllow the last small movement to settle within the established composition, with the environmental sound continuing around it. Preserve the quiet final composition.", + "latency-harbor-005": "## Harbor: the tide chart\nThe fabric cover around the folded tide chart is indigo, with a plain surface. This detail belongs to the existing object, stays at its established location and remains subordinate to the main action. Preserve its material and appearance through the camera movement.\nThe event takes place beside a small passenger ferry tied to a working harbor pier just before sunrise. An adult mechanic stands on the ferry's open side deck, and its captain stands beside the cabin entrance. The mechanic offers a folded tide chart and says, \"The channel is clear.\" The captain accepts it, answers, \"Then we can go,\" and looks out toward the harbor entrance. A slow lateral camera move reveals the channel beyond their shoulders. The boat remains moored throughout this brief exchange. Close voices, water against the hull, a loose halyard and a distant gull make the soundscape; there is no music.\nThe ferry is a practical coastal launch with a dark blue hull and a narrow cream band beneath its windows.\nIt carries a small enclosed cabin forward and an open passenger area behind it.\nThe camera is on the open deck, looking diagonally toward the cabin and the gap between the two people.\nThis angle places the pier along the left edge of the view and open water farther to the right.\nThe horizon is low enough that the upper part of the cabin has a clear silhouette against the pale sky.\nNothing in the composition suggests that the ferry is already underway.\nThe mechanic has finished a routine inspection rather than an emergency repair.\nTheir expression is alert but comfortable, with the slight tiredness of an early start.\nThey wear a plain work jacket over a warm shirt and carry no conspicuous badge or brand.\nA few old creases in the jacket show where the elbows bend.\nThe sleeve nearest the chart has a darker damp patch near its cuff from resting against the rail.\nKeep that patch in the same place as the arm moves.\nThe mechanic's free hand rests lightly on the top of a closed tool bag at hip level.\nThe captain is a different adult, dressed for a cool morning outside.\nA heavy knit sweater is visible beneath an open weatherproof coat.\nTheir hair is tidy but not freshly styled, and the light catches a few loose strands when they turn toward the water.\nTheir stance is balanced on the gently moving deck, with one foot slightly ahead of the other.\nThey are listening to the mechanic before the first line begins.\nThe captain does not interrupt or make a broad theatrical gesture.\nTheir reply is a small decision shared between people accustomed to working together.\nThe tide chart is a real paper object with several old folds.\nIt is partly folded into a rectangle that can be held in one hand, but one narrow flap remains loose.\nFaint printed lines and numbers are visible as a texture on its pale surface without becoming a readable title or a map inset.\nA soft graphite mark near one fold suggests that it has been used for planning.\nThe mechanic holds its lower edge between the thumb and fingers, keeping the paper clear of the damp rail.\nIts upper corner lifts slightly in the breeze before the captain takes it.\nBegin with both people already in the frame.\nThe mechanic's hand and the chart occupy the space between their bodies, below their faces.\nThis arrangement lets the first line and the handover belong to the same view.\nAs the mechanic speaks, the chart moves a short distance toward the captain.\nThe motion is neither a flourish nor an abrupt thrust.\nThe captain's receiving hand rises from beside the coat, touches the opposite edge and supports it before the mechanic releases their grip.\nThe paper bends a little between the two hands during that shared moment of support.\nThe mechanic's line is spoken in an ordinary low voice suitable for the quiet morning. The consonants remain clear, but the delivery does not sound like a public announcement. Their mouth and jaw make the small movements of the exact words, and their eyes remain on the captain. There is a slight release of breath after \"clear.\" The captain acknowledges the information with a very small nod before replying. The pause is long enough to register listening and short enough that the exchange feels familiar. Do not add another question, greeting or explanation of the voyage.\nWhen the captain says, \"Then we can go,\" the first part of the line is addressed to the mechanic. On the last words, their gaze begins to move toward the channel. The head follows the eyes through a modest turn, revealing more of the cheek nearest the exterior light. The chart settles against the front of the coat, still visibly held. The mechanic follows the captain's look with a quieter change of attention. Both remain in place. The ending is anticipation of departure, not departure itself: no engine surge, released rope or sudden movement of the ferry is needed.\nThe camera makes a restrained lateral movement toward the open-water side of the deck.\nKeep the final gesture restrained and preserve the surrounding atmosphere. Preserve the quiet final composition as the scene reaches its stated resolution." +} \ No newline at end of file From a491f15e6f4c702bea517101bfa29e86cd375b67 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 09:43:43 -0700 Subject: [PATCH 02/24] [wip]: document 4090 setup validation and reproducible baselines --- scripts/benchmarks/minimax_h3_4090/README.md | 81 ++++++++++++++++++++ 1 file changed, 81 insertions(+) create mode 100644 scripts/benchmarks/minimax_h3_4090/README.md diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md new file mode 100644 index 0000000000..24be1c201d --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -0,0 +1,81 @@ +# FastH3 on a single RTX 4090 + +Use the private `FastVideo/FastH3-Pruned-8Step-FP8-ckpt300` checkpoint. +Keep its `fastvideo_inference.json`: nine sigma-grid points produce eight +DMD forwards. Preserve VSA sparsity 0.8 and tile size 64. + +## Setup and validation + +The October 3, 2026 pod has one RTX 4090 (24,564 MiB), driver 580.126.20, +a 99,999,997,952-byte host cgroup limit, and 150 GB disk. Its runtime is +PyTorch 2.12.0+cu126, CUDA toolkit 12.6, and FlashInfer 0.7.1rc2. + +Exact-size host arenas replace the pinned allocator for layerwise blocks +and H3 module swaps. Each arena packs typed views at 256-byte offsets into +dedicated CUDA-registered pages. Live views retain their registration owner. +Mutation and hook detachment unregister old arenas. Registration failures +fall back to the ordinary pinned allocator with a warning. + +Validation on this pod: all 12 tests passed with the command below. +The tests cover repeated offloaded forwards, BF16, mixed-dtype exact copies, +owner lifetime, registration fallback, large-buffer mutation, detachment +after prefetch, and persistent H3 swaps with changing buffers. + +```bash +source /workspace/env.sh +source /workspace/venv/bin/activate +cd /workspace/fastvideo +python -P -m pytest fastvideo/tests/hooks/test_pinned_memory.py \ + fastvideo/tests/hooks/test_layerwise_offload.py -q +``` + +The supplied `handoff_4090/kernel_microbench/pinned_memory.py` measured +5.06 GiB extra cgroup usage for 2.87 GiB through `pin_memory()`, versus +2.87 GiB using direct host registration. Pinned H2D measured 25.9 GB/s; +pageable H2D measured 10.2 GB/s. These are microbenchmarks, not clip timings. + +The supplied `handoff_4090/kernel_microbench/t_fp8.py` measured: + +| K → N, M = 38,976 | Fused quantization | FP8 GEMM + scale epilogue | +| --- | --- | --- | +| 5376 → 5376 | 0.68 ms | 11.27 ms | +| 5376 → 28672 | 0.68 ms | 40.39 ms | +| 14336 → 5376 | 2.10 ms | 27.96 ms | + +Commands on the pod, run before clip benchmarks: + +```bash +cd /workspace +python -P kernel_microbench/pinned_memory.py +python -P kernel_microbench/t_fp8.py +``` + +## End-to-end baseline + +Run one warmup and at least two timed requests. The benchmark saves clips, +the exact Python command, runtime environment, sampling geometry, config, +source commit, each wall time, and the median to `results.json`. Stage logs +include GPU allocation peaks and conditioning, denoise, and decode timings. +Per-run host samples report both total cgroup memory and anonymous memory; +total usage includes checkpoint file cache. Both are sampled every 100 ms +and include all processes in the pod's cgroup. + +```bash +cd /workspace +export FASTVIDEO_SOURCE_COMMIT=ce877b200 +export FASTVIDEO_H3_PARK_MODULES=vae,audio_vae +export MAX_JOBS=4 +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + baseline-480p /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --lazy --no-vae-compile \ + --height 480 --width 832 --frames 243 --timed 2 +``` + +For the full-resolution baseline, change the name to `baseline-768p`, +height to 768, and width to 1344; retain 243 frames. Layerwise offload, +lazy component loading, and eager VAE decode are the starting configuration. +Do not compare these numbers with a different frame count or decoder. + +After a baseline works, measure FFN chunk sizes 16,384 and 8,192, then +increase resident DiT blocks within the measured GPU budget. Kernel or +decoder changes also need same-seed visual and auditory comparison. From 4ae791c3df9f4646afb895dd6371099c998eb225 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 09:46:18 -0700 Subject: [PATCH 03/24] [wip]: record checkpoint revision and GPU driver in 4090 benchmark results --- scripts/benchmarks/minimax_h3_4090/bench_pod.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index 417dcdc594..dc4349abad 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -9,6 +9,7 @@ import pathlib import shlex import statistics +import subprocess import sys import threading import time @@ -112,12 +113,20 @@ def main(): config = {"model_path": a.model, "engine": engine, "pipeline": {"experimental": experimental}} out_dir = a.output_root / a.name out_dir.mkdir(parents=True, exist_ok=True) + model_root = pathlib.Path(a.model) + revision_file = model_root / ".cache/huggingface/download/fastvideo_inference.json.metadata" + model_revision = revision_file.read_text().splitlines()[0] if revision_file.is_file() else None + hardware = subprocess.check_output( + ["nvidia-smi", "--query-gpu=name,memory.total,driver_version,pci.bus_id", "--format=csv,noheader"], text=True + ).strip() sampling = {"seed": 20260929, "height": a.height, "width": a.width, "num_frames": a.frames, "fps": 24, "num_inference_steps": 9, "guidance_scale": 1.0, "batch_cfg": False} results = {"name": a.name, "quant": a.quant, "command": shlex.join([sys.executable, "-P", *sys.argv]), "env": {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_")) or k in ("CUDA_VISIBLE_DEVICES", "MAX_JOBS")}, "torch": torch.__version__, "cuda": torch.version.cuda, + "hardware": hardware, "model_revision": model_revision, + "model_contract": json.loads((model_root / "fastvideo_inference.json").read_text()), "source_commit": os.environ.get("FASTVIDEO_SOURCE_COMMIT"), "gpu": torch.cuda.get_device_name(0), "config": config, "sampling": sampling, "runs": []} (out_dir / "results.json").write_text(json.dumps(results, indent=2)) From 21f98998da2b464cab732992cafc0523705d3a3d Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 09:49:48 -0700 Subject: [PATCH 04/24] [wip]: record completed 480p baseline and stage summary tooling --- scripts/benchmarks/minimax_h3_4090/README.md | 22 +++++++ .../benchmarks/minimax_h3_4090/summarize.py | 61 +++++++++++++++++++ 2 files changed, 83 insertions(+) create mode 100644 scripts/benchmarks/minimax_h3_4090/summarize.py diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 24be1c201d..f038c1811e 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -76,6 +76,28 @@ height to 768, and width to 1344; retain 243 frames. Layerwise offload, lazy component loading, and eager VAE decode are the starting configuration. Do not compare these numbers with a different frame count or decoder. +Completed baseline at source commit `ce877b200`, checkpoint revision +`f2ef54f9ff2091762ab8689b6514dcab5bc1d383`: + +| Configuration | Median e2e | Denoise stage | Video decode stage | Peak GPU allocated | Peak host anon | Peak total cgroup | +| --- | --- | --- | --- | --- | --- | --- | +| FP8, layerwise, lazy, eager H3 VAE, 832×480, 243 frames | 163.34 s | 91.72 s | 34.54 s | 17.47 GiB | 28.91 GiB | 76.95 GiB | + +Two timed requests took 163.67 s and 163.00 s after one warmup. +Stage times are medians and include deferred loading. Memory columns are +maxima across the timed requests. Total cgroup usage includes file cache. +The GPU peak occurs during conditioning. The saved clip contains 243 frames +at 24 fps (10.125 seconds) and an AAC audio track. A contact-sheet inspection +confirms a coherent pottery scene; speech and same-seed reference parity +still need review before claiming quality equivalence. + +Summarize a completed run while excluding warmup: + +```bash +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/summarize.py \ + /workspace/outputs/baseline-480p/results.json /workspace/baseline-480p.log +``` + After a baseline works, measure FFN chunk sizes 16,384 and 8,192, then increase resident DiT blocks within the measured GPU budget. Kernel or decoder changes also need same-seed visual and auditory comparison. diff --git a/scripts/benchmarks/minimax_h3_4090/summarize.py b/scripts/benchmarks/minimax_h3_4090/summarize.py new file mode 100644 index 0000000000..cffaa7988e --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/summarize.py @@ -0,0 +1,61 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Join benchmark results with per-request stage logs, excluding warmup runs.""" + +import argparse +import json +import re +import statistics +from pathlib import Path + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("results", type=Path) + parser.add_argument("log", type=Path) + args = parser.parse_args() + results = json.loads(args.results.read_text()) + completed = [] + stages = {} + peaks = [] + for line in args.log.read_text().splitlines(): + timing = re.search(r"\[(\w+)_stage\|[^\]]+\] Execution completed in ([\d.]+) ms", line) + if timing: + stages[timing[1]] = float(timing[2]) / 1000 + peak = re.search(r"Memory peak_allocated=([\d.]+) GiB", line) + if peak: + peaks.append(float(peak[1])) + if line.startswith("RUN "): + run = json.loads(line[4:]) + run["stage_s"] = stages + run["peak_gpu_allocated_gib"] = max(peaks) if peaks else None + completed.append(run) + stages = {} + peaks = [] + if len(completed) != len(results["runs"]): + raise ValueError("Log and results.json have different completed run counts") + timed = [run for run in completed if not run["warmup"]] + if len(timed) < 2: + raise ValueError("At least two timed runs are required for a baseline summary") + for run in timed: + if not {"denoising", "video_decoding", "audio_decoding"}.issubset(run["stage_s"]): + raise ValueError("Missing stage timings for a completed run") + summary = { + "name": results["name"], + "sampling": results["sampling"], + "source_commit": results["source_commit"], + "timed_runs": len(timed), + "median_e2e_s": statistics.median(run["wall_s"] for run in timed), + "median_stage_s": {name: statistics.median(run["stage_s"][name] for run in timed) + for name in ("conditioning", "denoising", "video_decoding", "audio_decoding")}, + "peak_gpu_allocated_gib": max(run["peak_gpu_allocated_gib"] for run in timed), + "peak_host_anon_gib": max(run["peak_host_anon_gib"] for run in timed), + "peak_host_cgroup_gib": max(run["peak_host_cgroup_gib"] for run in timed), + "notes": "Stage times include deferred component loading. Host peaks are pod-wide samples every 100 ms.", + "runs": timed, + } + args.results.with_name(f"{args.results.stem}-summary.json").write_text(json.dumps(summary, indent=2) + "\n") + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() From c716408ca4abb80a2259040a0441077d6e17fd5d Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:13:46 -0700 Subject: [PATCH 05/24] [perf]: reduce H3 attention activation copies and share FP8 input quantization --- fastvideo/models/dits/minimax_h3.py | 53 ++++++++--- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 71 ++++++++++++-- .../test_minimax_h3_tile_first.py | 94 +++++++++++++++++++ 3 files changed, 199 insertions(+), 19 deletions(-) create mode 100644 fastvideo/tests/transformers/test_minimax_h3_tile_first.py diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index ccda038983..bf8b226be1 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -37,7 +37,7 @@ from fastvideo.logger import init_logger from fastvideo.models.dits.base import BaseDiT from fastvideo.models.dits.minimax_h3_vsa_fp4 import (STAGES, vsa_fp4_attention, vsa_fp4_attention_sp, - vsa_fp4_requested) + vsa_fp4_requested, vsa_tile_first_attention) from fastvideo.models.dits.minimax_h3_fusions import ( HAVE_TRITON, fused_qknorm_rope, @@ -242,6 +242,7 @@ def __init__( # kernel (see minimax_h3_vsa_fp4); grad and compile keep the generic path. self._layer_idx = layer_idx_from_prefix(prefix, default=-1) self._vsa_fp4 = use_vsa and vsa_fp4_requested() + self._vsa_tile_first = use_vsa and os.environ.get("FASTVIDEO_H3_VSA_TILE_FIRST", "0") == "1" self.to_gate_compress: ReplicatedLinear | None = None # None = unchecked; the first forward tests the loaded weight once and # skips the gate branch entirely while it is structurally zero. @@ -336,9 +337,26 @@ def forward( with STAGES.span("out_proj"): hidden_states, _ = self.to_out(hidden_states) return hidden_states - query, _ = self.to_q(hidden_states) - key, _ = self.to_k(hidden_states) - value, _ = self.to_v(hidden_states) + if (self._vsa_tile_first and hidden_states.is_cuda and rotary_emb is not None + and not torch.is_grad_enabled() and not torch.compiler.is_compiling() + and (not model_parallel_is_initialized() or get_sp_world_size() == 1)): + meta = get_forward_context().attn_metadata + if isinstance(meta, MiniMaxH3VSAMetadata) and meta.tile_elems == 64: + use_fused_rope = self.fuse_qknorm_rope and _can_run_minimax_h3_fusion(hidden_states) + hidden_states = vsa_tile_first_attention(self, hidden_states, rotary_emb, meta, use_fused_rope) + with STAGES.span("out_proj"): + hidden_states, _ = self.to_out(hidden_states) + return hidden_states + with STAGES.span("qkv_proj"): + # All three projections see the same activations. Reuse their FP8 + # quantization when the loaded methods have identical granularity. + from fastvideo.models.dits.minimax_h3_vsa_fp4 import _shared_input_projections + if not torch.is_grad_enabled() and not torch.compiler.is_compiling(): + query, key, value = _shared_input_projections((self.to_q, self.to_k, self.to_v), hidden_states) + else: + query, _ = self.to_q(hidden_states) + key, _ = self.to_k(hidden_states) + value, _ = self.to_v(hidden_states) query = query.unflatten(-1, (self.num_attention_heads, self.attention_head_dim)) key = key.unflatten(-1, (self.num_attention_heads, self.attention_head_dim)) value = value.unflatten(-1, (self.num_attention_heads, self.attention_head_dim)) @@ -362,16 +380,18 @@ def forward( gate_compress, _ = self.to_gate_compress(hidden_states) extra_attention_kwargs["gate_compress"] = gate_compress.unflatten( -1, (self.num_attention_heads, self.attention_head_dim)) - hidden_states, _ = self.distributed_attention( - query, - key, - value, - original_seq_len=original_seq_len, - freqs_cis=None, - **extra_attention_kwargs, - ) + with STAGES.span("attention"): + hidden_states, _ = self.distributed_attention( + query, + key, + value, + original_seq_len=original_seq_len, + freqs_cis=None, + **extra_attention_kwargs, + ) hidden_states = hidden_states.flatten(2, 3).type_as(query) - hidden_states, _ = self.to_out(hidden_states) + with STAGES.span("out_proj"): + hidden_states, _ = self.to_out(hidden_states) return hidden_states @@ -644,6 +664,9 @@ def forward( 1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices) with nvtx_range("minimax_h3.transformer_block.self_attention"): attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len) + # The attention input is dead now. Keeping it until assignment below + # overlaps three full-width activations during the residual fusion. + del norm_hidden_states if use_modulate_fusion: with nvtx_range("minimax_h3.transformer_block.modulate_fusion"): hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate( @@ -662,6 +685,7 @@ def forward( norm_hidden_states = self.norm2(hidden_states) norm_hidden_states = norm_hidden_states * ( 1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices) + del attention_output with nvtx_range("minimax_h3.transformer_block.feed_forward"), STAGES.span("feed_forward"): feed_forward_output = self.ff(norm_hidden_states) if use_modulate_fusion and not torch.compiler.is_compiling(): @@ -1135,6 +1159,9 @@ def forward( # The eager driver owns profiling markers while each block's compiled # forward owns the graph that the marker surrounds. for block_index, block in enumerate(self.transformer_blocks): + if STAGES.enabled: + logger.info("H3_MEMORY_BLOCK %d allocated=%.3f GiB reserved=%.3f GiB", block_index, + torch.cuda.memory_allocated() / 2**30, torch.cuda.memory_reserved() / 2**30) with nvtx_range(f"minimax_h3.transformer_block.{block_index}"), STAGES.span("block_total"): packed_hidden_states = block( packed_hidden_states, diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index 5429f76a3f..b3ab154524 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -123,22 +123,81 @@ def _layout_for(meta: MiniMaxH3VSAMetadata, rotary_emb: tuple[torch.Tensor, torc def _shared_input_projections(linears: tuple[Any, ...], x: torch.Tensor) -> list[torch.Tensor]: - """Run projections of one input, quantizing it once when all are NVFP4 with the unit activation scale. + """Share compatible FP8 preparation or unit-scale NVFP4 activation quantization. - Only then is one quantized copy exactly what each layer would have produced; layers with a calibrated - or dynamic activation scale quantize their own input. + Calibrated or dynamic NVFP4 activation scales retain independent preparation. """ from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod + from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod methods = [linear.quant_method for linear in linears] - if not all( - type(m) is NVFP4QuantizeMethod and m.wants_prequantized_input() and m.uses_unit_activation_scale(linear) - for m, linear in zip(methods, linears, strict=True)): + same_nvfp4 = all(type(m) is NVFP4QuantizeMethod and m.wants_prequantized_input() + and m.uses_unit_activation_scale(linear) + for m, linear in zip(methods, linears, strict=True)) + same_fp8 = all(type(m) is FP8QuantizeMethod and m.granularity == methods[0].granularity for m in methods) + if not (same_nvfp4 or same_fp8): return [linear(x)[0] for linear in linears] pre = methods[0].quantize_input(x) return [m.apply(linear, x, linear.bias, pre_quantized=pre) for m, linear in zip(methods, linears, strict=True)] +def vsa_tile_first_attention(attn: Any, hidden_states: torch.Tensor, + rotary_emb: tuple[torch.Tensor, torch.Tensor], + meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> torch.Tensor: + """Single-rank BF16 VSA with one input scatter instead of a Q/K/V/gate stack. + + The existing backend computes the same tile-64 mask, valid-key handling, + fine attention and compression branch. Bias-free projections keep pad + rows zero. This path is inference-only and keeps the checkpoint layout. + """ + layout = _layout_for(meta, rotary_emb) + logical = layout.n_tiles * layout.tile + heads, dim = attn.num_attention_heads, attn.attention_head_dim + with STAGES.span("tile_input"): + x_tiles = layout.gather_in(hidden_states)[:, :logical] + with STAGES.span("qkv_proj"): + query, key, value = (t.unflatten(-1, (heads, dim)) + for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), x_tiles)) + with STAGES.span("qknorm_rope"): + cos, sin = layout.cos[:logical], layout.sin[:logical] + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + query = fused_qknorm_rope(query, attn.norm_q.weight, cos.to(query.dtype), sin.to(query.dtype), attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos.to(key.dtype), sin.to(key.dtype), attn.norm_k.eps) + else: + query = attn._apply_rotary_emb(attn.norm_q(query), (cos, sin)) + key = attn._apply_rotary_emb(attn.norm_k(key), (cos, sin)) + gate = None + if attn.to_gate_compress is not None and attn._gate_active(): + with STAGES.span("gate_proj"): + gate, _ = attn.to_gate_compress(x_tiles) + gate = gate.unflatten(-1, (heads, dim)) + capture_root = os.environ.get("FASTVIDEO_H3_CAPTURE_QKV") + if capture_root and attn._layer_idx in (0, 20, 41): + from pathlib import Path + root = Path(capture_root) + root.mkdir(parents=True, exist_ok=True) + capture = root / f"layer-{attn._layer_idx}.pt" + if not capture.exists(): + q_pooled = _pool_tiles(query, meta.variable_block_sizes, meta.tile_elems) + k_pooled = _pool_tiles(key, meta.variable_block_sizes, meta.tile_elems) + scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / dim**0.5 + sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity + mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, sparsity, meta.exempt) + # Two heads keep the artifact small while retaining all real keys, + # query rows, per-tile selections and partial-tile validity. + torch.save({"q": query[:, :, :2].transpose(1, 2).contiguous().cpu(), + "k": key[:, :, :2].transpose(1, 2).contiguous().cpu(), + "v": value[:, :, :2].transpose(1, 2).contiguous().cpu(), + "mask": mask[:, :2].cpu(), "vbs": meta.variable_block_sizes.cpu(), + "untile": meta.untile_combined_index.cpu()}, capture) + del q_pooled, k_pooled, scores, mask + with STAGES.span("attention"): + out = attn.distributed_attention.attn_impl.forward(query, key, value, gate, meta) + with STAGES.span("untile_output"): + return out.index_select(1, layout.untile).flatten(2, 3) + + def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], meta: MiniMaxH3VSAMetadata, use_fused_rope: bool) -> torch.Tensor: """Attention core for ``MiniMaxH3Attention``; returns the pre-``to_out`` ``[B, L, H*D]``.""" diff --git a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py new file mode 100644 index 0000000000..9b1ae79016 --- /dev/null +++ b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py @@ -0,0 +1,94 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Tile-first VSA parity on CUDA, including partial tiles and learned gates.""" +from __future__ import annotations + +from unittest.mock import patch + +import pytest +import torch + +from fastvideo.layers.quantization.fp8_config import FP8Config, FP8QuantizeMethod +from fastvideo.platforms import AttentionBackendEnum +from fastvideo.models.dits.minimax_h3_vsa_fp4 import _shared_input_projections + + +def _install_fp8_buffers(layer): + weight = layer.weight.data.float() + if layer.quant_method.granularity == "channel": + scale = (weight.abs().amax(dim=1, keepdim=True) / 448).clamp_min(1e-6) + else: + scale = (weight.abs().amax().reshape(1) / 448).clamp_min(1e-6) + layer.register_buffer("_fp8_weight", (weight / scale).to(torch.float8_e4m3fn)) + layer.register_buffer("_fp8_weight_scale", scale) + layer.register_parameter("weight", None) + + +@pytest.mark.parametrize("granularity", ["tensor", "channel"]) +def test_shared_fp8_projections_match_independent_quantization(granularity): + if not torch.cuda.is_available() or torch.cuda.get_device_capability() < (8, 9): + pytest.skip("sm89+ CUDA is required for FP8 GEMM") + from fastvideo.layers.linear import ReplicatedLinear + + torch.manual_seed(17) + layers = tuple(ReplicatedLinear(128, 256, bias=True, quant_config=FP8Config(granularity), + prefix=f"block.attn.to_{name}") for name in ("q", "k", "v")) + for layer in layers: + layer.to(device="cuda", dtype=torch.bfloat16) + layer.weight.data.normal_(std=0.1) + layer.bias.data.normal_(std=0.1) + _install_fp8_buffers(layer) + x = torch.randn(1, 272, 128, device="cuda", dtype=torch.bfloat16) + with torch.inference_mode(): + reference = [layer(x)[0] for layer in layers] + with patch.object(FP8QuantizeMethod, "quantize_input", autospec=True, + side_effect=FP8QuantizeMethod.quantize_input) as quant: + actual = _shared_input_projections(layers, x) + assert quant.call_count == 1 + for expected, output in zip(reference, actual, strict=True): + torch.testing.assert_close(output, expected, atol=0, rtol=0) + + +@pytest.mark.parametrize("fp8", [False, True]) +@pytest.mark.parametrize("gate_active", [False, True]) +@pytest.mark.parametrize("fused_rope", [False, True]) +def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distributed_setup, + fp8, gate_active, fused_rope): + if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported(): + pytest.skip("BF16 CUDA is required") + if fp8 and torch.cuda.get_device_capability() < (8, 9): + pytest.skip("sm89+ CUDA is required for FP8 GEMM") + from fastvideo.attention.backends.video_sparse_attn_h3 import MiniMaxH3VSAMetadataBuilder + from fastvideo.forward_context import set_forward_context + from fastvideo.models.dits.minimax_h3 import MiniMaxH3Attention + + monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "VIDEO_SPARSE_ATTN_H3") + monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1") + monkeypatch.setenv("FASTVIDEO_VSA_SM100A", "0") + monkeypatch.setenv("FASTVIDEO_H3_VSA_FP4", "0") + monkeypatch.setenv("FASTVIDEO_H3_VSA_TILE_FIRST", "0") + torch.manual_seed(21) + attn = MiniMaxH3Attention(256, 2, 128, 1e-5, (AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3,), FP8Config("channel") if fp8 else None, + "transformer_blocks.0.attn", fuse_qknorm_rope=fused_rope) + attn.to(device="cuda", dtype=torch.bfloat16) + for parameter in attn.parameters(): + parameter.data.normal_(std=0.1) + if not gate_active: + attn.to_gate_compress.weight.data.zero_() + if fp8: + for layer in (attn.to_q, attn.to_k, attn.to_v, attn.to_out): + _install_fp8_buffers(layer) + meta = MiniMaxH3VSAMetadataBuilder().build(999, (4, 6, 10), (1, 1, 1), 0.8, + (65, 97), torch.device("cuda"), tile_size=64) + length = meta.total_seq_length + x = torch.randn(1, length, 256, device="cuda", dtype=torch.bfloat16) + angles = torch.randn(length, 96, device="cuda") + rope = angles.cos(), angles.sin() + with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=meta): + reference = attn(x, rope, length) + attn._vsa_tile_first = True + actual = attn(x, rope, length) + # Row order can choose a different GEMM reduction; neither attention nor + # the VSA selection/padding semantics are approximated by this route. + error = (actual.float() - reference.float()).norm() / reference.float().norm() + assert error < (0.02 if fp8 else 0.005), float(error) + torch.testing.assert_close(actual, reference, rtol=0.03, atol=0.05) From bd713dddeefc996fbf343c76cbb348d01b9d17b6 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:34:11 -0700 Subject: [PATCH 06/24] [wip]: prototype tile-64 INT8 QK and FP8 PV attention on sm89 --- .../backends/minimax_h3_sparse_int8.py | 112 ++++++++++++++++++ .../attention/test_minimax_h3_sparse_int8.py | 48 ++++++++ scripts/benchmarks/minimax_h3_4090/README.md | 54 +++++++++ .../minimax_h3_4090/bench_sparse_qkv.py | 56 +++++++++ 4 files changed, 270 insertions(+) create mode 100644 fastvideo/attention/backends/minimax_h3_sparse_int8.py create mode 100644 fastvideo/tests/attention/test_minimax_h3_sparse_int8.py create mode 100644 scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py diff --git a/fastvideo/attention/backends/minimax_h3_sparse_int8.py b/fastvideo/attention/backends/minimax_h3_sparse_int8.py new file mode 100644 index 0000000000..f4c2546e80 --- /dev/null +++ b/fastvideo/attention/backends/minimax_h3_sparse_int8.py @@ -0,0 +1,112 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Experimental sm89 tile-64 VSA: INT8 QK and FP8 PV with FP32 accumulation. + +Retains each query tile's original key selection and masks partial key tiles. +Unlike a 128-query adapter, it adds no attention blocks. Q/K use per-token +scales; K centering is a softmax-invariant shift. V uses one scale per head +and channel, so its dequantization can be applied once in the epilogue. +Numerical validation and same-seed clip review are required before enabling. +""" +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _quantize_qk(X, Mean, VBS, Y, Scale, L: tl.constexpr, D: tl.constexpr, CENTER: tl.constexpr, ROWS: tl.constexpr): + hz = tl.program_id(1) + rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS) + cols = tl.arange(0, D) + x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :], rows[:, None] < L, 0).to(tl.float32) + if CENTER: + mean = tl.load(Mean + hz * D + cols) + valid_size = tl.load(VBS + rows // 64, rows < L, 0) + x = tl.where((rows % 64 < valid_size)[:, None], x - mean[None, :], 0.0) + scale = tl.maximum(tl.max(tl.abs(x), 1) / 127.0, 1e-8) + y = tl.floor(x / scale[:, None] + 0.5).to(tl.int8) + tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], y, rows[:, None] < L) + tl.store(Scale + hz * L + rows, scale, rows < L) + + +@triton.jit +def _quantize_v(X, Scale, Y, L: tl.constexpr, D: tl.constexpr, ROWS: tl.constexpr): + hz = tl.program_id(1) + rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS) + cols = tl.arange(0, D) + scale = tl.load(Scale + hz * D + cols) + x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :], rows[:, None] < L, 0).to(tl.float32) + tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], (x / scale[None, :]).to(tl.float8e4nv), rows[:, None] + < L) + + +@triton.autotune( + configs=[triton.Config({}, num_warps=w, num_stages=s) for w, s in ((4, 2), (4, 3), (4, 4), (8, 2), (8, 3))], + key=["L", "D"]) +@triton.jit +def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr): + tile, hz = tl.program_id(0), tl.program_id(1) + nt: tl.constexpr = L // 64 + rows = tile * 64 + tl.arange(0, 64) + cols = tl.arange(0, D) + q = tl.load(Q + (hz * L + rows[:, None]) * D + cols[None, :]) + qs = tl.load(QS + hz * L + rows) + nblocks = tl.load(Count + hz * nt + tile) + m = tl.full((64, ), -float("inf"), tl.float32) + den = tl.zeros((64, ), tl.float32) + acc = tl.zeros((64, D), tl.float32) + for block in range(nblocks): + kv = tl.load(Index + (hz * nt + tile) * nt + block) + key_rows = kv * 64 + tl.arange(0, 64) + k = tl.load(K + (hz * L + key_rows[None, :]) * D + cols[:, None]) + ks = tl.load(KS + hz * L + key_rows) + valid = tl.load(VBS + kv) + if valid > 0: + logits = tl.dot(q, k).to(tl.float32) * qs[:, None] * ks[None, :] * (1.4426950408889634 / D**0.5) + logits = tl.where((tl.arange(0, 64) < valid)[None, :], logits, -float("inf")) + new_m = tl.maximum(m, tl.max(logits, 1)) + p = tl.exp2(logits - new_m[:, None]) + alpha = tl.exp2(m - new_m) + den = den * alpha + tl.sum(p, 1) + acc = acc * alpha[:, None] + v = tl.load(V + (hz * L + key_rows[:, None]) * D + cols[None, :]) + acc += tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32) + m = new_m + vs = tl.load(VS + hz * D + cols) + result = acc / den[:, None] * (vs[None, :] / 448.0) + result = tl.where(den[:, None] > 0, result, 0.0) + tl.store(Out + (hz * L + rows[:, None]) * D + cols[None, :], result.to(Out.dtype.element_ty)) + + +def sparse_int8_fp8_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: torch.Tensor, + vbs: torch.Tensor) -> torch.Tensor: + """Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles.""" + if torch.is_grad_enabled(): + raise ValueError("Sparse INT8/FP8 attention is inference-only") + if not q.is_cuda or torch.cuda.get_device_capability(q.device) != (8, 9): + raise ValueError("Sparse INT8/FP8 attention requires sm89 CUDA") + if q.dtype != torch.bfloat16 or q.shape[-1] != 128 or q.shape != k.shape or q.shape != v.shape: + raise ValueError("Sparse INT8/FP8 attention requires matching BF16 Q/K/V with head dimension 128") + b, h, length, dim = q.shape + if length != vbs.numel() * 64 or mask.shape != (b, h, length // 64, length // 64): + raise ValueError("Sparse INT8/FP8 attention requires a tile-64 mask and validity vector") + from fastvideo_kernel.triton_kernels.index import map_to_index + + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous() + # Tile pads are zero by the VSA contract. Avoid a full FP32 copy for the reduction. + mean = k.sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1) + qi, ki = torch.empty_like(q, dtype=torch.int8), torch.empty_like(k, dtype=torch.int8) + qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32) + ks = torch.empty_like(qs) + grid = (triton.cdiv(length, 16), b * h) + _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, CENTER=False, ROWS=16, num_warps=4) + _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, CENTER=True, ROWS=16, num_warps=4) + vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) + vf = torch.empty_like(v, dtype=torch.float8_e4m3fn) + _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) + index, count = map_to_index(mask.contiguous()) + out = torch.empty_like(q) + _sparse_int8_fp8[(length // 64, b * h)](qi, ki, vf, qs, ks, vs, index, count, vbs, out, length, dim) + return out diff --git a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py new file mode 100644 index 0000000000..5900effc67 --- /dev/null +++ b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +"""sm89 INT8-QK/FP8-PV regression against dense masked BF16 attention.""" +from __future__ import annotations + +import pytest +import torch + + +def _cuda_sm89(): + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (8, 9): + pytest.skip("RTX 4090 / sm89 CUDA is required") + + +@pytest.mark.parametrize("partial", [False, True]) +def test_sparse_int8_preserves_tile_selection_and_valid_keys(partial): + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_int8_fp8_attention + + torch.manual_seed(42) + q, k, v = (torch.randn(1, 2, 256, 128, device="cuda", dtype=torch.bfloat16) for _ in range(3)) + vbs = torch.tensor([64, 7 if partial else 64, 31 if partial else 64, 64], device="cuda", dtype=torch.int32) + valid = torch.arange(256, device="cuda") % 64 < vbs.repeat_interleave(64) + k[..., ~valid, :] = 0 + v[..., ~valid, :] = 0 + # Adjacent query tiles deliberately select different keys. A paired-query + # OR adapter would fail this regression even with perfect quantization. + mask = torch.tensor([[1, 0, 0, 1], [0, 1, 0, 0], [1, 0, 1, 0], [0, 0, 1, 1]], + device="cuda", dtype=torch.bool)[None, None].expand(1, 2, -1, -1).contiguous() + dense_mask = mask.repeat_interleave(64, -2).repeat_interleave(64, -1) & valid[None, None, None, :] + with torch.inference_mode(): + expected = torch.nn.functional.scaled_dot_product_attention(q.float(), k.float(), v.float(), + attn_mask=dense_mask) + output = sparse_int8_fp8_attention(q, k, v, mask, vbs) + assert torch.isfinite(output).all() + relative_error = (output.float() - expected).norm() / expected.norm() + assert relative_error < 0.055, float(relative_error) + torch.testing.assert_close(output.float(), expected, atol=0.05, rtol=0.15) + + +def test_sparse_int8_handles_empty_selection(): + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_int8_fp8_attention + + q = torch.zeros(1, 1, 128, 128, device="cuda", dtype=torch.bfloat16) + with torch.inference_mode(): + out = sparse_int8_fp8_attention(q, q, q, torch.zeros(1, 1, 2, 2, device="cuda", dtype=torch.bool), + torch.tensor([64, 64], device="cuda", dtype=torch.int32)) + assert torch.count_nonzero(out) == 0 diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index f038c1811e..88064ef699 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -101,3 +101,57 @@ python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/summarize.py \ After a baseline works, measure FFN chunk sizes 16,384 and 8,192, then increase resident DiT blocks within the measured GPU budget. Kernel or decoder changes also need same-seed visual and auditory comparison. + +## Tile-first attention and full-resolution profiling + +Commit `c5d9f8132` shares compatible FP8 Q/K/V activation quantization and +releases dead block activations before residual modulation. The opt-in +`FASTVIDEO_H3_VSA_TILE_FIRST=1` scatters the attention input before projection, +then uses the existing BF16 VSA kernel. It retains tile-64 selection, partial +key validity and the learned compression branch. It supports eager, +single-rank inference; grad, compile, and multi-rank requests use the generic +path. Ten CUDA tests passed on the 4090, including mixed partial tiles, +active/zero gates, fused/unfused RoPE and FP8/nonquantized projections. + +The original 1344×768 baseline failed with a GPU OOM in post-attention +modulation before the FFN. No successful 768p timing is established yet. +A profiling run was launched with the following command; retrieve its +results when SSH access is restored. The server recognizes the public key; the +local passphrase-protected private key needs its agent/keychain identity loaded. Profiling/capture timings are +for diagnosis and must not be used as the final speed claim. + +```bash +FASTVIDEO_SOURCE_COMMIT=c5d9f8132 \ +FASTVIDEO_H3_PARK_MODULES=vae,audio_vae \ +FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 \ +FASTVIDEO_H3_VSA_TILE_FIRST=1 \ +FASTVIDEO_H3_CAPTURE_QKV=/workspace/qkv-768p MAX_JOBS=4 \ +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + tile-first-768p-profile /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --lazy --no-vae-compile \ + --height 768 --width 1344 --frames 243 --timed 2 --profile +``` + +`FASTVIDEO_H3_CAPTURE_QKV` saves the first inputs from layers 0, 20 and 41, +two heads each, with full real sequences, masks, valid tile sizes and packed +row indices. Disable both capture and profiling for final clip timings. + +The separate `minimax_h3_sparse_int8.py` prototype uses INT8 QK and FP8 PV, +FP32 accumulators and the original 64-token mask. It has no automatic pipeline +route. Offline compilation with Triton 3.8.0 for sm89 passed all four entry points +and emitted native INT8 and FP8 MMA instructions (20,480 bytes of shared +memory for the attention kernel). This is compilation evidence only. It must +pass CUDA tests and real-QKV/clip checks before integration. +Run its microbenchmark on an idle GPU: + +```bash +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py \ + /workspace/qkv-768p --output /workspace/sparse-qkv-results.json +``` + +SpargeAttn at `ae5b629ebb41e41f86b3ea2ab5a3283f13ac151a` built on the pod +with CUDA 12.8, `TORCH_CUDA_ARCH_LIST=8.9`, and `MAX_JOBS=4`. The upstream +`-Xcompiler -include,cassert` workaround was removed from `setup.py` to +avoid GCC 13 duplicate standard-library definitions. It is not selected by +the pipeline: its public 128-query/64-key adapter also needs correct masking +of partial H3 tiles before a meaningful parity comparison. diff --git a/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py new file mode 100644 index 0000000000..b1fe97d7db --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py @@ -0,0 +1,56 @@ +"""Validate and time candidate kernels on captured real H3 Q/K/V. + +Run on an idle sm89 GPU, separately from clip benchmarks. Captures come from +FASTVIDEO_H3_CAPTURE_QKV on the tile-first path; they retain two full heads. +""" +import argparse +import json +import pathlib + +import torch +import triton + +from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_int8_fp8_attention +from fastvideo_kernel.block_sparse_attn import block_sparse_attn + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("captures", type=pathlib.Path) + parser.add_argument("--output", type=pathlib.Path, required=True) + args = parser.parse_args() + records = [] + with torch.inference_mode(): + for capture in sorted(args.captures.glob("layer-*.pt")): + state = torch.load(capture, map_location="cuda", weights_only=True) + q, k, v, mask, vbs = (state[key] for key in ("q", "k", "v", "mask", "vbs")) + def baseline(): + return block_sparse_attn(q, k, v, mask, vbs)[0] + def candidate(): + return sparse_int8_fp8_attention(q, k, v, mask, vbs) + expected = baseline() + output = candidate() + valid_rows = state["untile"] + ref = expected.index_select(2, valid_rows).float() + actual = output.index_select(2, valid_rows).float() + delta = actual - ref + reference_ms = triton.testing.do_bench(baseline) + candidate_ms = triton.testing.do_bench(candidate) + record = {"capture": capture.name, "shape": list(q.shape), + "mask_density": float(mask.float().mean()), + "finite": bool(torch.isfinite(actual).all()), + "relative_l2": float(delta.norm() / ref.norm()), + "max_abs": float(delta.abs().max()), + "cosine": float(torch.nn.functional.cosine_similarity(actual.flatten(), ref.flatten(), dim=0)), + "bf16_ms": reference_ms, "int8_fp8_ms": candidate_ms, + "speedup": reference_ms / candidate_ms} + print(json.dumps(record), flush=True) + records.append(record) + if not records: + raise RuntimeError("No real Q/K/V captures found") + args.output.write_text(json.dumps({"gpu": torch.cuda.get_device_name(), "torch": torch.__version__, + "triton": triton.__version__, "records": records}, indent=2) + "\n") + + +if __name__ == "__main__": + main() From 5bf9804de85953207af6356d157fc68b74fe2635 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:55:28 -0700 Subject: [PATCH 07/24] [perf]: retain offloaded H3 VAEs on host until their consuming stages --- .../basic/minimax_h3/minimax_h3_pipeline.py | 9 +++++++-- .../stages/test_minimax_h3_sequential_start.py | 15 +++++++++++++++ 2 files changed, 22 insertions(+), 2 deletions(-) diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 005f981a9c..1edee22c4a 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -497,12 +497,17 @@ def _park_denoise_modules(self) -> None: torch.cuda.empty_cache() logger.info("Parked MiniMax-H3 denoise modules on CPU for text encode") - def _restore_denoise_modules(self) -> None: + def _restore_denoise_modules(self, fastvideo_args: FastVideoArgs) -> None: from fastvideo.pipelines import composed_pipeline_base device = composed_pipeline_base.get_local_torch_device() restored = False for name in _DENOISE_MODULE_NAMES: + # Encode/decode stages move each VAE to the device when consumed. + # Keeping offloaded VAEs on the host leaves room for DiT activations + # and resident blocks throughout the denoising loop. + if name in {"vae", "audio_vae"} and fastvideo_args.vae_cpu_offload: + continue module = self.get_module(name) if module is None: continue @@ -517,7 +522,7 @@ def _run_condition_then_denoise(self, batch: ForwardBatch, fastvideo_args: FastV self._release_text_encoder() self._load_denoise_modules(fastvideo_args) if not self._unified_memory_host(): - self._restore_denoise_modules() + self._restore_denoise_modules(fastvideo_args) if not self._denoise_stages_ready: self._add_denoise_stages(ref2va=self._ref2va) for name in ( diff --git a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index c513ff1239..3763e9159a 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -5,6 +5,7 @@ from contextlib import nullcontext from types import SimpleNamespace +import pytest import torch import fastvideo.pipelines.composed_pipeline_base as composed_pipeline_base @@ -51,6 +52,8 @@ def to(device): def _patch_pipeline_construction(monkeypatch, events: list, *, unified_memory: bool = False) -> None: + # These contract tests use lightweight objects, not tensor-bearing modules. + monkeypatch.setenv("FASTVIDEO_H3_PINNED_SWAP", "0") monkeypatch.setattr( composed_pipeline_base, "maybe_init_distributed_environment_and_model_parallel", @@ -523,3 +526,15 @@ def fake_load(self, fastvideo_args, loaded_modules=None): assert first is not None and second is not None assert len(loads) == 1 assert pipeline.get_module("text_encoder") is not None + + +@pytest.mark.parametrize("vae_offload", [True, False]) +def test_sequential_restore_keeps_offloaded_vaes_on_host_until_consumed(monkeypatch, vae_offload): + """Do not occupy denoise VRAM with decoders that stages load on demand.""" + pipeline = MiniMaxH3Pipeline.__new__(MiniMaxH3Pipeline) + pipeline.modules = {name: _stub_module(name) for name in _DENOISE_MODULE_NAMES} + moved = [] + monkeypatch.setattr(composed_pipeline_base, "get_local_torch_device", lambda: torch.device("cuda", 0)) + monkeypatch.setattr(pipeline, "_move_module", lambda module, device: moved.append(module.name) or True) + pipeline._restore_denoise_modules(SimpleNamespace(vae_cpu_offload=vae_offload)) + assert moved == (["transformer"] if vae_offload else list(_DENOISE_MODULE_NAMES)) From 59e6946efe9fb0471c5417231e12b54446b37789 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:55:28 -0700 Subject: [PATCH 08/24] [perf]: add opt-in sm89 tile-64 BF16 and INT8 QK attention --- .../backends/minimax_h3_sparse_int8.py | 77 ++++++++++++++----- .../backends/video_sparse_attn_h3.py | 30 ++++++-- .../attention/test_minimax_h3_sparse_int8.py | 8 +- .../test_minimax_h3_tile_first.py | 5 +- .../minimax_h3_4090/bench_sparse_qkv.py | 41 +++++----- 5 files changed, 108 insertions(+), 53 deletions(-) diff --git a/fastvideo/attention/backends/minimax_h3_sparse_int8.py b/fastvideo/attention/backends/minimax_h3_sparse_int8.py index f4c2546e80..c848f951d9 100644 --- a/fastvideo/attention/backends/minimax_h3_sparse_int8.py +++ b/fastvideo/attention/backends/minimax_h3_sparse_int8.py @@ -1,10 +1,11 @@ # SPDX-License-Identifier: Apache-2.0 -"""Experimental sm89 tile-64 VSA: INT8 QK and FP8 PV with FP32 accumulation. +"""Experimental sm89 tile-64 VSA with BF16 or INT8 QK and FP32 accumulation. Retains each query tile's original key selection and masks partial key tiles. Unlike a 128-query adapter, it adds no attention blocks. Q/K use per-token scales; K centering is a softmax-invariant shift. V uses one scale per head and channel, so its dequantization can be applied once in the epilogue. +BF16 PV is the default: FP8 PV had excessive error on real H3 inputs. Numerical validation and same-seed clip review are required before enabling. """ from __future__ import annotations @@ -45,13 +46,15 @@ def _quantize_v(X, Scale, Y, L: tl.constexpr, D: tl.constexpr, ROWS: tl.constexp configs=[triton.Config({}, num_warps=w, num_stages=s) for w, s in ((4, 2), (4, 3), (4, 4), (8, 2), (8, 3))], key=["L", "D"]) @triton.jit -def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr): +def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr, + INT8_QK: tl.constexpr, FP8_PV: tl.constexpr): tile, hz = tl.program_id(0), tl.program_id(1) nt: tl.constexpr = L // 64 rows = tile * 64 + tl.arange(0, 64) cols = tl.arange(0, D) q = tl.load(Q + (hz * L + rows[:, None]) * D + cols[None, :]) - qs = tl.load(QS + hz * L + rows) + if INT8_QK: + qs = tl.load(QS + hz * L + rows) nblocks = tl.load(Count + hz * nt + tile) m = tl.full((64, ), -float("inf"), tl.float32) den = tl.zeros((64, ), tl.float32) @@ -60,10 +63,14 @@ def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexp kv = tl.load(Index + (hz * nt + tile) * nt + block) key_rows = kv * 64 + tl.arange(0, 64) k = tl.load(K + (hz * L + key_rows[None, :]) * D + cols[:, None]) - ks = tl.load(KS + hz * L + key_rows) + if INT8_QK: + ks = tl.load(KS + hz * L + key_rows) valid = tl.load(VBS + kv) if valid > 0: - logits = tl.dot(q, k).to(tl.float32) * qs[:, None] * ks[None, :] * (1.4426950408889634 / D**0.5) + logits = tl.dot(q, k).to(tl.float32) + if INT8_QK: + logits = logits * qs[:, None] * ks[None, :] + logits = logits * (1.4426950408889634 / D**0.5) logits = tl.where((tl.arange(0, 64) < valid)[None, :], logits, -float("inf")) new_m = tl.maximum(m, tl.max(logits, 1)) p = tl.exp2(logits - new_m[:, None]) @@ -71,16 +78,27 @@ def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexp den = den * alpha + tl.sum(p, 1) acc = acc * alpha[:, None] v = tl.load(V + (hz * L + key_rows[:, None]) * D + cols[None, :]) - acc += tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32) + if FP8_PV: + acc += tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32) + else: + acc += tl.dot(p.to(tl.bfloat16), v, out_dtype=tl.float32) m = new_m - vs = tl.load(VS + hz * D + cols) - result = acc / den[:, None] * (vs[None, :] / 448.0) + result = acc / den[:, None] + if FP8_PV: + vs = tl.load(VS + hz * D + cols) + result = result * (vs[None, :] / 448.0) result = tl.where(den[:, None] > 0, result, 0.0) tl.store(Out + (hz * L + rows[:, None]) * D + cols[None, :], result.to(Out.dtype.element_ty)) -def sparse_int8_fp8_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: torch.Tensor, - vbs: torch.Tensor) -> torch.Tensor: +def sparse_sm89_attention(q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + mask: torch.Tensor, + vbs: torch.Tensor, + *, + int8_qk: bool = True, + fp8_pv: bool = False) -> torch.Tensor: """Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles.""" if torch.is_grad_enabled(): raise ValueError("Sparse INT8/FP8 attention is inference-only") @@ -95,18 +113,35 @@ def sparse_int8_fp8_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, q, k, v = q.contiguous(), k.contiguous(), v.contiguous() vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous() - # Tile pads are zero by the VSA contract. Avoid a full FP32 copy for the reduction. - mean = k.sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1) - qi, ki = torch.empty_like(q, dtype=torch.int8), torch.empty_like(k, dtype=torch.int8) - qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32) - ks = torch.empty_like(qs) grid = (triton.cdiv(length, 16), b * h) - _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, CENTER=False, ROWS=16, num_warps=4) - _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, CENTER=True, ROWS=16, num_warps=4) - vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) - vf = torch.empty_like(v, dtype=torch.float8_e4m3fn) - _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) + qi, ki, vf = q, k, v + qs, ks, vs = q, k, v # unused pointers in BF16 ablations + if int8_qk: + # Tile pads are zero by contract; avoid a full FP32 copy for the reduction. + mean = k.sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1) + qi, ki = torch.empty_like(q, dtype=torch.int8), torch.empty_like(k, dtype=torch.int8) + qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32) + ks = torch.empty_like(qs) + _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, CENTER=False, ROWS=16, num_warps=4) + _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, CENTER=True, ROWS=16, num_warps=4) + if fp8_pv: + vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) + vf = torch.empty_like(v, dtype=torch.float8_e4m3fn) + _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) index, count = map_to_index(mask.contiguous()) out = torch.empty_like(q) - _sparse_int8_fp8[(length // 64, b * h)](qi, ki, vf, qs, ks, vs, index, count, vbs, out, length, dim) + _sparse_int8_fp8[(length // 64, b * h)](qi, + ki, + vf, + qs, + ks, + vs, + index, + count, + vbs, + out, + length, + dim, + INT8_QK=int8_qk, + FP8_PV=fp8_pv) return out diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index 65401519bd..86629ade78 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -547,6 +547,9 @@ def __init__( self.prefix = prefix self.layer_idx = layer_idx_from_prefix(prefix, default=-1) self.head_size = head_size + self._sm89_kernel = os.environ.get("FASTVIDEO_H3_VSA_SM89_KERNEL", "original") + if self._sm89_kernel not in {"original", "bf16", "int8"}: + raise ValueError("FASTVIDEO_H3_VSA_SM89_KERNEL must be original, bf16, or int8") # Generic torch.compile must not specialize the shared VSA forward on # the Python ``layer_idx`` value of each of H3's 50 blocks. This # tensor is prepared after weights load and drives only the compiled @@ -880,13 +883,26 @@ def forward( # type: ignore[override] q_bhsd = q_bhsd[:, :, :logical_seq_len].contiguous() k_bhsd = k_bhsd[:, :, :logical_seq_len].contiguous() v_bhsd = v_bhsd[:, :, :logical_seq_len].contiguous() - out_bhsd, _ = block_sparse_attn_64_bhsd( - q_bhsd, - k_bhsd, - v_bhsd, - mask, - attn_metadata.variable_block_sizes, - ) + if (self._sm89_kernel != "original" and not torch.is_grad_enabled() and not compiling + and q_bhsd.dtype == torch.bfloat16 and q_bhsd.shape[-1] == 128 + and torch.cuda.get_device_capability(q_bhsd.device) == (8, 9)): + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + logger.info_once(f"MiniMax-H3 VSA tile-64 forward: sm89 {self._sm89_kernel} QK / BF16 PV") + out_bhsd = sparse_sm89_attention(q_bhsd, + k_bhsd, + v_bhsd, + mask, + attn_metadata.variable_block_sizes, + int8_qk=self._sm89_kernel == "int8", + fp8_pv=False) + else: + out_bhsd, _ = block_sparse_attn_64_bhsd( + q_bhsd, + k_bhsd, + v_bhsd, + mask, + attn_metadata.variable_block_sizes, + ) if has_sm100a_pair and use_sm100a: out_bhsd = out_bhsd[:, :, :logical_seq_len] out = out_bhsd.transpose(1, 2).contiguous() diff --git a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py index 5900effc67..1f22ad36ef 100644 --- a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py +++ b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py @@ -14,7 +14,7 @@ def _cuda_sm89(): @pytest.mark.parametrize("partial", [False, True]) def test_sparse_int8_preserves_tile_selection_and_valid_keys(partial): _cuda_sm89() - from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_int8_fp8_attention + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention torch.manual_seed(42) q, k, v = (torch.randn(1, 2, 256, 128, device="cuda", dtype=torch.bfloat16) for _ in range(3)) @@ -30,7 +30,7 @@ def test_sparse_int8_preserves_tile_selection_and_valid_keys(partial): with torch.inference_mode(): expected = torch.nn.functional.scaled_dot_product_attention(q.float(), k.float(), v.float(), attn_mask=dense_mask) - output = sparse_int8_fp8_attention(q, k, v, mask, vbs) + output = sparse_sm89_attention(q, k, v, mask, vbs) assert torch.isfinite(output).all() relative_error = (output.float() - expected).norm() / expected.norm() assert relative_error < 0.055, float(relative_error) @@ -39,10 +39,10 @@ def test_sparse_int8_preserves_tile_selection_and_valid_keys(partial): def test_sparse_int8_handles_empty_selection(): _cuda_sm89() - from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_int8_fp8_attention + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention q = torch.zeros(1, 1, 128, 128, device="cuda", dtype=torch.bfloat16) with torch.inference_mode(): - out = sparse_int8_fp8_attention(q, q, q, torch.zeros(1, 1, 2, 2, device="cuda", dtype=torch.bool), + out = sparse_sm89_attention(q, q, q, torch.zeros(1, 1, 2, 2, device="cuda", dtype=torch.bool), torch.tensor([64, 64], device="cuda", dtype=torch.int32)) assert torch.count_nonzero(out) == 0 diff --git a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py index 9b1ae79016..9ebbf8181d 100644 --- a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py +++ b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py @@ -48,11 +48,12 @@ def test_shared_fp8_projections_match_independent_quantization(granularity): torch.testing.assert_close(output, expected, atol=0, rtol=0) +@pytest.mark.parametrize("kernel", ["original", "bf16", "int8"]) @pytest.mark.parametrize("fp8", [False, True]) @pytest.mark.parametrize("gate_active", [False, True]) @pytest.mark.parametrize("fused_rope", [False, True]) def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distributed_setup, - fp8, gate_active, fused_rope): + fp8, gate_active, fused_rope, kernel): if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported(): pytest.skip("BF16 CUDA is required") if fp8 and torch.cuda.get_device_capability() < (8, 9): @@ -66,6 +67,7 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distribu monkeypatch.setenv("FASTVIDEO_VSA_SM100A", "0") monkeypatch.setenv("FASTVIDEO_H3_VSA_FP4", "0") monkeypatch.setenv("FASTVIDEO_H3_VSA_TILE_FIRST", "0") + monkeypatch.setenv("FASTVIDEO_H3_VSA_SM89_KERNEL", "original") torch.manual_seed(21) attn = MiniMaxH3Attention(256, 2, 128, 1e-5, (AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3,), FP8Config("channel") if fp8 else None, "transformer_blocks.0.attn", fuse_qknorm_rope=fused_rope) @@ -86,6 +88,7 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distribu with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=meta): reference = attn(x, rope, length) attn._vsa_tile_first = True + attn.distributed_attention.attn_impl._sm89_kernel = kernel actual = attn(x, rope, length) # Row order can choose a different GEMM reduction; neither attention nor # the VSA selection/padding semantics are approximated by this route. diff --git a/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py index b1fe97d7db..6ee2c2c1f9 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py @@ -10,7 +10,7 @@ import torch import triton -from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_int8_fp8_attention +from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention from fastvideo_kernel.block_sparse_attn import block_sparse_attn @@ -26,26 +26,27 @@ def main(): q, k, v, mask, vbs = (state[key] for key in ("q", "k", "v", "mask", "vbs")) def baseline(): return block_sparse_attn(q, k, v, mask, vbs)[0] - def candidate(): - return sparse_int8_fp8_attention(q, k, v, mask, vbs) expected = baseline() - output = candidate() - valid_rows = state["untile"] - ref = expected.index_select(2, valid_rows).float() - actual = output.index_select(2, valid_rows).float() - delta = actual - ref - reference_ms = triton.testing.do_bench(baseline) - candidate_ms = triton.testing.do_bench(candidate) - record = {"capture": capture.name, "shape": list(q.shape), - "mask_density": float(mask.float().mean()), - "finite": bool(torch.isfinite(actual).all()), - "relative_l2": float(delta.norm() / ref.norm()), - "max_abs": float(delta.abs().max()), - "cosine": float(torch.nn.functional.cosine_similarity(actual.flatten(), ref.flatten(), dim=0)), - "bf16_ms": reference_ms, "int8_fp8_ms": candidate_ms, - "speedup": reference_ms / candidate_ms} - print(json.dumps(record), flush=True) - records.append(record) + for int8_qk, fp8_pv in ((True, True), (True, False), (False, True), (False, False)): + def candidate(): + return sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=int8_qk, fp8_pv=fp8_pv) + output = candidate() + valid_rows = state["untile"] + ref = expected.index_select(2, valid_rows).float() + actual = output.index_select(2, valid_rows).float() + delta = actual - ref + reference_ms = triton.testing.do_bench(baseline) + candidate_ms = triton.testing.do_bench(candidate) + record = {"capture": capture.name, "shape": list(q.shape), "int8_qk": int8_qk, "fp8_pv": fp8_pv, + "mask_density": float(mask.float().mean()), + "finite": bool(torch.isfinite(actual).all()), + "relative_l2": float(delta.norm() / ref.norm()), + "max_abs": float(delta.abs().max()), + "cosine": float(torch.nn.functional.cosine_similarity(actual.flatten(), ref.flatten(), dim=0)), + "bf16_ms": reference_ms, "int8_fp8_ms": candidate_ms, + "speedup": reference_ms / candidate_ms} + print(json.dumps(record), flush=True) + records.append(record) if not records: raise RuntimeError("No real Q/K/V captures found") args.output.write_text(json.dumps({"gpu": torch.cuda.get_device_name(), "torch": torch.__version__, From b3ab6e1e743d8146d517b56a3af2e652f4139345 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:58:59 -0700 Subject: [PATCH 09/24] [docs]: record cached 4090 timings and sm89 precision validation --- scripts/benchmarks/minimax_h3_4090/README.md | 61 ++++++++++++++++++-- 1 file changed, 56 insertions(+), 5 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 88064ef699..27148ad561 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -114,10 +114,13 @@ path. Ten CUDA tests passed on the 4090, including mixed partial tiles, active/zero gates, fused/unfused RoPE and FP8/nonquantized projections. The original 1344×768 baseline failed with a GPU OOM in post-attention -modulation before the FFN. No successful 768p timing is established yet. -A profiling run was launched with the following command; retrieve its -results when SSH access is restored. The server recognizes the public key; the -local passphrase-protected private key needs its agent/keychain identity loaded. Profiling/capture timings are +modulation before the FFN. The tile-first/FFN-16384 profiling run subsequently completed all three +768p clips without OOM: diagnostic median 349.76 s, denoise stage 244.07 s, +video decode stage 65.94 s, peak GPU allocated 17.55 GiB. A run without +profiling is still required for a release speed claim. +The profiling run used the following command. For SSH on macOS, +`-o UseKeychain=yes` retrieves the stored passphrase when the agent has no +loaded identities. Profiling/capture timings are for diagnosis and must not be used as the final speed claim. ```bash @@ -138,7 +141,10 @@ row indices. Disable both capture and profiling for final clip timings. The separate `minimax_h3_sparse_int8.py` prototype uses INT8 QK and FP8 PV, FP32 accumulators and the original 64-token mask. It has no automatic pipeline -route. Offline compilation with Triton 3.8.0 for sm89 passed all four entry points +route. Initial real-QKV tests found approximately 1.9× fine-kernel speedup +but 4.3–13.4% relative L2 error and a strict partial-tile elementwise test +failure. Do not select it for shipping. The microbenchmark now includes +BF16 QK/PV ablations to isolate that error. Offline compilation with Triton 3.8.0 for sm89 passed all four entry points and emitted native INT8 and FP8 MMA instructions (20,480 bytes of shared memory for the attention kernel). This is compilation evidence only. It must pass CUDA tests and real-QKV/clip checks before integration. @@ -155,3 +161,48 @@ with CUDA 12.8, `TORCH_CUDA_ARCH_LIST=8.9`, and `MAX_JOBS=4`. The upstream avoid GCC 13 duplicate standard-library definitions. It is not selected by the pipeline: its public 128-query/64-key adapter also needs correct masking of partial H3 tiles before a meaningful parity comparison. + +## Cached-component 480p result + +At `4d9846573`, retain components between requests (omit `--lazy`), and set +`FASTVIDEO_H3_VSA_TILE_FIRST=1` and `FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384`. +Keep all other baseline settings, including the eager light VAE and original +BF16 attention kernel. After one warmup the two timed requests took 115.87 s +and 115.03 s, median **115.45 s** (29.3% less wall time than the lazy baseline). +Conditioning/denoise/video-decode stage medians were 10.85/72.62/24.98 s. +Peak GPU allocation was 19.02 GiB, host anon 43.96 GiB, total cgroup 92.72 GiB +including file cache. This recipe requires more host RAM than the 32 GB target; +its minimum RAM has not been tested under a smaller host limit. + +Decoded raw video and PCM audio SHA256 hashes match the baseline exactly for +both ceramics and harbor at seed 20260929. This establishes output identity +for these two prompts; it does not establish the checkpoint's BF16-reference +quality on other prompts. Raw results, clips and hash evidence are saved in +`output/fasth3-4090-20261003/` beside the workspace. + +## sm89 kernel precision choices + +`FASTVIDEO_H3_VSA_SM89_KERNEL=bf16` opts into the new entirely BF16 tile-64 +fine kernel. `int8` uses per-token INT8 QK with BF16 PV. `original` is the +unchanged default. Resolution happens when the backend is constructed; +unsupported devices, grad and compile requests retain the original route. +Both preserve the original tile selection, partial key masks and gated +compression. The rejected FP8-PV experiment is only exposed in the diagnostic +microbenchmark, never the pipeline route. + +Two-head real-QKV captures at 1344×768, layers 0/20/41, measured: + +| QK / PV | Fine-kernel speedup including input quantization | Relative L2 vs original BF16 | +| --- | --- | --- | +| BF16 / BF16 | 1.23× | 0.005–0.008% | +| INT8 / BF16 | 1.59× | 0.58–0.62% | +| INT8 / FP8 (rejected) | 1.91× | 4.3–13.4% | + +These are fine-kernel microbenchmarks, not end-to-end clip speedups. Same-seed +clip checks are required for the INT8 route. All 44 targeted CUDA/CPU checks +passed for native/tile-first routing, partial tiles, learned compression, +shared FP8 projections and sequential component restoration. + +At `9a8465ac4`, CPU-offloaded VAEs also remain on the host during denoising; +the encode/decode stages move them on demand. This frees room for resident +DiT blocks without changing any model arithmetic. From 994c220fdbff76ce0b3d9f39742beffbffb2b98b Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:13:54 -0700 Subject: [PATCH 10/24] [wip]: validate tilewise FP8 values and dynamic probability scales --- .../backends/minimax_h3_sparse_int8.py | 46 +++++++++++++++---- .../attention/test_minimax_h3_sparse_int8.py | 24 ++++++++++ .../minimax_h3_4090/bench_sparse_qkv.py | 11 +++-- 3 files changed, 70 insertions(+), 11 deletions(-) diff --git a/fastvideo/attention/backends/minimax_h3_sparse_int8.py b/fastvideo/attention/backends/minimax_h3_sparse_int8.py index c848f951d9..4c8c5ddad1 100644 --- a/fastvideo/attention/backends/minimax_h3_sparse_int8.py +++ b/fastvideo/attention/backends/minimax_h3_sparse_int8.py @@ -42,12 +42,23 @@ def _quantize_v(X, Scale, Y, L: tl.constexpr, D: tl.constexpr, ROWS: tl.constexp < L) +@triton.jit +def _quantize_v_tiles(X, Y, Scale, L: tl.constexpr, D: tl.constexpr): + tile, hz = tl.program_id(0), tl.program_id(1) + rows = tile * 64 + tl.arange(0, 64) + cols = tl.arange(0, D) + x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :]).to(tl.float32) + scale = tl.maximum(tl.max(tl.abs(x), 0) / 448.0, 1e-8) + tl.store(Scale + (hz * (L // 64) + tile) * D + cols, scale) + tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], (x / scale[None, :]).to(tl.float8e4nv)) + + @triton.autotune( configs=[triton.Config({}, num_warps=w, num_stages=s) for w, s in ((4, 2), (4, 3), (4, 4), (8, 2), (8, 3))], key=["L", "D"]) @triton.jit def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr, - INT8_QK: tl.constexpr, FP8_PV: tl.constexpr): + INT8_QK: tl.constexpr, FP8_PV: tl.constexpr, V_TILE: tl.constexpr, P_DYNAMIC: tl.constexpr): tile, hz = tl.program_id(0), tl.program_id(1) nt: tl.constexpr = L // 64 rows = tile * 64 + tl.arange(0, 64) @@ -72,19 +83,30 @@ def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexp logits = logits * qs[:, None] * ks[None, :] logits = logits * (1.4426950408889634 / D**0.5) logits = tl.where((tl.arange(0, 64) < valid)[None, :], logits, -float("inf")) - new_m = tl.maximum(m, tl.max(logits, 1)) + block_max = tl.max(logits, 1) + new_m = tl.maximum(m, block_max) p = tl.exp2(logits - new_m[:, None]) alpha = tl.exp2(m - new_m) den = den * alpha + tl.sum(p, 1) acc = acc * alpha[:, None] v = tl.load(V + (hz * L + key_rows[:, None]) * D + cols[None, :]) if FP8_PV: - acc += tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32) + if P_DYNAMIC: + pscale = tl.maximum(tl.exp2(block_max - new_m) / 448.0, 1e-30) + pv = tl.dot((p / pscale[:, None]).to(tl.float8e4nv), v, out_dtype=tl.float32) + pv = pv * (pscale[:, None] * 448.0) + else: + pv = tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32) + if V_TILE: + scale = tl.load(VS + (hz * nt + kv) * D + cols) + acc += pv * (scale[None, :] / 448.0) + else: + acc += pv else: acc += tl.dot(p.to(tl.bfloat16), v, out_dtype=tl.float32) m = new_m result = acc / den[:, None] - if FP8_PV: + if FP8_PV and not V_TILE: vs = tl.load(VS + hz * D + cols) result = result * (vs[None, :] / 448.0) result = tl.where(den[:, None] > 0, result, 0.0) @@ -98,7 +120,9 @@ def sparse_sm89_attention(q: torch.Tensor, vbs: torch.Tensor, *, int8_qk: bool = True, - fp8_pv: bool = False) -> torch.Tensor: + fp8_pv: bool = False, + fp8_v_tiles: bool = False, + fp8_dynamic_p: bool = False) -> torch.Tensor: """Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles.""" if torch.is_grad_enabled(): raise ValueError("Sparse INT8/FP8 attention is inference-only") @@ -125,9 +149,13 @@ def sparse_sm89_attention(q: torch.Tensor, _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, CENTER=False, ROWS=16, num_warps=4) _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, CENTER=True, ROWS=16, num_warps=4) if fp8_pv: - vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) vf = torch.empty_like(v, dtype=torch.float8_e4m3fn) - _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) + if fp8_v_tiles: + vs = torch.empty((b, h, length // 64, dim), device=q.device, dtype=torch.float32) + _quantize_v_tiles[(length // 64, b * h)](v, vf, vs, length, dim, num_warps=8) + else: + vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) + _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) index, count = map_to_index(mask.contiguous()) out = torch.empty_like(q) _sparse_int8_fp8[(length // 64, b * h)](qi, @@ -143,5 +171,7 @@ def sparse_sm89_attention(q: torch.Tensor, length, dim, INT8_QK=int8_qk, - FP8_PV=fp8_pv) + FP8_PV=fp8_pv, + V_TILE=fp8_v_tiles, + P_DYNAMIC=fp8_dynamic_p) return out diff --git a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py index 1f22ad36ef..122aa08f11 100644 --- a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py +++ b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py @@ -46,3 +46,27 @@ def test_sparse_int8_handles_empty_selection(): out = sparse_sm89_attention(q, q, q, torch.zeros(1, 1, 2, 2, device="cuda", dtype=torch.bool), torch.tensor([64, 64], device="cuda", dtype=torch.int32)) assert torch.count_nonzero(out) == 0 + + +def test_fp8_dynamic_probability_scale_preserves_small_blocks(): + """A large earlier max must not erase a later block's small P but large V.""" + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + + q = torch.zeros(1, 1, 128, 128, device="cuda", dtype=torch.bfloat16) + k, v = torch.zeros_like(q), torch.zeros_like(q) + q[..., 0] = 16 + k[..., :64, 0] = 14 + k[..., 64:, 0] = 2 + v[..., 64:, :] = 1e7 + mask = torch.ones(1, 1, 2, 2, device="cuda", dtype=torch.bool) + vbs = torch.tensor([64, 64], device="cuda", dtype=torch.int32) + with torch.inference_mode(): + reference = torch.nn.functional.scaled_dot_product_attention(q.float(), k.float(), v.float()) + fixed = sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=False, fp8_pv=True, + fp8_v_tiles=True, fp8_dynamic_p=False) + dynamic = sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=False, fp8_pv=True, + fp8_v_tiles=True, fp8_dynamic_p=True) + assert reference.abs().min() > 0.1 + assert torch.count_nonzero(fixed) == 0 + torch.testing.assert_close(dynamic.float(), reference, rtol=0.02, atol=0.02) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py index 6ee2c2c1f9..b03b6f70c5 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_sparse_qkv.py @@ -27,9 +27,14 @@ def main(): def baseline(): return block_sparse_attn(q, k, v, mask, vbs)[0] expected = baseline() - for int8_qk, fp8_pv in ((True, True), (True, False), (False, True), (False, False)): + for int8_qk, fp8_pv, v_tiles, dynamic_p in ((True, True, False, False), + (True, True, True, False), + (True, True, True, True), + (True, False, False, False), + (False, True, True, True), + (False, False, False, False)): def candidate(): - return sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=int8_qk, fp8_pv=fp8_pv) + return sparse_sm89_attention(q, k, v, mask, vbs, int8_qk=int8_qk, fp8_pv=fp8_pv, fp8_v_tiles=v_tiles, fp8_dynamic_p=dynamic_p) output = candidate() valid_rows = state["untile"] ref = expected.index_select(2, valid_rows).float() @@ -37,7 +42,7 @@ def candidate(): delta = actual - ref reference_ms = triton.testing.do_bench(baseline) candidate_ms = triton.testing.do_bench(candidate) - record = {"capture": capture.name, "shape": list(q.shape), "int8_qk": int8_qk, "fp8_pv": fp8_pv, + record = {"capture": capture.name, "shape": list(q.shape), "int8_qk": int8_qk, "fp8_pv": fp8_pv, "fp8_v_tiles": v_tiles, "fp8_dynamic_p": dynamic_p, "mask_density": float(mask.float().mean()), "finite": bool(torch.isfinite(actual).all()), "relative_l2": float(delta.norm() / ref.norm()), From 6b18535ffbc2128d36dd1432936d2fb3744d4b67 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:18:01 -0700 Subject: [PATCH 11/24] [docs]: record resident-block timings and encoder constraints --- scripts/benchmarks/minimax_h3_4090/README.md | 50 ++++++++++++++++++++ 1 file changed, 50 insertions(+) diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 27148ad561..3563d6d4d0 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -206,3 +206,53 @@ shared FP8 projections and sequential component restoration. At `9a8465ac4`, CPU-offloaded VAEs also remain on the host during denoising; the encode/decode stages move them on demand. This frees room for resident DiT blocks without changing any model arithmetic. + +## Six resident blocks and encoder priorities + +At `6d3c4cda5`, the opt-in BF16 fine kernel with six resident DiT blocks, +cached components and the settings below measured **111.27 s** median for +832×480, 243 frames. Timed requests were 111.61/110.92 s after one warmup. +Conditioning, denoise and video decode medians were 11.20/67.56/25.49 s. +Peak GPU allocation was 21.63 GiB, host anonymous memory 42.10 GiB and +total cgroup usage 90.94 GiB, including file cache. + +```bash +FASTVIDEO_SOURCE_COMMIT=6d3c4cda5 \ +FASTVIDEO_H3_PARK_MODULES=vae,audio_vae \ +FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 \ +FASTVIDEO_H3_VSA_TILE_FIRST=1 FASTVIDEO_H3_VSA_SM89_KERNEL=bf16 \ +FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=6 MAX_JOBS=4 \ +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + sm89-bf16-480p-resident6 /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --no-vae-compile \ + --height 480 --width 832 --frames 243 --timed 2 +``` + +The BF16 kernel changes floating-point reductions: the ceramics clip is +visually coherent in the sampled contact sheet but is not identical to the +original-kernel clip (decoded-video SSIM 0.597653). This is a speed candidate, +not proof of quality equivalence. The original kernel remains the default. + +At `26390848b`, per-key-tile V scaling reduced experimental INT8-QK/FP8-PV +real-tensor error to 0.84–1.54%, with 1.77× fine-kernel speedup. Dynamic +per-query, per-key-block P scaling also preserves contributions that would +underflow with the fixed P scale; it measured 0.80–1.52% error and 1.71× +speedup. All 30 focused CUDA kernel/routing tests passed. These FP8-PV routes +remain microbenchmark-only and need clip validation. + +The current text encoder is the trimmed 50-layer Qwen3-VL with serialized +NVFP4 weights, dequantized to BF16 per linear on sm89. The existing serialized +blockwise FP8 encoder requires sm100+ and FlashInfer's Blackwell GEMM; it +cannot run on the 4090 as written. An Ada FP8 implementation would also need +encoder streaming because its weights are larger. First try fused NVFP4 +dequantization and avoid per-linear GPU scalar synchronization; then compare +a native sm89 FP8 encoder at equal prompts. Conditioning is only about +11.2 s of the current 111.3 s clip, so encoder work alone cannot dominate the +end-to-end gain. + +Remaining speed experiments: INT8-QK/BF16-PV same-seed clips; VAE compilation +and tile-batch tuning; more resident blocks after streaming the encoder; +fused FP8 GEMM epilogues and norm/activation quantization. The 16 GiB cap +still needs encoder streaming and a completed memory-capped run. The cached +recipe's 42 GiB anonymous host peak does not establish a 32 GB system-RAM +minimum. From 92fc26c9c683326fe987c5ccced3ea04f85c1725 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:23:24 -0700 Subject: [PATCH 12/24] [docs]: record five-second 4090 timing and FP8 encoder footprint --- scripts/benchmarks/minimax_h3_4090/README.md | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 3563d6d4d0..39f4983c01 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -233,6 +233,15 @@ visually coherent in the sampled contact sheet but is not identical to the original-kernel clip (decoded-video SSIM 0.597653). This is a speed candidate, not proof of quality equivalence. The original kernel remains the default. +The same cached/six-resident-block BF16 recipe at `26390848b`, with name +`sm89-bf16-480p-5s-resident6` and `--frames 124`, measured **65.27 s** median. +The legal frame count represents 5.167 s at 24 fps. Timed requests were +65.08/65.46 s after a 110.34 s warmup. Conditioning/denoise/video-decode +medians were 11.74/34.22/13.44 s. Peak GPU allocation was 21.62 GiB, +host anon 41.14 GiB and total cgroup 88.45 GiB. These generation wall times +include decode/export and exclude initial generator construction. No +profiling or QKV capture was enabled. + At `26390848b`, per-key-tile V scaling reduced experimental INT8-QK/FP8-PV real-tensor error to 0.84–1.54%, with 1.77× fine-kernel speedup. Dynamic per-query, per-key-block P scaling also preserves contributions that would @@ -244,7 +253,12 @@ The current text encoder is the trimmed 50-layer Qwen3-VL with serialized NVFP4 weights, dequantized to BF16 per linear on sm89. The existing serialized blockwise FP8 encoder requires sm100+ and FlashInfer's Blackwell GEMM; it cannot run on the 4090 as written. An Ada FP8 implementation would also need -encoder streaming because its weights are larger. First try fused NVFP4 +encoder streaming because its weights are larger. Reading the current +checkpoint tensor shapes gives 15.33 GiB total encoder weights, including +11.35 GiB packed values and 1.42 GiB block scales. Replacing those packed +values with FP8 while retaining the other tensors projects about 25.3 GiB +before activations (the FP8 block-scale overhead is small). This is a storage +estimate, not a measured FP8 encoder. First try fused NVFP4 dequantization and avoid per-linear GPU scalar synchronization; then compare a native sm89 FP8 encoder at equal prompts. Conditioning is only about 11.2 s of the current 111.3 s clip, so encoder work alone cannot dominate the From d986589b5e6a68414ac7534293d89472acd4ffb3 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:52:39 -0700 Subject: [PATCH 13/24] [feat]: stream the H3 encoder for consumer VRAM limits --- fastvideo/hooks/layerwise_offload.py | 26 +++-- .../encoders/minimax_h3_checkpoint_nvfp4.py | 6 +- .../models/encoders/minimax_h3_qwen3_vl.py | 29 +++++- fastvideo/models/loader/component_loader.py | 8 ++ .../basic/minimax_h3/minimax_h3_pipeline.py | 4 + .../stages/minimax_h3_conditioning.py | 3 +- .../test_minimax_h3_encoder_layerwise.py | 94 +++++++++++++++++++ 7 files changed, 158 insertions(+), 12 deletions(-) create mode 100644 fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index bc9ac0903e..6d33327464 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -172,7 +172,11 @@ def mutate_params_scope(self): self.state.on_init(self.state.module_ref) # pyright: ignore -def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): +def enable_layerwise_offload(model: nn.Module, + is_replace: bool = False, + *, + resident_blocks: int | None = None, + cyclic: bool = True): if torch.cuda.is_available(): device = torch.device("cuda", torch.cuda.current_device()) else: @@ -183,12 +187,15 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): # The first N entries skip offloading and stay wherever the model is placed (normally the # GPU), so a GPU with spare memory streams only the remainder over PCIe. import os - try: - resident = max(0, int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0"))) - except ValueError: - logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%r", - os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS")) - resident = 0 + if resident_blocks is not None: + resident = max(0, resident_blocks) + else: + try: + resident = max(0, int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0"))) + except ValueError: + logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%r", + os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS")) + resident = 0 for name, submodule in model.named_children(): if isinstance(submodule, nn.ModuleList): for idx, module_entry in enumerate(submodule): @@ -214,6 +221,7 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): return raise ValueError("No nn.ModuleList found in the model for layerwise offloading.") - # circular linking of states + # Repeated DiT steps prefetch the first block after the last. A once-per-request + # encoder can skip that unused copy and release every layer after its forward. for i in range(len(state_list)): - state_list[i].next_state = state_list[(i + 1) % len(state_list)] + state_list[i].next_state = state_list[(i + 1) % len(state_list)] if cyclic or i + 1 < len(state_list) else None diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 686346e8ba..6d0e0e6deb 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -417,6 +417,10 @@ def process_weights_after_loading(self, layer: nn.Module) -> None: # ``mm_fp4`` folds both global scales into one multiplier. Activations use a # unit global scale, so the multiplier is the inverse weight global scale. device = weight_scale.device + # Serialized weights are immutable between post-load hooks. Keeping the + # validated scalar on the host avoids a CUDA synchronization per linear + # on the BF16 fallback used by consumer GPUs. + layer._nvfp4_dequant_global_scale = global_scale layer.register_buffer("_nvfp4_alpha", torch.tensor(1.0 / global_scale, dtype=torch.float32, device=device), persistent=False) layer.register_buffer("_nvfp4_x_global_scale", torch.ones((), dtype=torch.float32, device=device), @@ -429,7 +433,7 @@ def _apply_finalized(layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor # Pre-Blackwell GPUs have no FP4 GEMM: expand this layer's weight to bf16 for the one call. # The encoder runs once per request, so the transient weight is cheaper than keeping a bf16 copy. weight = dequantize_serialized_nvfp4(layer.weight_packed, layer.weight_scale, - float(layer.weight_global_scale.item()), x.dtype) + layer._nvfp4_dequant_global_scale, x.dtype) return torch.nn.functional.linear(x, weight, None if bias is None else bias.to(x.dtype)) original_shape = x.shape if x.numel() == 0: diff --git a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py index c348140d1d..4886ae1db8 100644 --- a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py +++ b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py @@ -690,8 +690,17 @@ def encode_ids( if (pixel_values_videos is None) != (video_grid_thw is None): raise ValueError("pixel_values_videos and video_grid_thw must be provided together") + stream_device = getattr(self, "_h3_encoder_layerwise_device", None) + if stream_device is not None and (pixel_values is not None or pixel_values_videos is not None): + raise ValueError("Layerwise H3 encoder currently supports text-only conditioning; " + "disable FASTVIDEO_H3_ENCODER_LAYERWISE for visual references") + input_ids = input_ids.unsqueeze(0) - inputs_embeds = self.language_model.embed_tokens(input_ids) + embedding_ids = input_ids.to("cpu") if stream_device is not None else input_ids + inputs_embeds = self.language_model.embed_tokens(embedding_ids) + if stream_device is not None: + inputs_embeds = inputs_embeds.to(stream_device) + input_ids = input_ids.to(stream_device) image_mask = None video_mask = None @@ -746,6 +755,24 @@ def encode_ids( raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}") return hidden_states[0] + def prepare_layerwise_offload(self, device: torch.device) -> None: + """Stream language layers for text-only CUDA inference, retaining embeddings on CPU.""" + if getattr(self, "_h3_encoder_layerwise_device", None) is not None: + return + if device.type != "cuda": + raise ValueError("Layerwise H3 encoder requires CUDA") + from fastvideo.distributed import get_tp_world_size + from fastvideo.hooks.layerwise_offload import enable_layerwise_offload + + if get_tp_world_size() != 1: + raise ValueError("Layerwise H3 encoder requires tensor parallel size 1") + self.to("cpu") + self.language_model.rotary_emb.to(device) + if self.language_model.norm is not None: + self.language_model.norm.to(device) + enable_layerwise_offload(self.language_model, resident_blocks=0, cyclic=False) + self._h3_encoder_layerwise_device = device + def forward( self, input_ids: torch.Tensor, diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 7e668a1709..ab7bf07715 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -467,6 +467,14 @@ def load_model( # Explicitly move model to target device after loading weights model = model.to(target_device) + prepare_layerwise = getattr(model, "prepare_layerwise_offload", None) + if os.environ.get("FASTVIDEO_H3_ENCODER_LAYERWISE", "0") == "1" and callable(prepare_layerwise): + if target_device.type != "cpu": + raise ValueError("Layerwise H3 encoder requires text_encoder_cpu_offload=True") + prepare_layerwise(runtime_device) + use_cpu_offload = False + logger.info("Enabled text-only layerwise H3 encoder with CPU token embeddings") + from fastvideo.platforms import current_platform if use_cpu_offload and checkpoint_quant_config is not None: diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 1edee22c4a..9b2cb09433 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -464,6 +464,10 @@ def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: stage.conditioner = self.get_module("text_encoder") def _move_module(self, module: Any, device: str | torch.device) -> bool: + if getattr(module, "_h3_encoder_layerwise_device", None) is not None: + # Layer hooks own placement; moving the whole encoder would restore + # every weight at once and defeat its VRAM bound. + return True if _module_has_dtensor_params(module): return False if os.environ.get("FASTVIDEO_H3_PINNED_SWAP", "1") == "1": diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py index d39128e91d..1e26ee5409 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_conditioning.py @@ -261,7 +261,8 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward device = get_local_torch_device() first_param = next(self.conditioner.parameters(), None) moved_for_forward = (fastvideo_args.text_encoder_cpu_offload and first_param is not None - and not isinstance(first_param, DTensor)) + and not isinstance(first_param, DTensor) + and getattr(self.conditioner, "_h3_encoder_layerwise_device", None) is None) if moved_for_forward: self.conditioner.to(device) try: diff --git a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py new file mode 100644 index 0000000000..3661ddda02 --- /dev/null +++ b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py @@ -0,0 +1,94 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Text-only H3 encoder streaming parity and placement contracts.""" +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch + +from fastvideo.hooks.hooks import ModuleHookManager +from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLArchConfig, MiniMaxH3Qwen3VLConfig +from fastvideo.models.encoders.minimax_h3_checkpoint_nvfp4 import MiniMaxH3SerializedNVFP4Config +from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner +from fastvideo.models.loader.text_encoder_quantization import _process_quantized_text_encoder_weights +from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import MiniMaxH3Pipeline +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage +from fastvideo.pipelines.pipeline_batch_info import ForwardBatch + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for encoder streaming") +@pytest.mark.parametrize("quantized", [False, True]) +def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, quantized): + # DiT residency must not accidentally keep encoder layers resident too. + monkeypatch.setenv("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "6") + config = MiniMaxH3Qwen3VLConfig() + config.arch_config = MiniMaxH3Qwen3VLArchConfig( + vocab_size=64, hidden_size=128, intermediate_size=256, + num_hidden_layers=3, num_hidden_layers_override=2, output_hidden_state_index=2, + num_attention_heads=1, num_key_value_heads=1, head_dim=128, + rope_scaling={"mrope_interleaved": True, "mrope_section": [32, 16, 16], "rope_type": "default"}, + vision_depth=1, vision_hidden_size=64, vision_intermediate_size=128, + vision_num_heads=1, vision_deepstack_visual_indexes=(), vision_out_hidden_size=128, + ) + config.quant_config = MiniMaxH3SerializedNVFP4Config() if quantized else None + torch.manual_seed(81) + model = MiniMaxH3Qwen3VLConditioner(config).to(dtype=torch.bfloat16).eval() + for name, parameter in model.named_parameters(): + if name.endswith("weight_packed"): + parameter.data.random_(0, 256) + elif name.endswith("weight_scale"): + parameter.data.fill_(0x38) + elif name.endswith("weight_global_scale"): + parameter.data.fill_(2.0) + else: + parameter.data.normal_(std=0.02) + if quantized: + _process_quantized_text_encoder_weights(model, torch.device("cuda")) + linear = model.language_model.layers[0].self_attn.q_proj.to("cuda") + x = torch.randn(3, 128, device="cuda", dtype=torch.bfloat16) + expected_linear = linear(x)[0] + with patch.object(torch.Tensor, "item", side_effect=AssertionError("Unexpected device scalar read")): + actual_linear = linear(x)[0] + torch.testing.assert_close(actual_linear, expected_linear, rtol=0, atol=0) + ids = torch.tensor([1, 7, 4, 21, 5, 31, 18], device="cuda") + model.to("cuda") + expected = model.encode_ids(ids) + assert torch.isfinite(expected).all() + model.to("cpu") + model.prepare_layerwise_offload(torch.device("cuda")) + model.prepare_layerwise_offload(torch.device("cuda")) # repeated setup is harmless + assert model.language_model.embed_tokens.weight.device.type == "cpu" + assert next(model.visual.parameters()).device.type == "cpu" + for _ in range(2): + actual = model.encode_ids(ids) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + for layer in model.language_model.layers: + assert all(parameter.numel() == 0 for parameter in layer.parameters()) + manager = ModuleHookManager.get_from(layer) + assert manager is not None + assert not manager.forward_hooks["LayerwiseOffloadHook"].state.gpu_named_parameters + with pytest.raises(ValueError, match="text-only"): + model.encode_ids(ids, pixel_values=torch.zeros(1, device="cuda"), + image_grid_thw=torch.ones(1, 3, device="cuda", dtype=torch.int64)) + + +@pytest.mark.parametrize("device", ["cpu", "cuda"]) +def test_pipeline_does_not_move_streamed_encoder_whole(device): + module = SimpleNamespace(_h3_encoder_layerwise_device=torch.device("cuda")) + module.to = lambda *_: pytest.fail("Whole encoder move defeats streaming") + assert MiniMaxH3Pipeline._move_module(None, module, device) + + +def test_conditioning_stage_keeps_streamed_encoder_placement(monkeypatch): + import fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning as conditioning + + module = SimpleNamespace(_h3_encoder_layerwise_device=torch.device("cuda")) + module.parameters = lambda: iter([torch.empty(1)]) + module.to = lambda *_: pytest.fail("Conditioning must retain layerwise placement") + stage = MiniMaxH3ConditioningStage.__new__(MiniMaxH3ConditioningStage) + stage.conditioner, stage.ref2va = module, False + stage._encode_fl2va = lambda *_: (torch.zeros(1, 2, 128), torch.zeros(2, dtype=torch.int32)) + monkeypatch.setattr(conditioning, "get_local_torch_device", lambda: torch.device("cpu")) + batch = ForwardBatch(data_type="video", prompt="streaming parity") + output = stage.forward(batch, SimpleNamespace(text_encoder_cpu_offload=True)) + assert output.prompt_embeds[0].shape == (1, 2, 128) From 887deaab12be75ae7a44e6c635825de510330f3c Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 12:06:21 -0700 Subject: [PATCH 14/24] [perf]: fuse serialized NVFP4 encoder weight dequantization --- .../layers/quantization/nvfp4_dequant.py | 54 +++++++++++++++++ .../encoders/minimax_h3_checkpoint_nvfp4.py | 9 ++- .../test_minimax_h3_encoder_layerwise.py | 11 +++- .../ops/quantization/test_nvfp4_dequant.py | 20 +++++++ .../minimax_h3_4090/bench_encoder_dequant.py | 59 +++++++++++++++++++ 5 files changed, 148 insertions(+), 5 deletions(-) create mode 100644 fastvideo/layers/quantization/nvfp4_dequant.py create mode 100644 fastvideo/tests/ops/quantization/test_nvfp4_dequant.py create mode 100644 scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py diff --git a/fastvideo/layers/quantization/nvfp4_dequant.py b/fastvideo/layers/quantization/nvfp4_dequant.py new file mode 100644 index 0000000000..5bc3c2b342 --- /dev/null +++ b/fastvideo/layers/quantization/nvfp4_dequant.py @@ -0,0 +1,54 @@ +# SPDX-License-Identifier: Apache-2.0 +"""One-pass serialized NVFP4 weight expansion for BF16 consumer-GPU compute.""" +import torch +import triton +import triton.language as tl + + +@triton.jit +def _dequantize_nvfp4(P, S, OUT, N: tl.constexpr, K: tl.constexpr, INVERSE_SCALE, BLOCK: tl.constexpr): + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + valid = offsets < N * K + row, col = offsets // K, offsets % K + packed = tl.load(P + row * (K // 2) + col // 2, valid, other=0).to(tl.uint32) + code = (packed >> ((col % 2) * 4)) & 15 + magnitude = (code & 7).to(tl.float32) + value = tl.where(magnitude < 4, magnitude * 0.5, tl.where(magnitude < 6, magnitude - 2, magnitude * 2 - 8)) + value = value * tl.where((code & 8) != 0, -1.0, 1.0) + group = col // 16 + # FlashInfer layout_128x4: [row_tile, col_tile, row%32, row//32%4, col%4]. + scale_index = ((((row // 128) * (K // 64) + group // 4) * 32 + row % 32) * 4 + (row // 32) % 4) * 4 + group % 4 + scale = tl.load(S + scale_index, valid, other=0.0).to(tl.float32) + output = (value * scale) * INVERSE_SCALE + tl.store(OUT + offsets, output, valid) + + +def dequantize_nvfp4_cuda(packed: torch.Tensor, + scales: torch.Tensor, + global_scale: float, + dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + """Expand E2M1 nibbles and swizzled E4M3 scales without full FP32 intermediates.""" + if not packed.is_cuda or scales.device != packed.device: + raise ValueError("NVFP4 fused dequantization requires tensors on the same CUDA device") + if packed.ndim != 2 or packed.dtype != torch.uint8 or scales.dtype != torch.uint8: + raise ValueError("NVFP4 fused dequantization requires packed uint8 weights and scales") + if not packed.is_contiguous() or not scales.is_contiguous(): + raise ValueError("NVFP4 fused dequantization requires contiguous tensors") + rows, cols = packed.shape[0], packed.shape[1] * 2 + if rows % 128 or cols % 64 or scales.numel() != rows * cols // 16: + raise ValueError("NVFP4 fused dequantization requires exact 128x4 scale geometry") + if dtype not in (torch.bfloat16, torch.float16, torch.float32): + raise ValueError("NVFP4 fused dequantization requires a floating output dtype") + output = torch.empty((rows, cols), dtype=dtype, device=packed.device) + _dequantize_nvfp4[(triton.cdiv(rows * cols, 1024), )]( + packed, + scales.view(torch.float8_e4m3fn), + output, + rows, + cols, + # Match Torch's CPU-scalar division: form the + # reciprocal in double, then cast to FP32. + 1.0 / global_scale, + BLOCK=1024, + num_warps=4) + return output diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 6d0e0e6deb..2d5d2897e6 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -37,6 +37,7 @@ content checks a copied tensor can still fail. """ +import os from typing import Any import torch @@ -421,6 +422,7 @@ def process_weights_after_loading(self, layer: nn.Module) -> None: # validated scalar on the host avoids a CUDA synchronization per linear # on the BF16 fallback used by consumer GPUs. layer._nvfp4_dequant_global_scale = global_scale + layer._nvfp4_fused_dequant = os.environ.get("FASTVIDEO_H3_ENCODER_FUSED_DEQUANT", "0") == "1" layer.register_buffer("_nvfp4_alpha", torch.tensor(1.0 / global_scale, dtype=torch.float32, device=device), persistent=False) layer.register_buffer("_nvfp4_x_global_scale", torch.ones((), dtype=torch.float32, device=device), @@ -432,8 +434,11 @@ def _apply_finalized(layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor if not _fp4_gemm_supported(layer.weight_packed.device): # Pre-Blackwell GPUs have no FP4 GEMM: expand this layer's weight to bf16 for the one call. # The encoder runs once per request, so the transient weight is cheaper than keeping a bf16 copy. - weight = dequantize_serialized_nvfp4(layer.weight_packed, layer.weight_scale, - layer._nvfp4_dequant_global_scale, x.dtype) + dequantize = dequantize_serialized_nvfp4 + if layer._nvfp4_fused_dequant: + from fastvideo.layers.quantization.nvfp4_dequant import dequantize_nvfp4_cuda + dequantize = dequantize_nvfp4_cuda + weight = dequantize(layer.weight_packed, layer.weight_scale, layer._nvfp4_dequant_global_scale, x.dtype) return torch.nn.functional.linear(x, weight, None if bias is None else bias.to(x.dtype)) original_shape = x.shape if x.numel() == 0: diff --git a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py index 3661ddda02..a13712a7b3 100644 --- a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py +++ b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py @@ -17,10 +17,11 @@ @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for encoder streaming") -@pytest.mark.parametrize("quantized", [False, True]) -def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, quantized): +@pytest.mark.parametrize("quantized,fused", [(False, False), (True, False), (True, True)]) +def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, quantized, fused): # DiT residency must not accidentally keep encoder layers resident too. monkeypatch.setenv("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "6") + monkeypatch.setenv("FASTVIDEO_H3_ENCODER_FUSED_DEQUANT", "0") config = MiniMaxH3Qwen3VLConfig() config.arch_config = MiniMaxH3Qwen3VLArchConfig( vocab_size=64, hidden_size=128, intermediate_size=256, @@ -39,7 +40,7 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup elif name.endswith("weight_scale"): parameter.data.fill_(0x38) elif name.endswith("weight_global_scale"): - parameter.data.fill_(2.0) + parameter.data.fill_(2.7) else: parameter.data.normal_(std=0.02) if quantized: @@ -54,6 +55,10 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup model.to("cuda") expected = model.encode_ids(ids) assert torch.isfinite(expected).all() + if fused: + for layer in model.modules(): + if hasattr(layer, "_nvfp4_fused_dequant"): + layer._nvfp4_fused_dequant = True model.to("cpu") model.prepare_layerwise_offload(torch.device("cuda")) model.prepare_layerwise_offload(torch.device("cuda")) # repeated setup is harmless diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_dequant.py b/fastvideo/tests/ops/quantization/test_nvfp4_dequant.py new file mode 100644 index 0000000000..6eb9f6aa87 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_dequant.py @@ -0,0 +1,20 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused weight expansion against the independent serialized Torch decoder.""" +import pytest +import torch + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for Triton NVFP4 decoder") +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +@pytest.mark.parametrize("global_scale", [1.0, 2.7, 438.912]) +def test_fused_nvfp4_expands_all_codes_and_scale_tiles(dtype, global_scale): + from fastvideo.layers.quantization.nvfp4_dequant import dequantize_nvfp4_cuda + from fastvideo.models.encoders.minimax_h3_checkpoint_nvfp4 import dequantize_serialized_nvfp4 + + torch.manual_seed(23) + # Multiple row/column tiles distinguish the swizzle from a row-major decoder. + packed = torch.arange(256, device="cuda", dtype=torch.uint8).repeat(256, 1) + scales = torch.randint(0, 127, (256, 32), device="cuda", dtype=torch.uint8) + reference = dequantize_serialized_nvfp4(packed, scales, global_scale, dtype) + actual = dequantize_nvfp4_cuda(packed, scales, global_scale, dtype) + torch.testing.assert_close(actual, reference, rtol=0, atol=0) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py b/scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py new file mode 100644 index 0000000000..798793e678 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/bench_encoder_dequant.py @@ -0,0 +1,59 @@ +"""Compare serialized NVFP4 weight expansion on an idle consumer GPU.""" +import argparse +import json +import os +import pathlib +import time + +import torch + +from fastvideo.layers.quantization.nvfp4_dequant import dequantize_nvfp4_cuda +from fastvideo.models.encoders.minimax_h3_checkpoint_nvfp4 import dequantize_serialized_nvfp4 + + +def measure(fn, packed, scales, global_scale): + for _ in range(3): + fn(packed, scales, global_scale) + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + start, stop = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + wall = time.perf_counter() + start.record() + for _ in range(10): + fn(packed, scales, global_scale) + stop.record() + stop.synchronize() + return {"gpu_ms": start.elapsed_time(stop) / 10, + "wall_ms": (time.perf_counter() - wall) * 100, + "peak_extra_gib": (torch.cuda.max_memory_allocated() - baseline) / 2**30} + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--output", type=pathlib.Path, required=True) + args = ap.parse_args() + rows = [] + global_scale = float(torch.tensor(317.224, dtype=torch.float32)) + torch.manual_seed(63) + for n, k in [(8192, 8192), (25600, 8192), (8192, 25600)]: + packed = torch.randint(0, 256, (n, k // 2), device="cuda", dtype=torch.uint8) + scales = torch.randint(0, 127, (n, k // 16), device="cuda", dtype=torch.uint8) + reference = dequantize_serialized_nvfp4(packed, scales, global_scale) + fused = dequantize_nvfp4_cuda(packed, scales, global_scale) + torch.testing.assert_close(fused, reference, rtol=0, atol=0) + del reference, fused + row = {"n": n, "k": k, "bf16_exact": True, + "torch": measure(dequantize_serialized_nvfp4, packed, scales, global_scale), + "fused": measure(dequantize_nvfp4_cuda, packed, scales, global_scale)} + row["speedup"] = row["torch"]["gpu_ms"] / row["fused"]["gpu_ms"] + print(json.dumps(row), flush=True) + rows.append(row) + del packed, scales + args.output.write_text(json.dumps({"gpu": torch.cuda.get_device_name(), "torch": torch.__version__, + "source_commit": os.environ.get("FASTVIDEO_SOURCE_COMMIT"), + "global_scale": global_scale, "rows": rows}, indent=2)) + + +if __name__ == "__main__": + main() From 2aa19c4b986ad264829871df8a12c12c404d6f55 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 12:30:13 -0700 Subject: [PATCH 15/24] [perf]: release H3 fine-attention copies before the gated merge --- fastvideo/attention/backends/video_sparse_attn_h3.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index 86629ade78..db82b76f60 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -906,6 +906,9 @@ def forward( # type: ignore[override] if has_sm100a_pair and use_sm100a: out_bhsd = out_bhsd[:, :, :logical_seq_len] out = out_bhsd.transpose(1, 2).contiguous() + # Fine attention is complete. Release its layout copies before the + # gated compression merge creates full-sequence temporaries. + del q_bhsd, k_bhsd, v_bhsd, out_bhsd else: out, _ = block_sparse_attn_256_bshd( logical_query, From 9c9f1edcd07c346b11f00638bfcf79910494501a Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 12:30:13 -0700 Subject: [PATCH 16/24] [perf]: share VAE INT8 input preparation and avoid weight copies --- .../models/vaes/minimax_h3_int8_convrot.py | 63 +++++++++++----- fastvideo/models/vaes/minimax_h3_video.py | 11 ++- .../tests/vaes/test_minimax_h3_int8_shared.py | 73 +++++++++++++++++++ 3 files changed, 127 insertions(+), 20 deletions(-) create mode 100644 fastvideo/tests/vaes/test_minimax_h3_int8_shared.py diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py index 649541fd19..c4f366dd7d 100644 --- a/fastvideo/models/vaes/minimax_h3_int8_convrot.py +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -104,6 +104,7 @@ def __init__( self.out_features = out_features self.convrot = convrot self.group_size = group_size + self._transpose_view = os.environ.get("FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW", "0") == "1" self.register_buffer("weight", torch.empty(out_features, in_features, dtype=torch.int8)) self.register_buffer("weight_scale", torch.empty(out_features, 1, dtype=torch.float32)) if bias: @@ -123,32 +124,60 @@ def _dequant_int8_gemm( # int32 acc is ~K·127² and overflows fp16 before 1/127 scales land. return acc.float() * x_scale.float() * weight_scale.t().float() - def forward(self, x: torch.Tensor) -> torch.Tensor: - original_shape = x.shape - x_2d = x.reshape(-1, original_shape[-1]).contiguous() + def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + x_2d = x.reshape(-1, x.shape[-1]).contiguous() if self.convrot: x_2d = rotate_activation(x_2d, self.group_size) - if x_2d.device.type == "cuda" and x_2d.shape[-1] % 8 == 0: - row_max = x_2d.abs().amax(dim=-1, keepdim=True).clamp_min(1e-30) - x_scale = row_max / 127.0 - x_q = (x_2d / x_scale).round().clamp(-128, 127).to(torch.int8) - # torch._int_mm requires M > 16. VAE decode is far above that; - # pad only the leftover short rows. - rows = x_q.shape[0] - if rows <= 16: - pad = 17 - rows - x_q = F.pad(x_q, (0, 0, 0, pad)) - x_scale = F.pad(x_scale, (0, 0, 0, pad)) - acc = torch._int_mm(x_q, self.weight.t().contiguous())[:rows] - x_scale = x_scale[:rows] - out = self._dequant_int8_gemm(acc, x_scale, self.weight_scale) + row_max = x_2d.abs().amax(dim=-1, keepdim=True).clamp_min(1e-30) + x_scale = row_max / 127.0 + x_q = (x_2d / x_scale).round().clamp(-128, 127).to(torch.int8) + rows = x_q.shape[0] + if rows <= 16: + pad = 17 - rows + x_q = F.pad(x_q, (0, 0, 0, pad)) + x_scale = F.pad(x_scale, (0, 0, 0, pad)) + return x_q, x_scale + + def forward_quantized(self, x_q: torch.Tensor, x_scale: torch.Tensor, + original_shape: tuple[int, ...], dtype: torch.dtype) -> torch.Tensor: + rows = math.prod(original_shape[:-1]) + weight = self.weight.t() + if not self._transpose_view: + weight = weight.contiguous() + acc = torch._int_mm(x_q, weight)[:rows] + out = self._dequant_int8_gemm(acc, x_scale[:rows], self.weight_scale) + if self.bias is not None: + out = out + self.bias.float() + return out.to(dtype=dtype).view(*original_shape[:-1], self.out_features) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + original_shape = x.shape + if x.device.type == "cuda" and x.shape[-1] % 8 == 0: + return self.forward_quantized(*self.quantize_input(x), original_shape, x.dtype) else: + x_2d = x.reshape(-1, original_shape[-1]).contiguous() + if self.convrot: + x_2d = rotate_activation(x_2d, self.group_size) out = F.linear(x_2d.float(), self._dequant_weight(torch.float32)) if self.bias is not None: out = out + self.bias.float() return out.to(dtype=x.dtype).view(*original_shape[:-1], self.out_features) +def shared_int8_projections(layers: tuple[nn.Module, ...], x: torch.Tensor) -> tuple[torch.Tensor, ...]: + """Reuse identical ConvRot/row quantization while retaining each projection's INT8 GEMM.""" + first = layers[0] + compatible = (x.is_cuda and x.shape[-1] % 8 == 0 + and all(isinstance(layer, Int8ConvRotLinear) for layer in layers)) + if compatible: + compatible = all((layer.in_features, layer.convrot, layer.group_size) + == (first.in_features, first.convrot, first.group_size) for layer in layers) + if not compatible: + return tuple(layer(x) for layer in layers) + x_q, x_scale = first.quantize_input(x) + return tuple(layer.forward_quantized(x_q, x_scale, x.shape, x.dtype) for layer in layers) + + def _int8_linear_from_tensors( weight: torch.Tensor, scale: torch.Tensor, diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index 21be0864a9..223aae259e 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -301,6 +301,7 @@ def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: self.heads = heads self.dim_head = dim_head self.use_bias = bias + self._share_int8_qkv = os.environ.get("FASTVIDEO_H3_VAE_INT8_SHARED_QKV", "0") == "1" inner_dim = heads * dim_head self.norm_q = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=False) self.norm_k = nn.RMSNorm(dim_head, eps=eps, elementwise_affine=False) @@ -340,9 +341,13 @@ def forward( 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)) + if self._share_int8_qkv and not torch.is_grad_enabled() and not torch.compiler.is_compiling(): + from fastvideo.models.vaes.minimax_h3_int8_convrot import shared_int8_projections + projections = shared_int8_projections((self.to_q, self.to_k, self.to_v), hidden_states) + else: + projections = tuple(layer(hidden_states) for layer in (self.to_q, self.to_k, self.to_v)) + query, key, value = (projection.unflatten(2, (self.heads, -1)) for projection in projections) + del projections query = self.norm_q(query.float()).to(query.dtype) key = self.norm_k(key.float()).to(key.dtype) diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py new file mode 100644 index 0000000000..88bb79e23b --- /dev/null +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py @@ -0,0 +1,73 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exact parity of shared VAE input preparation and strided INT8 weight GEMMs.""" +from unittest.mock import patch + +import pytest +import torch + +from fastvideo.models.vaes.minimax_h3_int8_convrot import Int8ConvRotLinear, shared_int8_projections + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for INT8 GEMM") +@pytest.mark.parametrize("rows", [3, 17, 129]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("convrot", [False, True]) +def test_shared_int8_and_transpose_views_are_exact(rows, dtype, convrot, monkeypatch): + monkeypatch.setenv("FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW", "0") + torch.manual_seed(73) + layers = tuple(Int8ConvRotLinear(256, out, bias=index != 1, convrot=convrot, group_size=256) + .to("cuda") for index, out in enumerate([128, 256, 64])) + for layer in layers: + layer.weight.random_(-127, 128) + layer.weight_scale.uniform_(0.0001, 0.03) + if layer.bias is not None: + layer.bias.data.normal_() + x = torch.randn(1, rows, 256, device="cuda", dtype=dtype) + x[0, 0].zero_() # clamp/padding semantics must also survive sharing + with torch.inference_mode(): + expected = tuple(layer(x) for layer in layers) + with patch.object(Int8ConvRotLinear, "quantize_input", autospec=True, + side_effect=Int8ConvRotLinear.quantize_input) as quant: + shared = shared_int8_projections(layers, x) + assert quant.call_count == 1 + for layer in layers: + layer._transpose_view = True + views = shared_int8_projections(layers, x) + for ref, actual, view in zip(expected, shared, views, strict=True): + assert torch.isfinite(ref).all() + torch.testing.assert_close(actual, ref, rtol=0, atol=0) + torch.testing.assert_close(view, ref, rtol=0, atol=0) + + +def test_shared_int8_keeps_cpu_fallback_exact(): + layers = tuple(Int8ConvRotLinear(16, 8, bias=False, convrot=False, group_size=16) for _ in range(3)) + for layer in layers: + layer.weight.fill_(1) + layer.weight_scale.fill_(0.01) + x = torch.ones(3, 16) + for actual, ref in zip(shared_int8_projections(layers, x), (layer(x) for layer in layers), strict=True): + torch.testing.assert_close(actual, ref, rtol=0, atol=0) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for VAE attention parity") +def test_vae_attention_shares_quantized_projections_exactly(distributed_setup, monkeypatch): + from fastvideo.models.vaes.minimax_h3_video import MiniMaxH3VideoAttention + + monkeypatch.setenv("FASTVIDEO_H3_VAE_INT8_SHARED_QKV", "0") + monkeypatch.setenv("FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW", "0") + torch.manual_seed(49) + attention = MiniMaxH3VideoAttention(256, 2, 128).to("cuda").eval() + for name in ("to_q", "to_k", "to_v"): + layer = Int8ConvRotLinear(256, 256, bias=True, convrot=True, group_size=256).to("cuda") + layer.weight.random_(-8, 9) + layer.weight_scale.fill_(0.01) + layer.bias.data.normal_(std=0.1) + setattr(attention, name, layer) + x = torch.randn(2, 33, 256, device="cuda") + with torch.inference_mode(): + expected = attention(x) + attention._share_int8_qkv = True + for layer in (attention.to_q, attention.to_k, attention.to_v): + layer._transpose_view = True + actual = attention(x) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) From e457b68341fce46f1b20e6b2f935cd9f692b1b5b Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 12:47:53 -0700 Subject: [PATCH 17/24] [docs]: record 12 GiB and streamed 4090 results --- scripts/benchmarks/minimax_h3_4090/README.md | 58 ++++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 39f4983c01..905f70d0e6 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -270,3 +270,61 @@ fused FP8 GEMM epilogues and norm/activation quantization. The 16 GiB cap still needs encoder streaming and a completed memory-capped run. The cached recipe's 42 GiB anonymous host peak does not establish a 32 GB system-RAM minimum. + + +## Streamed encoder and smaller VRAM caps + +At `753e560f6`, `FASTVIDEO_H3_ENCODER_LAYERWISE=1` streams the language +layers separately from DiT residency, retaining token embeddings and unused +vision modules on the CPU. This route currently supports text-only T2VA; +visual references fail explicitly. The pipeline preserves the streamed +placement. Twenty-eight offload/encoder/stage tests passed, including exact +repeated BF16 and NVFP4 parity. + +At `78540b635`, `FASTVIDEO_H3_ENCODER_FUSED_DEQUANT=1` expands packed NVFP4 +weights with one Triton pass. Fifteen strict tests passed across FP32, FP16, +BF16, swizzled scales and repeated encoder forwards. On three large stress +matrices the expansion was 25.3–25.9× faster and used 8× less temporary GPU +memory than Torch expansion. This is a dequantization microbenchmark; the +whole conditioning stage measured 0.55–0.62 seconds in the clip runs below. +The current encoder remains NVFP4 storage with BF16 GEMMs on Ada. + +All rows use 832×480, 243 frames, eight DMD forwards, sparsity 0.8, tile 64, +cached components, INT8 QK/BF16 PV and the eager light H3 VAE. Each median +has one warmup and two timed requests. Source is `78540b635` except the +16 GiB row (`753e560f6`, before fused dequantization). + +| 4090 configuration | Median e2e | Timed requests | Denoise | Video decode | Peak GPU allocated | Peak host anon | +| --- | --- | --- | --- | --- | --- | --- | +| 16 GiB cap, 6 resident | 104.03 s | 100.10 / 107.96 s | 72.42 s | 25.54 s | 11.28 GiB | 39.86 GiB | +| 12 GiB cap, 0 resident | 107.27 s | 111.02 / 103.52 s | 77.38 s | 25.48 s | 8.68 GiB | 42.35 GiB | +| Uncapped, 30 resident | 98.97 s | 98.88 / 99.05 s | 69.12 s | 25.47 s | 21.69 GiB | 29.44 GiB | + +Set `FASTVIDEO_CUDA_MEMORY_CAP_GIB=12` for the 12 GiB recipe; unset it for +the full card. Set resident blocks to the table value. Both fused rows use: + +```bash +FASTVIDEO_SOURCE_COMMIT=78540b635 \ +FASTVIDEO_H3_PARK_MODULES=vae,audio_vae \ +FASTVIDEO_H3_ENCODER_LAYERWISE=1 FASTVIDEO_H3_ENCODER_FUSED_DEQUANT=1 \ +FASTVIDEO_H3_VSA_TILE_FIRST=1 FASTVIDEO_H3_VSA_SM89_KERNEL=int8 \ +FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=30 FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 \ +FASTVIDEO_H3_VAE_TILE_BATCH=28 MAX_JOBS=4 \ +python -P /workspace/fastvideo/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + sm89-int8-480p-resident30-fused /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --no-vae-compile --height 480 --width 832 --frames 243 --timed 2 +``` + +The 12 GiB and 30-resident ceramics clips have identical decoded-video and +PCM audio hashes to the six-resident INT8 clip. These placement and dequant +changes preserve the candidate's output; quality equivalence of INT8 +attention to the original BF16 checkpoint still requires motion/speech review. +Allocator caps emulate available VRAM on a 4090, not another card's speed. +Host peaks include all pod processes; actual 32 GB host-limit support has +not been established. The 30-resident warmup reached 30.26 GiB anonymous +memory and the timed runs reached 77.50 GiB total cgroup usage including cache. + +The 8 GiB cap at `fcdba37fc` completed its warmup but OOMed on the timed +harbor prompt in fine attention. Do not report it as supported. A 34-resident +experiment completed denoising but OOMed during VAE INT8 epilogue allocation. +Both failures motivate subsequent memory work rather than speed claims. From 7a0d7d33b4751dca71d5fa4d33bcbf82de36ae7d Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:03:45 -0700 Subject: [PATCH 18/24] [perf]: read H3 INT8 sparse attention from existing layouts --- .../backends/minimax_h3_sparse_int8.py | 34 ++++++++++----- .../backends/video_sparse_attn_h3.py | 19 ++++++--- .../attention/test_minimax_h3_sparse_int8.py | 41 +++++++++++++++++++ 3 files changed, 77 insertions(+), 17 deletions(-) diff --git a/fastvideo/attention/backends/minimax_h3_sparse_int8.py b/fastvideo/attention/backends/minimax_h3_sparse_int8.py index 4c8c5ddad1..5854d25196 100644 --- a/fastvideo/attention/backends/minimax_h3_sparse_int8.py +++ b/fastvideo/attention/backends/minimax_h3_sparse_int8.py @@ -16,11 +16,13 @@ @triton.jit -def _quantize_qk(X, Mean, VBS, Y, Scale, L: tl.constexpr, D: tl.constexpr, CENTER: tl.constexpr, ROWS: tl.constexpr): +def _quantize_qk(X, Mean, VBS, Y, Scale, L: tl.constexpr, D: tl.constexpr, H: tl.constexpr, XB: tl.constexpr, + XH: tl.constexpr, XS: tl.constexpr, XD: tl.constexpr, CENTER: tl.constexpr, ROWS: tl.constexpr): hz = tl.program_id(1) rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS) cols = tl.arange(0, D) - x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :], rows[:, None] < L, 0).to(tl.float32) + offset = (hz // H) * XB + (hz % H) * XH + rows[:, None] * XS + cols[None, :] * XD + x = tl.load(X + offset, rows[:, None] < L, 0).to(tl.float32) if CENTER: mean = tl.load(Mean + hz * D + cols) valid_size = tl.load(VBS + rows // 64, rows < L, 0) @@ -57,8 +59,9 @@ def _quantize_v_tiles(X, Y, Scale, L: tl.constexpr, D: tl.constexpr): configs=[triton.Config({}, num_warps=w, num_stages=s) for w, s in ((4, 2), (4, 3), (4, 4), (8, 2), (8, 3))], key=["L", "D"]) @triton.jit -def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr, - INT8_QK: tl.constexpr, FP8_PV: tl.constexpr, V_TILE: tl.constexpr, P_DYNAMIC: tl.constexpr): +def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr, H: tl.constexpr, + VB: tl.constexpr, VH: tl.constexpr, VS_ROW: tl.constexpr, VD: tl.constexpr, INT8_QK: tl.constexpr, + FP8_PV: tl.constexpr, V_TILE: tl.constexpr, P_DYNAMIC: tl.constexpr): tile, hz = tl.program_id(0), tl.program_id(1) nt: tl.constexpr = L // 64 rows = tile * 64 + tl.arange(0, 64) @@ -89,7 +92,7 @@ def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexp alpha = tl.exp2(m - new_m) den = den * alpha + tl.sum(p, 1) acc = acc * alpha[:, None] - v = tl.load(V + (hz * L + key_rows[:, None]) * D + cols[None, :]) + v = tl.load(V + (hz // H) * VB + (hz % H) * VH + key_rows[:, None] * VS_ROW + cols[None, :] * VD) if FP8_PV: if P_DYNAMIC: pscale = tl.maximum(tl.exp2(block_max - new_m) / 448.0, 1e-30) @@ -135,19 +138,26 @@ def sparse_sm89_attention(q: torch.Tensor, raise ValueError("Sparse INT8/FP8 attention requires a tile-64 mask and validity vector") from fastvideo_kernel.triton_kernels.index import map_to_index - q, k, v = q.contiguous(), k.contiguous(), v.contiguous() + # The production INT8-QK/BF16-PV route reads BSHD-backed views directly. + # Quantized Q/K and the output remain contiguous BHSD. Other ablations + # retain their established layout and arithmetic. + if not int8_qk or fp8_pv: + q, k, v = q.contiguous(), k.contiguous(), v.contiguous() vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous() grid = (triton.cdiv(length, 16), b * h) qi, ki, vf = q, k, v qs, ks, vs = q, k, v # unused pointers in BF16 ablations if int8_qk: # Tile pads are zero by contract; avoid a full FP32 copy for the reduction. - mean = k.sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1) - qi, ki = torch.empty_like(q, dtype=torch.int8), torch.empty_like(k, dtype=torch.int8) + # Preserve the exact reduction used by the old contiguous adapter; + # its temporary copy dies before Q/K quantization and fine attention. + mean = k.contiguous().sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1) + qi = torch.empty(q.shape, device=q.device, dtype=torch.int8) + ki = torch.empty(k.shape, device=k.device, dtype=torch.int8) qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32) ks = torch.empty_like(qs) - _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, CENTER=False, ROWS=16, num_warps=4) - _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, CENTER=True, ROWS=16, num_warps=4) + _quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, h, *q.stride(), CENTER=False, ROWS=16, num_warps=4) + _quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, h, *k.stride(), CENTER=True, ROWS=16, num_warps=4) if fp8_pv: vf = torch.empty_like(v, dtype=torch.float8_e4m3fn) if fp8_v_tiles: @@ -157,7 +167,7 @@ def sparse_sm89_attention(q: torch.Tensor, vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8) _quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4) index, count = map_to_index(mask.contiguous()) - out = torch.empty_like(q) + out = torch.empty(q.shape, device=q.device, dtype=q.dtype) _sparse_int8_fp8[(length // 64, b * h)](qi, ki, vf, @@ -170,6 +180,8 @@ def sparse_sm89_attention(q: torch.Tensor, out, length, dim, + h, + *vf.stride(), INT8_QK=int8_qk, FP8_PV=fp8_pv, V_TILE=fp8_v_tiles, diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index db82b76f60..1dddf21948 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -781,9 +781,14 @@ def forward( # type: ignore[override] # kernels' granularity. These entries take BHSD ([B, H, S_pad, D]); # mirror block_sparse_attn_256_bshd's Triton branch and transpose # around the call. - q_bhsd = query.transpose(1, 2).contiguous() - k_bhsd = key.transpose(1, 2).contiguous() - v_bhsd = value.transpose(1, 2).contiguous() + sm89_strided = (self._sm89_kernel == "int8" and not torch.is_grad_enabled() and not compiling + and query.dtype == torch.bfloat16 and query.shape[-1] == 128 + and torch.cuda.get_device_capability(query.device) == (8, 9)) + q_bhsd = query.transpose(1, 2) + k_bhsd = key.transpose(1, 2) + v_bhsd = value.transpose(1, 2) + if not sm89_strided: + q_bhsd, k_bhsd, v_bhsd = (t.contiguous() for t in (q_bhsd, k_bhsd, v_bhsd)) sm100a_mask = mask sm100a_variable_block_sizes = attn_metadata.variable_block_sizes @@ -880,9 +885,11 @@ def forward( # type: ignore[override] ) else: if has_sm100a_pair: - q_bhsd = q_bhsd[:, :, :logical_seq_len].contiguous() - k_bhsd = k_bhsd[:, :, :logical_seq_len].contiguous() - v_bhsd = v_bhsd[:, :, :logical_seq_len].contiguous() + q_bhsd = q_bhsd[:, :, :logical_seq_len] + k_bhsd = k_bhsd[:, :, :logical_seq_len] + v_bhsd = v_bhsd[:, :, :logical_seq_len] + if not sm89_strided: + q_bhsd, k_bhsd, v_bhsd = (t.contiguous() for t in (q_bhsd, k_bhsd, v_bhsd)) if (self._sm89_kernel != "original" and not torch.is_grad_enabled() and not compiling and q_bhsd.dtype == torch.bfloat16 and q_bhsd.shape[-1] == 128 and torch.cuda.get_device_capability(q_bhsd.device) == (8, 9)): diff --git a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py index 122aa08f11..431c5279d4 100644 --- a/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py +++ b/fastvideo/tests/attention/test_minimax_h3_sparse_int8.py @@ -70,3 +70,44 @@ def test_fp8_dynamic_probability_scale_preserves_small_blocks(): assert reference.abs().min() > 0.1 assert torch.count_nonzero(fixed) == 0 torch.testing.assert_close(dynamic.float(), reference, rtol=0.02, atol=0.02) + + +@pytest.mark.parametrize("batch,heads", [(1, 2), (2, 3)]) +@pytest.mark.parametrize("partner_pad", [False, True]) +def test_int8_bshd_views_match_contiguous_and_reduce_peak(batch, heads, partner_pad): + """Read production BSHD views without retaining three BHSD copies.""" + _cuda_sm89() + from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention + + torch.manual_seed(113) + length, dim = 1024, 128 + storage_length = length + (64 if partner_pad else 0) + tensors = [torch.randn(batch, storage_length, heads, dim, device="cuda", dtype=torch.bfloat16) + for _ in range(3)] + q, k, v = [tensor[:, :length].transpose(1, 2) for tensor in tensors] + vbs = torch.full((length // 64,), 64, device="cuda", dtype=torch.int32) + vbs[1], vbs[4] = 7, 31 + valid = torch.arange(length, device="cuda") % 64 < vbs.repeat_interleave(64) + k[:, :, ~valid] = 0 + v[:, :, ~valid] = 0 + mask = torch.rand(batch, heads, length // 64, length // 64, device="cuda") > 0.8 + mask[:, :, 0] = False + with torch.inference_mode(): + # Populate autotuning/compilation caches before measuring allocations. + warm = sparse_sm89_attention(q, k, v, mask, vbs) + del warm + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + reference = sparse_sm89_attention(q.contiguous(), k.contiguous(), v.contiguous(), mask, vbs) + torch.cuda.synchronize() + copy_peak = torch.cuda.max_memory_allocated() - baseline + expected = reference.cpu() + del reference + torch.cuda.reset_peak_memory_stats() + actual = sparse_sm89_attention(q, k, v, mask, vbs) + torch.cuda.synchronize() + view_peak = torch.cuda.max_memory_allocated() - baseline + torch.testing.assert_close(actual.cpu(), expected, rtol=0, atol=0) + tensor_bytes = batch * heads * length * dim * 2 + assert copy_peak - view_peak >= tensor_bytes, (copy_peak, view_peak) From 3c0668f6c7182c335eda410f057320c8a5cf5924 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:03:45 -0700 Subject: [PATCH 19/24] [perf]: fuse H3 INT8 VAE scaling and bias without intermediates --- .../models/vaes/minimax_h3_int8_convrot.py | 5 ++ .../models/vaes/minimax_h3_int8_kernels.py | 46 +++++++++++++++++++ .../tests/vaes/test_minimax_h3_int8_shared.py | 20 ++++++++ 3 files changed, 71 insertions(+) create mode 100644 fastvideo/models/vaes/minimax_h3_int8_kernels.py diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py index c4f366dd7d..86c0b0d2d7 100644 --- a/fastvideo/models/vaes/minimax_h3_int8_convrot.py +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -105,6 +105,7 @@ def __init__( self.convrot = convrot self.group_size = group_size self._transpose_view = os.environ.get("FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW", "0") == "1" + self._fused_dequant = os.environ.get("FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT", "0") == "1" self.register_buffer("weight", torch.empty(out_features, in_features, dtype=torch.int8)) self.register_buffer("weight_scale", torch.empty(out_features, 1, dtype=torch.float32)) if bias: @@ -145,6 +146,10 @@ def forward_quantized(self, x_q: torch.Tensor, x_scale: torch.Tensor, if not self._transpose_view: weight = weight.contiguous() acc = torch._int_mm(x_q, weight)[:rows] + if self._fused_dequant and not torch.is_grad_enabled() and not torch.compiler.is_compiling(): + from fastvideo.models.vaes.minimax_h3_int8_kernels import fused_int8_dequant_bias + return fused_int8_dequant_bias(acc, x_scale[:rows], self.weight_scale, self.bias, dtype).view( + *original_shape[:-1], self.out_features) out = self._dequant_int8_gemm(acc, x_scale[:rows], self.weight_scale) if self.bias is not None: out = out + self.bias.float() diff --git a/fastvideo/models/vaes/minimax_h3_int8_kernels.py b/fastvideo/models/vaes/minimax_h3_int8_kernels.py new file mode 100644 index 0000000000..e1e5a9fad4 --- /dev/null +++ b/fastvideo/models/vaes/minimax_h3_int8_kernels.py @@ -0,0 +1,46 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Eager INT8 VAE epilogue with the reference's separate FP32 operations.""" +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _dequant_bias(Acc, XScale, WScale, Bias, Out, M: tl.constexpr, N: tl.constexpr, + XS: tl.constexpr, WS: tl.constexpr, HAS_BIAS: tl.constexpr, BLOCK: tl.constexpr): + offsets = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + rows, cols = offsets // N, offsets % N + valid = rows < M + acc = tl.load(Acc + offsets, valid, 0).to(tl.float32) + x_scale = tl.load(XScale + rows * XS, valid, 0).to(tl.float32) + weight_scale = tl.load(WScale + cols * WS).to(tl.float32) + out = acc * x_scale + out = out * weight_scale + if HAS_BIAS: + out = out + tl.load(Bias + cols).to(tl.float32) + tl.store(Out + offsets, out, valid) + + +def fused_int8_dequant_bias(acc: torch.Tensor, x_scale: torch.Tensor, + weight_scale: torch.Tensor, bias: torch.Tensor | None, + dtype: torch.dtype) -> torch.Tensor: + """Avoid full-size FP32 scaling intermediates; retain both rounding steps.""" + rows, cols = acc.shape + if not acc.is_cuda or acc.dtype != torch.int32 or not acc.is_contiguous(): + raise ValueError("INT8 VAE epilogue requires a contiguous CUDA INT32 matrix") + if x_scale.shape != (rows, 1) or weight_scale.shape != (cols, 1): + raise ValueError("INT8 VAE epilogue requires per-row and per-output-channel scales") + if any(t.device != acc.device for t in (x_scale, weight_scale)): + raise ValueError("INT8 VAE epilogue scales must be on the accumulator device") + if bias is not None and (bias.device != acc.device or bias.shape != (cols,) or not bias.is_contiguous()): + raise ValueError("INT8 VAE epilogue bias must be contiguous on the accumulator device") + if dtype not in (torch.float32, torch.float16, torch.bfloat16): + raise ValueError("INT8 VAE epilogue supports FP32, FP16 and BF16 outputs") + out = torch.empty((rows, cols), device=acc.device, dtype=dtype) + _dequant_bias[(triton.cdiv(rows * cols, 1024),)]( + acc, x_scale, weight_scale, bias if bias is not None else acc, out, + rows, cols, x_scale.stride(0), weight_scale.stride(0), bias is not None, + BLOCK=1024, num_warps=4, enable_fp_fusion=False) + return out diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py index 88bb79e23b..3874ab3eec 100644 --- a/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py @@ -32,6 +32,7 @@ def test_shared_int8_and_transpose_views_are_exact(rows, dtype, convrot, monkeyp assert quant.call_count == 1 for layer in layers: layer._transpose_view = True + layer._fused_dequant = True views = shared_int8_projections(layers, x) for ref, actual, view in zip(expected, shared, views, strict=True): assert torch.isfinite(ref).all() @@ -69,5 +70,24 @@ def test_vae_attention_shares_quantized_projections_exactly(distributed_setup, m attention._share_int8_qkv = True for layer in (attention.to_q, attention.to_k, attention.to_v): layer._transpose_view = True + layer._fused_dequant = True actual = attention(x) torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for fused INT8 epilogue") +@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("has_bias", [False, True]) +def test_fused_int8_epilogue_large_accumulators_and_small_scales(dtype, has_bias): + from fastvideo.models.vaes.minimax_h3_int8_kernels import fused_int8_dequant_bias + + torch.manual_seed(19) + acc = torch.randint(-400_000_000, 400_000_000, (129, 264), device="cuda", dtype=torch.int32) + x_scale = torch.logspace(-30, -3, 129, device="cuda").view(-1, 1) + w_scale = torch.logspace(-6, -2, 264, device="cuda").view(-1, 1) + bias = torch.randn(264, device="cuda") if has_bias else None + expected = acc.float() * x_scale.float() * w_scale.t().float() + if bias is not None: + expected = expected + bias.float() + actual = fused_int8_dequant_bias(acc, x_scale, w_scale, bias, dtype) + torch.testing.assert_close(actual, expected.to(dtype), rtol=0, atol=0) From a7f7ede7a7fba0f099858131b6a70ae954ceef1d Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:08:48 -0700 Subject: [PATCH 20/24] [bugfix]: retain consumer attention environment parsing after core rebase --- fastvideo/attention/backends/video_sparse_attn_h3.py | 1 + 1 file changed, 1 insertion(+) diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index 1dddf21948..a4ea213206 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -61,6 +61,7 @@ import functools import math +import os from dataclasses import dataclass from typing import Any From 87b22a5d8890c257595728cf2c353bfd24121d5b Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:13:52 -0700 Subject: [PATCH 21/24] [test]: adapt consumer VSA validation to packed segment metadata --- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 3 ++- .../transformers/test_minimax_h3_tile_first.py | 15 ++++++++++++--- 2 files changed, 14 insertions(+), 4 deletions(-) diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index b3ab154524..ee13941824 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -183,7 +183,8 @@ def vsa_tile_first_attention(attn: Any, hidden_states: torch.Tensor, k_pooled = _pool_tiles(key, meta.variable_block_sizes, meta.tile_elems) scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / dim**0.5 sparsity = 0.0 if attn._layer_idx in meta.dense_layers else meta.VSA_sparsity - mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, sparsity, meta.exempt) + mask = _build_block_mask(scores, meta.num_prefix_tiles, sparsity, meta.exempt, + meta.video_tile_spans, meta.span_sparsities) # Two heads keep the artifact small while retaining all real keys, # query rows, per-tile selections and partial-tile validity. torch.save({"q": query[:, :, :2].transpose(1, 2).contiguous().cpu(), diff --git a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py index 9ebbf8181d..752a709ce6 100644 --- a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py +++ b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py @@ -52,7 +52,7 @@ def test_shared_fp8_projections_match_independent_quantization(granularity): @pytest.mark.parametrize("fp8", [False, True]) @pytest.mark.parametrize("gate_active", [False, True]) @pytest.mark.parametrize("fused_rope", [False, True]) -def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distributed_setup, +def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distributed_setup, tmp_path, fp8, gate_active, fused_rope, kernel): if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported(): pytest.skip("BF16 CUDA is required") @@ -68,6 +68,9 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distribu monkeypatch.setenv("FASTVIDEO_H3_VSA_FP4", "0") monkeypatch.setenv("FASTVIDEO_H3_VSA_TILE_FIRST", "0") monkeypatch.setenv("FASTVIDEO_H3_VSA_SM89_KERNEL", "original") + capture = kernel == "int8" and fp8 and gate_active and fused_rope + if capture: + monkeypatch.setenv("FASTVIDEO_H3_CAPTURE_QKV", str(tmp_path)) torch.manual_seed(21) attn = MiniMaxH3Attention(256, 2, 128, 1e-5, (AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3,), FP8Config("channel") if fp8 else None, "transformer_blocks.0.attn", fuse_qknorm_rope=fused_rope) @@ -79,8 +82,9 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distribu if fp8: for layer in (attn.to_q, attn.to_k, attn.to_v, attn.to_out): _install_fp8_buffers(layer) - meta = MiniMaxH3VSAMetadataBuilder().build(999, (4, 6, 10), (1, 1, 1), 0.8, - (65, 97), torch.device("cuda"), tile_size=64) + meta = MiniMaxH3VSAMetadataBuilder().build(current_timestep=999, patch_size=(1, 1, 1), + VSA_sparsity=0.8, packed_segments=(65, 97, (4, 6, 10)), + device=torch.device("cuda"), tile_size=64) length = meta.total_seq_length x = torch.randn(1, length, 256, device="cuda", dtype=torch.bfloat16) angles = torch.randn(length, 96, device="cuda") @@ -95,3 +99,8 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distribu error = (actual.float() - reference.float()).norm() / reference.float().norm() assert error < (0.02 if fp8 else 0.005), float(error) torch.testing.assert_close(actual, reference, rtol=0.03, atol=0.05) + if capture: + data = torch.load(tmp_path / "layer-0.pt", weights_only=True) + torch.testing.assert_close(data["vbs"], meta.variable_block_sizes.cpu(), rtol=0, atol=0) + assert data["q"].shape == (1, 2, meta.variable_block_sizes.numel() * 64, 128) + assert data["mask"].dtype == torch.bool From fb92af176a405a47bcff203404951d86e72eed26 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:13:52 -0700 Subject: [PATCH 22/24] [bench]: sample total GPU memory for consumer budget checks --- scripts/benchmarks/minimax_h3_4090/bench_pod.py | 17 +++++++++++++++++ scripts/benchmarks/minimax_h3_4090/summarize.py | 2 ++ 2 files changed, 19 insertions(+) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index dc4349abad..a905e1989b 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -22,7 +22,18 @@ def __init__(self): self.stop = threading.Event() self.peak_bytes = 0 self.peak_anon_bytes = 0 + self.peak_gpu_bytes = None + self._gpu_used = None + self._nvml_shutdown = None self.thread = threading.Thread(target=self._sample, daemon=True) + try: + import pynvml + pynvml.nvmlInit() + self._nvml_shutdown = pynvml.nvmlShutdown + handle = pynvml.nvmlDeviceGetHandleByIndex(0) + self._gpu_used = lambda: pynvml.nvmlDeviceGetMemoryInfo(handle).used + except Exception as exc: + print(f"GPU memory sampler unavailable: {exc}", flush=True) def _sample(self): while not self.stop.is_set(): @@ -31,6 +42,8 @@ def _sample(self): self.peak_bytes = max(self.peak_bytes, int((root / "memory.current").read_text())) stats = dict(line.split() for line in (root / "memory.stat").read_text().splitlines()) self.peak_anon_bytes = max(self.peak_anon_bytes, int(stats["anon"])) + if self._gpu_used is not None: + self.peak_gpu_bytes = max(self.peak_gpu_bytes or 0, self._gpu_used()) except (OSError, KeyError, ValueError): return self.stop.wait(0.1) @@ -42,6 +55,8 @@ def __enter__(self): def __exit__(self, *_args): self.stop.set() self.thread.join() + if self._nvml_shutdown is not None: + self._nvml_shutdown() def main(): @@ -147,6 +162,8 @@ def main(): wall = round(time.perf_counter() - t, 2) results["runs"].append({"prompt": pid, "warmup": i < a.warmup, "wall_s": wall, "clip": request["output"]["output_path"], + "peak_gpu_used_gib": (round(host_peak.peak_gpu_bytes / 2**30, 3) + if host_peak.peak_gpu_bytes is not None else None), "peak_host_cgroup_gib": round(host_peak.peak_bytes / 2**30, 3), "peak_host_anon_gib": round(host_peak.peak_anon_bytes / 2**30, 3)}) timed = [run["wall_s"] for run in results["runs"] if not run["warmup"]] diff --git a/scripts/benchmarks/minimax_h3_4090/summarize.py b/scripts/benchmarks/minimax_h3_4090/summarize.py index cffaa7988e..fffecde50e 100644 --- a/scripts/benchmarks/minimax_h3_4090/summarize.py +++ b/scripts/benchmarks/minimax_h3_4090/summarize.py @@ -48,6 +48,8 @@ def main(): "median_stage_s": {name: statistics.median(run["stage_s"][name] for run in timed) for name in ("conditioning", "denoising", "video_decoding", "audio_decoding")}, "peak_gpu_allocated_gib": max(run["peak_gpu_allocated_gib"] for run in timed), + "peak_gpu_used_gib": max((run["peak_gpu_used_gib"] for run in timed + if run.get("peak_gpu_used_gib") is not None), default=None), "peak_host_anon_gib": max(run["peak_host_anon_gib"] for run in timed), "peak_host_cgroup_gib": max(run["peak_host_cgroup_gib"] for run in timed), "notes": "Stage times include deferred component loading. Host peaks are pod-wide samples every 100 ms.", From 5ac85ada47f382779b66e1a9812acc7617e1d647 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:32:01 -0700 Subject: [PATCH 23/24] [docs]: record 42-second 4090 clips and consumer memory validation --- scripts/benchmarks/minimax_h3_4090/README.md | 101 +++++++++++++++++-- 1 file changed, 92 insertions(+), 9 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 905f70d0e6..529e97be61 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -260,16 +260,13 @@ values with FP8 while retaining the other tensors projects about 25.3 GiB before activations (the FP8 block-scale overhead is small). This is a storage estimate, not a measured FP8 encoder. First try fused NVFP4 dequantization and avoid per-linear GPU scalar synchronization; then compare -a native sm89 FP8 encoder at equal prompts. Conditioning is only about -11.2 s of the current 111.3 s clip, so encoder work alone cannot dominate the -end-to-end gain. +a native sm89 FP8 encoder at equal prompts. The later streamed/fused conditioning stage is about 0.55 s, so a new +encoder export must be measured against that implementation. -Remaining speed experiments: INT8-QK/BF16-PV same-seed clips; VAE compilation -and tile-batch tuning; more resident blocks after streaming the encoder; -fused FP8 GEMM epilogues and norm/activation quantization. The 16 GiB cap -still needs encoder streaming and a completed memory-capped run. The cached -recipe's 42 GiB anonymous host peak does not establish a 32 GB system-RAM -minimum. +Remaining speed experiments include VAE compilation and tile-batch tuning, +INT8 decoder epilogue fusion, direct strided fine-attention reads, and fused +norm/activation quantization. See the completed smaller-VRAM results below. +The cached recipe's host peak does not establish a 32 GB system-RAM minimum. ## Streamed encoder and smaller VRAM caps @@ -328,3 +325,89 @@ The 8 GiB cap at `fcdba37fc` completed its warmup but OOMed on the timed harbor prompt in fine attention. Do not report it as supported. A 34-resident experiment completed denoising but OOMed during VAE INT8 epilogue allocation. Both failures motivate subsequent memory work rather than speed claims. + + +## Consumer kernel memory work on the release core + +The consumer branch is rebased onto release core `a97d23f09` (fork PR #45). +Historical measured commits remain reachable through tag +`h3-consumer-fp8-measured-20261003`; rebase changes their branch commit IDs. + +`7a0d7d33b` lets the INT8-QK/BF16-PV fine kernel read BSHD-backed views +without retaining three full BF16 layout copies. Q/K quantization writes +contiguous INT8 arrays and the output stays BHSD; the K-mean reduction keeps +the reference's contiguous reduction order. Four CUDA tests require exact +output equality for partial tiles, empty selections, multiple batches and +partner padding, together with a lower peak allocation. + +`3c0668f6c` adds the opt-in eager +`FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT=1`. One Triton pass applies the INT32 +GEMM's row scale, output-channel scale and optional bias, then casts the +result. FP32 operations retain separate rounding steps (FP fusion disabled). +Strict tests cover zero rows, small input batches, FP32/FP16/BF16, bias, +large INT32 accumulators and tiny scales. Compiled and grad paths retain +the reference implementation. Combine it with shared QKV and transpose +views using `FASTVIDEO_H3_VAE_INT8_SHARED_QKV=1` and +`FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW=1`. + +After the core rebase, 120 focused tests passed (one unrelated GPU cudagraph +check excluded) and pre-commit passed. `87b22a5d8` adapts capture and tests +to packed-segment metadata while retaining the core's calibrated NVFP4 +activation-scale guard. `fb92af176` adds optional NVML total-device-memory +samples every 100 ms, alongside host samples. Summaries distinguish sampled +total GPU usage from PyTorch's allocated peak. Allocator caps omit driver +and external CUDA memory, so the 8 GiB total-budget experiment uses a +7.25 GiB allocator cap and must also satisfy the observed NVML budget. +These are 4090 simulations; real lower-VRAM and 30-series performance still +needs those devices. + + +## Warmed release-core clip results + +At `fb92af176`, after one warmup and two timed requests: + +| Clip | Median e2e | Timed requests | Conditioning | Denoise | Video decode | Peak allocated | Sampled total GPU | Peak host anon | +| --- | --- | --- | --- | --- | --- | --- | --- | --- | +| 832×480, 124 frames / 5.167 s | **41.75 s** | 42.16 / 41.33 s | 0.55 s | 32.69 s | 6.63 s | 22.01 GiB | 23.983 GiB | 24.78 GiB | +| 832×480, 243 frames / 10.125 s | **79.67 s** | 79.64 / 79.70 s | 0.55 s | 63.51 s | 13.23 s | 20.29 GiB | 23.985 GiB | 27.02 GiB | + +Warmups took 84.40 / 125.16 s. These wall times include audio, frame export +and MP4 saving, and exclude initial generator construction. The 5 s recipe +keeps 34 DiT blocks resident; the 10 s recipe keeps 30. Both use the same +checkpoint revision and eight DMD forwards as the earlier rows. No profiling, +QKV capture, alternate decoder or sparsity increase is enabled. + +```bash +source /workspace/env.sh +source /workspace/venv/bin/activate +cd /workspace +export PYTHONPATH="/workspace/fastvideo-core:${PYTHONPATH:-}" +export FASTVIDEO_SOURCE_COMMIT=fb92af176 +export FASTVIDEO_H3_PARK_MODULES=vae,audio_vae +export FASTVIDEO_H3_ENCODER_LAYERWISE=1 FASTVIDEO_H3_ENCODER_FUSED_DEQUANT=1 +export FASTVIDEO_H3_VSA_TILE_FIRST=1 FASTVIDEO_H3_VSA_SM89_KERNEL=int8 +export FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=34 FASTVIDEO_H3_FFN_CHUNK_TOKENS=16384 +export FASTVIDEO_H3_VAE_TILE_BATCH=28 +export FASTVIDEO_H3_VAE_INT8_SHARED_QKV=1 FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW=1 +export FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT=1 MAX_JOBS=4 +unset FASTVIDEO_CUDA_MEMORY_CAP_GIB +python -P /workspace/fastvideo-core/scripts/benchmarks/minimax_h3_4090/bench_pod.py \ + sm89-int8-fast2-480p-5s /workspace/vol/pruned_fp8_300 fp8 \ + --offload-buffers --no-vae-compile --height 480 --width 832 --frames 124 --timed 2 +``` + +For 10 s, set resident blocks to 30, name to `sm89-int8-fast2-480p-10s`, +and frames to 243. Both decoded video and PCM audio hash-identically to the +previous INT8 candidate for ceramics and harbor at 10 s. This validates these +memory/decode changes on those prompts, while the INT8 attention candidate +still differs from the original BF16 attention clips and needs full quality +review. The sampled 5 s contact sheet is coherent. Raw clips, hashes and +results live in `output/fasth3-4090-20261003/` beside the worktree. + +The earlier unprofiled 1344×768, 243-frame run at `fcdba37fc` completed at +279.94 s median (289.06 / 270.82 s, 306.80 s warmup), with 12 resident blocks +and shared VAE QKV/transpose views, before direct-layout attention and fused +VAE epilogues. Its stage medians were 0.55 s conditioning, 227.21 s denoise, +44.82 s video decode and 0.48 s audio. Peak allocation was 21.19 GiB and host +anonymous memory 39.81 GiB. The updated 768p and 8 GiB total-budget runs are +in progress; they must complete before claiming their speed or support. From 98cdc7d3975c76b0de41b27583f0e0ada52cbf19 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 13:42:11 -0700 Subject: [PATCH 24/24] [docs]: record Track B PR and strict 8 GiB budget limit --- scripts/benchmarks/minimax_h3_4090/README.md | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 529e97be61..11c8b75e8a 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -358,6 +358,9 @@ samples every 100 ms, alongside host samples. Summaries distinguish sampled total GPU usage from PyTorch's allocated peak. Allocator caps omit driver and external CUDA memory, so the 8 GiB total-budget experiment uses a 7.25 GiB allocator cap and must also satisfy the observed NVML budget. +That trial completed its warmup and first timed request, but sampled total +GPU usage reached about 8.28 GiB. It therefore does **not** meet a strict +8 GiB device target. A tighter allocator cap still needs validation. These are 4090 simulations; real lower-VRAM and 30-series performance still needs those devices. @@ -409,5 +412,11 @@ The earlier unprofiled 1344×768, 243-frame run at `fcdba37fc` completed at and shared VAE QKV/transpose views, before direct-layout attention and fused VAE epilogues. Its stage medians were 0.55 s conditioning, 227.21 s denoise, 44.82 s video decode and 0.48 s audio. Peak allocation was 21.19 GiB and host -anonymous memory 39.81 GiB. The updated 768p and 8 GiB total-budget runs are -in progress; they must complete before claiming their speed or support. +anonymous memory 39.81 GiB. The updated 768p run was queued after the +7.25 GiB allocator-cap trial. SSH became unreachable before the final +results could be collected; updated 768p speed and strict 8 GiB support +remain unverified. + +Track B is staged in [draft PR #46](https://github.com/aryan5v/FastVideo/pull/46), +stacked on the shared release core in #45. Historical timing sources are +preserved by the `h3-consumer-fp8-measured-20261003` tag.