From 2d0ebf1eefa2d09fb668aff658382690af75ef1e Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 14:25:03 -0700 Subject: [PATCH 001/107] [feat]: load packed FastH3 NVFP4 and Comfy int8 VAE on Blackwell Keep encoder/DiT/VAE off disk between clips, overlay the Comfy int8-convrot decoder for quality, and let the playground queue prompts while switching clips. TAEH3 stays an optional preview config only. (cherry picked from commit 15e5517169b47908463f674363e2bd2078f7e171) --- .../basic/basic_fasth3_5090_nvfp4.yaml | 57 ++++ .../serving/openai_fasth3_5090_nvfp4.yaml | 58 ++++ .../openai_fasth3_5090_nvfp4_taeh3.yaml | 57 ++++ .../serving/openai_fasth3_pro6000_nvfp4.yaml | 57 ++++ fastvideo/api/compat.py | 12 +- fastvideo/api/schema.py | 3 + .../entrypoints/openai/static/playground.css | 3 + .../entrypoints/openai/static/playground.html | 10 +- .../entrypoints/openai/static/playground.js | 261 +++++++++------ fastvideo/layers/quantization/nvfp4_config.py | 159 ++++++++- fastvideo/models/loader/component_loader.py | 21 +- fastvideo/models/loader/fsdp_load.py | 60 +++- .../models/vaes/minimax_h3_int8_convrot.py | 315 ++++++++++++++++++ .../basic/minimax_h3/minimax_h3_pipeline.py | 88 ++++- fastvideo/tests/api/test_typed_quant_flow.py | 15 + .../entrypoints/test_openai_video_client.py | 1 + .../ops/quantization/test_nvfp4_config.py | 8 + .../quantization/test_nvfp4_h3_dit_export.py | 88 +++++ .../test_nvfp4_minimax_h3_wiring.py | 29 ++ .../test_minimax_h3_sequential_start.py | 157 ++++++++- .../vaes/test_minimax_h3_int8_convrot.py | 174 ++++++++++ 21 files changed, 1496 insertions(+), 137 deletions(-) create mode 100644 examples/inference/basic/basic_fasth3_5090_nvfp4.yaml create mode 100644 examples/serving/openai_fasth3_5090_nvfp4.yaml create mode 100644 examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml create mode 100644 examples/serving/openai_fasth3_pro6000_nvfp4.yaml create mode 100644 fastvideo/models/vaes/minimax_h3_int8_convrot.py create mode 100644 fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py create mode 100644 fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py diff --git a/examples/inference/basic/basic_fasth3_5090_nvfp4.yaml b/examples/inference/basic/basic_fasth3_5090_nvfp4.yaml new file mode 100644 index 0000000000..068cd79c89 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_5090_nvfp4.yaml @@ -0,0 +1,57 @@ +# FastH3 42-block NVFP4 DiT on one RTX 5090 (sm_120), 832x480 4-step. +# +# Requires the packed ``transformer/nvfp4_weights.safetensors`` export next to +# the dense shards, and ``engine.quantization.layer_profile: h3_dit`` so the +# loader overlays those buffers instead of materializing the 32 GB GEMMs. +# Full decode uses the Comfy int8-convrot VAE when +# ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` sits beside the dense +# shards. Live preview is ``video_decode_backend: taeh3`` on the serve config. +# +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_5090_nvfp4.yaml +generator: + model_path: /workspace/models/fasth3-42-live + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: true + vae: true + pin_cpu_memory: true + lazy_module_load: true + compile: + enabled: false + vae_enabled: false + pipeline: + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: true + inference_torch_compile: false + vae_parallel_decode: false +request: + prompt: >- + A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain + peaks above a blue valley filled with thin morning mist. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_5090_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/serving/openai_fasth3_5090_nvfp4.yaml b/examples/serving/openai_fasth3_5090_nvfp4.yaml new file mode 100644 index 0000000000..63339f9290 --- /dev/null +++ b/examples/serving/openai_fasth3_5090_nvfp4.yaml @@ -0,0 +1,58 @@ +# OpenAI-compatible FastH3 server for one RTX 5090 (sm_120). +# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE (quality decode). +# lazy_module_load is off so DiT/VAE stay resident; sequential load keeps the +# encoder in host RAM. TAEH3 live preview: openai_fasth3_5090_nvfp4_taeh3.yaml +# +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a \ +# fastvideo serve --config examples/serving/openai_fasth3_5090_nvfp4.yaml +generator: + model_path: /workspace/models/fasth3-42-live + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: true + image_encoder: true + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: true + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae + +server: + host: 127.0.0.1 + port: 8000 + output_dir: outputs/openai_fasth3_5090_nvfp4 + served_model_name: fasth3 + +default_request: + negative_prompt: "" + sampling: + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + seed: 2026 + output: + return_frames: false diff --git a/examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml b/examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml new file mode 100644 index 0000000000..cb9cea4228 --- /dev/null +++ b/examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml @@ -0,0 +1,57 @@ +# FastH3 5090 live preview: packed NVFP4 DiT/encoder + TAEH3 decode. +# Quality path is openai_fasth3_5090_nvfp4.yaml (Comfy int8 VAE). This one +# skips the 10 GB video VAE so clip time can approach DiT + tiny decode. +# +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a \ +# fastvideo serve --config examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml +generator: + model_path: /workspace/models/fasth3-42-live + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: true + image_encoder: true + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: true + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: taeh3 + +server: + host: 127.0.0.1 + port: 8000 + output_dir: outputs/openai_fasth3_5090_nvfp4_taeh3 + served_model_name: fasth3 + +default_request: + negative_prompt: "" + sampling: + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + seed: 2026 + output: + return_frames: false diff --git a/examples/serving/openai_fasth3_pro6000_nvfp4.yaml b/examples/serving/openai_fasth3_pro6000_nvfp4.yaml new file mode 100644 index 0000000000..698c083a7c --- /dev/null +++ b/examples/serving/openai_fasth3_pro6000_nvfp4.yaml @@ -0,0 +1,57 @@ +# OpenAI-compatible FastH3 server for one RTX PRO 6000 Blackwell (sm_120, 96 GB). +# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. 96 GB holds encoder, +# DiT, and VAE together, so sequential ping-pong stays off. +# +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FLASHINFER_CUDA_ARCH_LIST=12.0a \ +# fastvideo serve --config examples/serving/openai_fasth3_pro6000_nvfp4.yaml +generator: + model_path: /workspace/models/fasth3-42-live + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: true + vae: false + pin_cpu_memory: true + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae + +server: + host: 127.0.0.1 + port: 8000 + output_dir: outputs/openai_fasth3_pro6000_nvfp4 + served_model_name: fasth3 + +default_request: + negative_prompt: "" + sampling: + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + seed: 2026 + output: + return_frames: false diff --git a/fastvideo/api/compat.py b/fastvideo/api/compat.py index 07d4762cbf..32c13ccedb 100644 --- a/fastvideo/api/compat.py +++ b/fastvideo/api/compat.py @@ -292,8 +292,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An # typed surface accepts a string and does the wiring here so # downstream code can rely on a single source of truth. from fastvideo.layers.quantization import get_quantization_config + from fastvideo.layers.quantization.nvfp4_config import NVFP4Config _resolved_quant_cls = get_quantization_config(quantization.transformer_quant) - kwargs["transformer_quant"] = _resolved_quant_cls() + quant_instance = _resolved_quant_cls() + if quantization.layer_profile is not None: + if not isinstance(quant_instance, NVFP4Config): + raise ValueError("engine.quantization.layer_profile is only valid with transformer_quant NVFP4, " + f"got {type(quant_instance).__name__}") + quant_instance = NVFP4Config( + layer_profile=quantization.layer_profile, + retain_original_weights=quant_instance.retain_original_weights, + ) + kwargs["transformer_quant"] = quant_instance components = normalized.pipeline.components if components.pipeline_config_path is not None: diff --git a/fastvideo/api/schema.py b/fastvideo/api/schema.py index 77224d0368..2fda230ebb 100644 --- a/fastvideo/api/schema.py +++ b/fastvideo/api/schema.py @@ -76,6 +76,9 @@ class CompileConfig: class QuantizationConfig: text_encoder_quant: str | None = None transformer_quant: str | None = None + # Forwarded onto ``NVFP4Config`` when ``transformer_quant`` is ``NVFP4``. + # ``h3_dit`` selects the packed MiniMax-H3 attention+FFN export. + layer_profile: str | None = None @dataclass diff --git a/fastvideo/entrypoints/openai/static/playground.css b/fastvideo/entrypoints/openai/static/playground.css index cb7a0ee09b..9119d18252 100644 --- a/fastvideo/entrypoints/openai/static/playground.css +++ b/fastvideo/entrypoints/openai/static/playground.css @@ -81,6 +81,9 @@ pre { padding: 16px; font-size: 12px; background: var(--inset); border-radius: 8 video { width: 100%; max-height: 420px; } #job-status, .error { margin-block-start: 16px; font-size: 14px; } .error { color: var(--error); overflow-wrap: anywhere; } +.clip-nav { display: flex; flex-wrap: wrap; gap: 12px; margin-block-start: 16px; } +.clip-nav[hidden] { display: none !important; } +.clip-nav button { flex: 1 1 140px; } .result-actions { display: flex; flex-wrap: wrap; gap: 16px; align-items: center; margin-block-start: 16px; font-size: 14px; } .job-id { color: var(--muted); font: 12px/1.5 ui-monospace, monospace; overflow-wrap: anywhere; margin-block-start: 12px; } .history { border-block-start: 1px solid var(--line); margin-block-start: 28px; padding-block-start: 20px; } diff --git a/fastvideo/entrypoints/openai/static/playground.html b/fastvideo/entrypoints/openai/static/playground.html index 7dace42f86..66264428eb 100644 --- a/fastvideo/entrypoints/openai/static/playground.html +++ b/fastvideo/entrypoints/openai/static/playground.html @@ -18,7 +18,7 @@

Text to video + audio

A new prompt. The same running server.

-

Generate, watch, and refine. This page and your API clients share one running FastVideo server.

+

Generate, watch, and queue the next prompt. Clips stay in Recent jobs so you can switch between them while the server keeps working.

Connecting to server… @@ -37,7 +37,7 @@

Create a video

Resolution and sampling come from the server configuration.

-

Closing this page does not cancel a job.

+

Queue the next prompt while a clip plays. Closing this page does not cancel a job.

Use this prompt with cURL @@ -55,6 +55,10 @@

Create a video

No generation starts until you submit a prompt.

+
@@ -63,7 +67,7 @@

Create a video

Recent jobs

-

Shared with other clients. Job history clears when the server restarts.

+

Select any clip to play it. The prompt box is left alone so you can queue the next one. History clears when the server restarts.

    Loading recent jobs…

    diff --git a/fastvideo/entrypoints/openai/static/playground.js b/fastvideo/entrypoints/openai/static/playground.js index eeba66ddb1..2479aa957e 100644 --- a/fastvideo/entrypoints/openai/static/playground.js +++ b/fastvideo/entrypoints/openai/static/playground.js @@ -6,9 +6,11 @@ const active = (job) => ["queued", "in_progress"].includes(job.status); let model = ""; let readyLabel = "Server ready"; - let busy = false; - let currentJob = null; + let watching = null; + let jobs = []; + let pollTimer = 0; let pollVersion = 0; + let submitting = false; const api = async (path, options = {}) => { const response = await fetch(path, { ...options, signal: AbortSignal.timeout(15000) }); @@ -16,10 +18,19 @@ if (!response.ok) throw new Error(body.error?.message || `Request failed (HTTP ${response.status}).`); return body; }; - const setBusy = (value) => { - busy = value; - $("generate").disabled = value || !model; - $("generate").textContent = value ? "Waiting for this job…" : "Generate video ↗"; + const jobPath = (id) => `/v1/videos/${encodeURIComponent(id)}`; + const syncGenerateButton = () => { + $("generate").disabled = !model || submitting; + $("generate").textContent = jobs.some(active) ? "Queue next clip ↗" : "Generate video ↗"; + }; + const watchingIndex = () => jobs.findIndex((job) => job.id === watching?.id); + const syncClipNav = () => { + const index = watchingIndex(); + const show = jobs.length > 1 && index >= 0; + $("clip-nav").hidden = !show; + if (!show) return; + $("older-clip").disabled = index >= jobs.length - 1; + $("newer-clip").disabled = index <= 0; }; const showError = (message = "") => { $("error").textContent = message; @@ -36,137 +47,172 @@ $("copy-curl").disabled = !model || !$("prompt").value.trim() || !$("seed").validity.valid; $("copy-status").textContent = ""; }; - const refreshJobs = async () => { - $("refresh-jobs").disabled = true; - try { - const result = await api("/v1/videos?limit=8&order=desc"); - $("jobs").replaceChildren(); - result.data.forEach((job) => { - const item = document.createElement("li"); - const button = document.createElement("button"); - button.type = "button"; - button.setAttribute("aria-current", String(job.id === currentJob?.id)); - const prompt = document.createElement("span"); - prompt.className = "job-prompt"; - prompt.textContent = job.prompt || job.id; - const state = document.createElement("span"); - state.className = "job-state"; - state.textContent = labels[job.status] || job.status; - button.append(prompt, state); - button.addEventListener("click", () => { - if (busy) { - showError("Wait for this job, or open another playground tab to inspect a different job."); - return; - } - $("prompt").value = job.prompt || ""; - // Jobs do not report the seed; do not pair an old seed with this prompt. - $("seed").value = ""; - updateCurl(); - followJob(job.id); - }); - item.append(button); - $("jobs").append(item); - }); - $("history-status").textContent = result.data.length ? "Showing the latest jobs from this server." : "No jobs yet. Submit a prompt to start."; - } catch (error) { - $("history-status").textContent = `Could not load jobs. Check the server and select Refresh jobs. ${error.message}`; - } finally { - $("refresh-jobs").disabled = false; + const setVideoSource = (job) => { + const video = $("video"); + const complete = job.status === "completed"; + video.hidden = !complete; + $("empty-preview").hidden = complete; + $("download").hidden = !complete; + if (!complete) { + video.pause(); + video.removeAttribute("src"); + delete video.dataset.job; + $("empty-preview").querySelector("h3").textContent = job.status === "failed" ? "Generation failed" : "Your job is on the server"; + $("empty-preview").querySelector("p").textContent = job.status === "failed" + ? "Read the error below before trying again." + : "Queue another prompt or open a finished clip from Recent jobs. This one keeps generating."; + return; + } + $("download").href = `${jobPath(job.id)}/content`; + $("download").download = `${job.id}.mp4`; + if (video.dataset.job !== job.id) { + video.src = `${jobPath(job.id)}/content`; + video.dataset.job = job.id; } }; - const renderJob = (job) => { - currentJob = job; + const renderWatching = (job) => { + watching = job; + const url = new URL(location.href); + url.searchParams.set("job", job.id); + history.replaceState({}, "", url); $("job-state").textContent = labels[job.status] || job.status; $("job-state").dataset.status = job.status; $("job-id").textContent = `Job ${job.id}`; - const path = `/v1/videos/${encodeURIComponent(job.id)}`; - $("job-link").href = path; + $("job-link").href = jobPath(job.id); $("job-link").hidden = false; - const complete = job.status === "completed"; - $("video").hidden = !complete; - $("empty-preview").hidden = complete; - $("download").hidden = !complete; - if (complete) { - $("video").src = `${path}/content`; - $("download").href = `${path}/content`; - $("download").download = `${job.id}.mp4`; - $("job-status").textContent = "Video ready. Change the prompt and generate again without restarting the server."; + setVideoSource(job); + if (job.status === "completed") { + $("job-status").textContent = "Playing this clip. Queue the next prompt without leaving it, or pick another job from Recent jobs."; + } else if (job.status === "queued") { + $("job-status").textContent = "Queued. The server runs one generation at a time. Finished clips stay playable while you wait."; + } else if (job.status === "failed") { + $("job-status").textContent = "The server could not complete this job."; } else { - $("video").pause(); - $("video").removeAttribute("src"); - $("empty-preview").querySelector("h3").textContent = job.status === "failed" ? "Generation failed" : "Your job is on the server"; - $("empty-preview").querySelector("p").textContent = job.status === "failed" ? "Read the error below before trying again." : "You can keep editing the next prompt while you wait."; - $("job-status").textContent = job.status === "queued" ? "Queued. The server runs one generation at a time." : job.status === "failed" ? "The server could not complete this job." : "Generating video and audio. This page checks the job status automatically."; + $("job-status").textContent = "Generating. Open any finished clip from Recent jobs; this job keeps running."; } if (job.status === "failed") showError(job.error?.message || "Check the server logs before submitting another job."); + syncClipNav(); + for (const button of $("jobs").querySelectorAll("button[data-job-id]")) { + if (button.dataset.jobId === job.id) button.setAttribute("aria-current", "true"); + else button.removeAttribute("aria-current"); + } }; - const followJob = async (id) => { - const version = ++pollVersion; - const url = new URL(location.href); - url.searchParams.set("job", id); - history.replaceState({}, "", url); + const selectJob = (job) => { showError(); - setBusy(true); - $("check-status").hidden = true; - const deadline = Date.now() + 30 * 60 * 1000; - let restorePrompt = !$("prompt").value.trim(); - try { - while (version === pollVersion) { - const job = await api(`/v1/videos/${encodeURIComponent(id)}`); - if (version !== pollVersion) return; + renderWatching(job); + }; + const renderHistory = () => { + $("jobs").replaceChildren(); + jobs.forEach((job) => { + const item = document.createElement("li"); + const button = document.createElement("button"); + button.type = "button"; + button.dataset.jobId = job.id; + if (job.id === watching?.id) button.setAttribute("aria-current", "true"); + const prompt = document.createElement("span"); + prompt.className = "job-prompt"; + prompt.textContent = job.prompt || job.id; + const state = document.createElement("span"); + state.className = "job-state"; + state.textContent = labels[job.status] || job.status; + button.append(prompt, state); + button.addEventListener("click", () => selectJob(job)); + item.append(button); + $("jobs").append(item); + }); + const queued = jobs.filter((job) => job.status === "queued").length; + const running = jobs.some((job) => job.status === "in_progress"); + if (!jobs.length) { + $("history-status").textContent = "No jobs yet. Submit a prompt to start."; + } else if (queued || running) { + $("history-status").textContent = running + ? `One clip generating${queued ? `, ${queued} queued` : ""}. Select any row to play or inspect it.` + : `${queued} clip${queued === 1 ? "" : "s"} queued. Select any row to play a finished clip.`; + } else { + $("history-status").textContent = "Select any clip to play it. The prompt box is left alone so you can queue the next one."; + } + }; + const schedulePoll = () => { + clearTimeout(pollTimer); + if (!jobs.some(active)) return; + const version = ++pollVersion; + pollTimer = setTimeout(async () => { + if (version !== pollVersion) return; + try { + await refreshJobs(); $("connection").textContent = readyLabel; - if (restorePrompt) { - $("prompt").value = job.prompt || ""; - updateCurl(); - restorePrompt = false; - } - renderJob(job); - if (!active(job)) break; - if (Date.now() >= deadline) throw new Error("Stopped checking after 30 minutes. The job may still be running."); - await new Promise((resolve) => setTimeout(resolve, 2000)); + $("check-status").hidden = true; + } catch (error) { + $("connection").textContent = "Job status unavailable · check the server"; + showError(`${error.message} Select Check status to reconnect. A connection error does not cancel generation.`); + $("check-status").hidden = false; } + }, 2000); + }; + const refreshJobs = async () => { + $("refresh-jobs").disabled = true; + try { + const result = await api("/v1/videos?limit=32&order=desc"); + jobs = result.data; + const latestWatching = jobs.find((job) => job.id === watching?.id); + if (latestWatching) watching = latestWatching; + renderHistory(); + if (latestWatching) renderWatching(latestWatching); + syncGenerateButton(); + syncClipNav(); + schedulePoll(); } catch (error) { - if (version !== pollVersion) return; - $("connection").textContent = "Job status unavailable · check the server"; - showError(`${error.message} Select Check status to reconnect. A connection error does not cancel generation. If the job is missing, the server may have restarted or the job was deleted.`); - $("check-status").hidden = false; + $("history-status").textContent = `Could not load jobs. Check the server and select Refresh jobs. ${error.message}`; + throw error; } finally { - if (version === pollVersion) { - setBusy(false); - await refreshJobs(); - } + $("refresh-jobs").disabled = false; } }; $("generate-form").addEventListener("submit", async (event) => { event.preventDefault(); - if (busy || !model) return; + if (!model || submitting) return; if (!$("prompt").value.trim()) { showError("Write a prompt before generating a video."); $("prompt").focus(); return; } showError(); - setBusy(true); - $("job-status").textContent = "Submitting the prompt…"; + submitting = true; + syncGenerateButton(); + const keepClip = watching?.status === "completed"; + $("job-status").textContent = keepClip ? "Queuing the next clip…" : "Submitting the prompt…"; try { const job = await api("/v1/videos", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify(payload()), }); - renderJob(job); - await followJob(job.id); + await refreshJobs(); + if (keepClip) { + $("job-status").textContent = "Queued. This clip keeps playing while the server generates the next one."; + } else { + selectJob(job); + } } catch (error) { showError(`${error.message} Check Recent jobs before submitting again; the server may have received the prompt. Submissions are never retried automatically.`); $("job-status").textContent = "Could not confirm the submission."; - setBusy(false); - await refreshJobs(); + await refreshJobs().catch(() => undefined); + } finally { + submitting = false; + syncGenerateButton(); } }); $("prompt").addEventListener("input", updateCurl); $("seed").addEventListener("input", updateCurl); - $("refresh-jobs").addEventListener("click", refreshJobs); + $("refresh-jobs").addEventListener("click", () => refreshJobs().catch(() => undefined)); + $("older-clip").addEventListener("click", () => { + const index = watchingIndex(); + if (index >= 0 && index < jobs.length - 1) selectJob(jobs[index + 1]); + }); + $("newer-clip").addEventListener("click", () => { + const index = watchingIndex(); + if (index > 0) selectJob(jobs[index - 1]); + }); $("check-status").addEventListener("click", () => { - const id = new URL(location.href).searchParams.get("job"); - if (id) followJob(id); + refreshJobs().catch(() => undefined); }); $("video").addEventListener("error", () => { if (!$("video").hasAttribute("src")) return; @@ -189,7 +235,7 @@ $("connection").textContent = readyLabel; $("lifetime").textContent = config.runtime === "mlx" ? "MLX keeps the server and prompt cache available, but releases model components between phases to limit unified-memory use. Closing this page does not cancel a job." - : "The model stays loaded until you stop the server. Closing this page does not cancel a job."; + : "The model stays loaded until you stop the server. Queue the next prompt while a clip plays; closing this page does not cancel a job."; $("model").textContent = model; const d = config.defaults; const facts = []; @@ -198,11 +244,18 @@ if (d.fps) facts.push(`${d.fps} fps`); if (d.seed != null) $("seed").placeholder = String(d.seed); $("settings").textContent = facts.length ? `Server defaults · ${facts.join(" · ")}` : "Resolution and sampling come from the server configuration."; - setBusy(false); + syncGenerateButton(); updateCurl(); + await refreshJobs().catch(() => undefined); const id = new URL(location.href).searchParams.get("job"); - if (id) await followJob(id); - else await refreshJobs(); + const linked = jobs.find((job) => job.id === id); + if (linked) { + if (!$("prompt").value.trim() && linked.prompt) { + $("prompt").value = linked.prompt; + updateCurl(); + } + selectJob(linked); + } } catch (error) { $("connection").textContent = "Server unavailable"; $("history-status").textContent = "Start the H3 server, then reload this page."; diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 284622d595..60c81c72a0 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -8,8 +8,9 @@ variants that may land later (e.g. AMD's MX-FP4 or vendor-neutral e3m0). -The registered config targets the curated LTX-2 deployment set and the -main MiniMax-H3 transformer-block FFN linears. +The registered config targets the curated LTX-2 deployment set, the +main MiniMax-H3 transformer-block FFN linears, and the packed MiniMax-H3 +DiT export (``layer_profile="h3_dit"``) covering attention plus FFN. `flashinfer` is imported lazily inside the call paths that need it. This keeps ``import fastvideo`` cheap on hosts where flashinfer is @@ -19,6 +20,7 @@ from __future__ import annotations import logging +import os import re from typing import Any @@ -76,6 +78,16 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: for suffix in _LTX2_NVFP4_BLOCK_LINEAR_SUFFIXES) | frozenset( ("ltx2.adaln_single.linear", )) _MINIMAX_H3_NVFP4_FF_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.ff\.(?:fc_in|fc_out)$") +_MINIMAX_H3_NVFP4_DIT_PREFIX = re.compile( + r"(?:^|\.)transformer_blocks\.\d+\.(?:attn\.to_(?:q|k|v|out)|ff\.(?:fc_in|fc_out))$") +H3_NVFP4_DIT_EXPORT_FILENAME = "nvfp4_weights.safetensors" +H3_NVFP4_DIT_KEY_SEP = "::" +H3_NVFP4_DIT_BUFFER_NAMES = ( + "_nvfp4_weight", + "_nvfp4_weight_scale", + "_nvfp4_alpha", + "_weight_global_sf", +) def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: @@ -88,6 +100,42 @@ def is_minimax_h3_nvfp4_linear_prefix(prefix: str) -> bool: return _MINIMAX_H3_NVFP4_FF_PREFIX.search(prefix) is not None +def is_minimax_h3_nvfp4_dit_linear_prefix(prefix: str) -> bool: + """Return whether *prefix* is a MiniMax-H3 DiT attention or FFN linear. + + This is the packed NVFP4H3 export set: ``attn.to_{q,k,v,out}`` and + ``ff.{fc_in,fc_out}`` in each main transformer block. Token-refiner, + AdaLN, and embedding linears stay dense. + """ + return _MINIMAX_H3_NVFP4_DIT_PREFIX.search(prefix) is not None + + +def is_minimax_h3_nvfp4_dit_export_path(path: str) -> bool: + """Return whether *path* is the packed MiniMax-H3 DiT NVFP4 export file.""" + return os.path.basename(path) == H3_NVFP4_DIT_EXPORT_FILENAME + + +def find_minimax_h3_nvfp4_dit_export(weight_paths: list[str]) -> str | None: + """Locate ``nvfp4_weights.safetensors`` next to a transformer shard list.""" + seen: list[str] = [] + for path in weight_paths: + if is_minimax_h3_nvfp4_dit_export_path(path) and os.path.isfile(path): + return path + directory = path if os.path.isdir(path) else os.path.dirname(path) + if directory and directory not in seen: + seen.append(directory) + for directory in seen: + candidate = os.path.join(directory, H3_NVFP4_DIT_EXPORT_FILENAME) + if os.path.isfile(candidate): + return candidate + return None + + +def dense_transformer_safetensors(weight_paths: list[str]) -> list[str]: + """Drop the packed NVFP4 DiT export so it is not loaded as bf16 weights.""" + return [path for path in weight_paths if not is_minimax_h3_nvfp4_dit_export_path(path)] + + def _is_ltx2_refine_only_prefix(prefix: str) -> bool: return any(prefix.endswith(suffix) for suffix in _LTX2_REFINE_ONLY_SUFFIXES) @@ -431,14 +479,19 @@ class NVFP4Config(QuantizationConfig): NVFP4 is NVIDIA's block-scaled FP4 (e2m1 mantissa, fp32 alpha, ``layout_128x4`` scale layout, group size 16). LTX-2 uses its curated - attention and FFN deployment set. MiniMax-H3 uses only ``fc_in`` and - ``fc_out`` in each main transformer-block FFN. + attention and FFN deployment set. MiniMax-H3's default profile uses only + ``fc_in`` and ``fc_out`` in each main transformer-block FFN. + ``layer_profile="h3_dit"`` expands that to the packed NVFP4H3 DiT set + (attention ``to_{q,k,v,out}`` plus those FFN linears). """ def __init__(self, layer_profile: str = "refine", retain_original_weights: bool | None = None): super().__init__() - # ``base``: stage-1 set (no attn2.to_out, no cross-modal AV - # projections). ``refine``: full stage-2 set. + if layer_profile not in ("base", "refine", "h3_dit"): + raise ValueError("NVFP4Config.layer_profile must be one of 'base', 'refine', or 'h3_dit', " + f"got {layer_profile!r}") + # ``base`` / ``refine``: LTX-2 stage-1 vs stage-2 layer sets. + # ``h3_dit``: packed MiniMax-H3 attention+FFN export. self.layer_profile = layer_profile # Original bf16 ``layer.weight`` retention after FP4 conversion. # Default (None/False): purge the purgeable originals -- every @@ -473,10 +526,16 @@ def from_config(cls, config: dict[str, Any]) -> NVFP4Config: def get_quant_method(self, layer: torch.nn.Module, prefix: str): from fastvideo.layers.linear import LinearBase + if not isinstance(layer, LinearBase): + return None # LTX-2 switches its active subset by stage at runtime. MiniMax-H3 - # uses the fixed main-transformer FFN set selected by its prefix. - if isinstance(layer, LinearBase) and (is_ltx2_nvfp4_linear_prefix(prefix) - or is_minimax_h3_nvfp4_linear_prefix(prefix)): + # uses the fixed main-transformer FFN set unless ``h3_dit`` selects + # the packed attention+FFN export. + if self.layer_profile == "h3_dit": + tagged = is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + else: + tagged = is_ltx2_nvfp4_linear_prefix(prefix) or is_minimax_h3_nvfp4_linear_prefix(prefix) + if tagged: method = NVFP4QuantizeMethod(layer_prefix=prefix) method._retain_original_weights = self.retain_original_weights return method @@ -553,10 +612,92 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: ) +def nvfp4_linear_weight_param_names(model: torch.nn.Module) -> set[str]: + """State-dict names of ``weight`` on layers tagged with ``NVFP4QuantizeMethod``.""" + names: set[str] = set() + for module_name, module in model.named_modules(): + if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod): + names.add(f"{module_name}.weight" if module_name else "weight") + return names + + +def _module_by_nvfp4_export_prefix(modules: dict[str, torch.nn.Module], prefix: str) -> torch.nn.Module | None: + module = modules.get(prefix) + if module is not None: + return module + if prefix.startswith("minimax_h3."): + return modules.get(prefix[len("minimax_h3."):]) + return modules.get(f"minimax_h3.{prefix}") + + +def load_minimax_h3_nvfp4_dit_export( + model: torch.nn.Module, + path: str, + device: torch.device | str, +) -> int: + """Load a packed NVFP4H3 DiT export onto already-tagged NVFP4 linears. + + Keys are ``::`` with the four buffers + ``convert_model_to_nvfp4`` registers. The bf16 ``weight`` is dropped so a + 32 GB card never materializes the dense GEMMs. Every export prefix must + match an NVFP4 linear, and every NVFP4 linear must appear in the export. + """ + from safetensors import safe_open + + modules = dict(model.named_modules()) + tagged = { + name + for name, module in modules.items() if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod) + } + groups: dict[str, dict[str, str]] = {} + with safe_open(path, framework="pt", device="cpu") as reader: + for key in reader.keys(): # noqa: SIM118 + if H3_NVFP4_DIT_KEY_SEP not in key: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export key {key!r} is missing {H3_NVFP4_DIT_KEY_SEP!r}") + prefix, buffer_name = key.split(H3_NVFP4_DIT_KEY_SEP, 1) + groups.setdefault(prefix, {})[buffer_name] = key + + loaded_names: set[str] = set() + for prefix, buffers in groups.items(): + missing_buffers = [name for name in H3_NVFP4_DIT_BUFFER_NAMES if name not in buffers] + if missing_buffers: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} is missing {missing_buffers}") + extra_buffers = sorted(set(buffers) - set(H3_NVFP4_DIT_BUFFER_NAMES)) + if extra_buffers: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} has unknown buffers {extra_buffers}") + module = _module_by_nvfp4_export_prefix(modules, prefix) + if module is None: + raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} is not in the model") + if not isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod): + raise RuntimeError("MiniMax-H3 NVFP4 DiT export layer " + f"{prefix!r} is not an NVFP4 linear; set NVFP4Config(layer_profile='h3_dit')") + for buffer_name in H3_NVFP4_DIT_BUFFER_NAMES: + tensor = reader.get_tensor(buffers[buffer_name]).to(device=device) + module.register_buffer(buffer_name, tensor, persistent=False) + module.register_parameter("weight", None) + loaded_names.add(next(name for name, candidate in modules.items() if candidate is module)) + + missing_layers = tagged - loaded_names + extra_layers = loaded_names - tagged + if missing_layers or extra_layers: + raise RuntimeError("MiniMax-H3 NVFP4 DiT export does not cover the tagged linear set; " + f"missing={sorted(missing_layers)[:8]} extra={sorted(extra_layers)[:8]}") + logger.info("Loaded MiniMax-H3 NVFP4 DiT export: %d linears from %s", len(loaded_names), path) + return len(loaded_names) + + __all__ = [ + "H3_NVFP4_DIT_BUFFER_NAMES", + "H3_NVFP4_DIT_EXPORT_FILENAME", "NVFP4Config", "NVFP4QuantizeMethod", "convert_model_to_nvfp4", + "dense_transformer_safetensors", + "find_minimax_h3_nvfp4_dit_export", "is_ltx2_nvfp4_linear_prefix", + "is_minimax_h3_nvfp4_dit_export_path", + "is_minimax_h3_nvfp4_dit_linear_prefix", "is_minimax_h3_nvfp4_linear_prefix", + "load_minimax_h3_nvfp4_dit_export", + "nvfp4_linear_weight_param_names", ] diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 148a6d1bb0..ae50f6172f 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -469,7 +469,13 @@ def load_model( from fastvideo.platforms import current_platform - if use_cpu_offload: + if use_cpu_offload and checkpoint_quant_config is not None: + logger.info( + "Skipping FSDP CPU offload for serialized %s text encoder; " + "packed uint8 weights are not FSDP-shardable", + checkpoint_quant_config.get_name(), + ) + elif use_cpu_offload: pin_cpu_memory = fastvideo_args.pin_cpu_memory and is_pin_memory_available() # Disable FSDP for MPS as it's not compatible if current_platform.is_mps(): @@ -876,6 +882,14 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # Find all safetensors files safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) + int8_convrot_path = None + if class_name == "AutoencoderKLMiniMaxH3": + from fastvideo.models.vaes.minimax_h3_int8_convrot import ( + dense_vae_safetensors, + find_int8_convrot_vae_path, + ) + int8_convrot_path = find_int8_convrot_vae_path(model_path) + safetensors_list = dense_vae_safetensors(safetensors_list) if not safetensors_list: raise ValueError(f"No safetensors files found in {model_path}") # Common case: a single `.safetensors` checkpoint file. @@ -911,6 +925,9 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # strictly so missing/unexpected keys are surfaced early. strict_load = class_name in {"AutoencoderKL", "AutoencoderKLMiniMaxH3"} vae.load_state_dict(loaded, strict=strict_load) + if class_name == "AutoencoderKLMiniMaxH3" and int8_convrot_path is not None: + from fastvideo.models.vaes.minimax_h3_int8_convrot import overlay_minimax_h3_int8_convrot_decoder + overlay_minimax_h3_int8_convrot_decoder(vae, int8_convrot_path) if (class_name == "AutoencoderKLWan" and getattr(vae.config, "use_light_vae", False) and target_device.type == "cuda" and hasattr(vae, "optimize_memory_format")): vae.optimize_memory_format() @@ -1068,6 +1085,8 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors")) if not safetensors_list: raise ValueError(f"No safetensors files found in {model_path}") + from fastvideo.layers.quantization.nvfp4_config import dense_transformer_safetensors + safetensors_list = dense_transformer_safetensors(safetensors_list) # arch_config can infer architecture from weight keys (e.g. Flux2 layer counts) update_fn = getattr(dit_config.arch_config, "update_from_weight_keys", None) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 956850a508..ad1867a4b3 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -80,6 +80,11 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor for mod in model.modules(): qm = getattr(mod, "quant_method", None) if isinstance(qm, NVFP4QuantizeMethod): + if any( + getattr(module, "_nvfp4_weight", None) is not None for module in model.modules() + if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod)): + logger.info("NVFP4 packed export already populated; skipping runtime weight conversion") + return if defer_weight_conversion_until_lora_merge: logger.info("Deferring NVFP4 weight conversion until the inference LoRA merge completes") return @@ -214,6 +219,22 @@ def maybe_load_fsdp_model( """ _validate_fsdp_inference_quantization(init_params, fsdp_inference) + from fastvideo.layers.quantization.nvfp4_config import ( + dense_transformer_safetensors, + find_minimax_h3_nvfp4_dit_export, + load_minimax_h3_nvfp4_dit_export, + nvfp4_linear_weight_param_names, + ) + # Always drop the packed export from the dense shard list so ``::`` keys are + # never treated as bf16 parameters. Overlay it only for the H3 DiT profile. + packed_candidate = find_minimax_h3_nvfp4_dit_export(weight_dir_list) + weight_dir_list = dense_transformer_safetensors(weight_dir_list) + quant_config = getattr(init_params.get("config"), "quant_config", None) + packed_nvfp4_export = packed_candidate if getattr(quant_config, "layer_profile", None) == "h3_dit" else None + if packed_nvfp4_export is not None and lora_path is not None: + raise ValueError("Packed MiniMax-H3 NVFP4 DiT export cannot be combined with lora_path; " + "merge the adapter before exporting, or load without the packed file.") + # NOTE(will): cast_forward_inputs=True shouldn't be needed as we are # manually casting the inputs to the model mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=False) @@ -278,8 +299,21 @@ def maybe_load_fsdp_model( # Host offload is already disabled on unified memory (GB10). Staging the # 35B FastH3 DiT on CPU and then copying to CUDA doubled that working set # and took minutes. Follow cpu_offload: read onto the accelerator. - weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=cpu_offload) - logger.info("Loading transformer weights with to_cpu=%s", cpu_offload) + # Packed NVFP4H3 DiT exports still read the remaining dense shards on CPU + # so the 32 GB GEMMs never land on a 32 GB card before they are skipped. + nvfp4_skip_param_names: set[str] = set() + if packed_nvfp4_export is not None: + nvfp4_skip_param_names = nvfp4_linear_weight_param_names(model) + if not nvfp4_skip_param_names: + logger.warning( + "Found %s next to the transformer shards but no NVFP4 linears; " + "ignoring the packed export. Set NVFP4Config(layer_profile='h3_dit').", + packed_nvfp4_export, + ) + packed_nvfp4_export = None + load_weights_to_cpu = cpu_offload or packed_nvfp4_export is not None + weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=load_weights_to_cpu) + logger.info("Loading transformer weights with to_cpu=%s", load_weights_to_cpu) param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping) dense_lora_patch = DenseLoRAPatch.from_adapter( lora_path, @@ -307,7 +341,10 @@ def maybe_load_fsdp_model( cpu_offload=cpu_offload, param_names_mapping=param_names_mapping_fn, dense_lora_patch=dense_lora_patch, + skip_param_names=nvfp4_skip_param_names or None, ) + if packed_nvfp4_export is not None: + load_minimax_h3_nvfp4_dit_export(model, packed_nvfp4_export, device=device) if hasattr(model, "materialize_non_persistent_buffers"): model.materialize_non_persistent_buffers(device=device, dtype=default_dtype) for n, p in chain(model.named_parameters(), model.named_buffers()): @@ -585,6 +622,12 @@ def shard_model( fully_shard(model, **root_kwargs) +def _drop_state_dict_parameter(model: nn.Module, param_name: str) -> None: + module_name, _, attr_name = param_name.rpartition(".") + module = model.get_submodule(module_name) if module_name else model + module.register_parameter(attr_name, None) + + # TODO(PY): device mesh for cfg parallel def load_model_from_full_model_state_dict( model: FSDPModule | torch.nn.Module, @@ -596,6 +639,7 @@ def load_model_from_full_model_state_dict( param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None, training_mode: bool = True, dense_lora_patch: DenseLoRAPatch | None = None, + skip_param_names: set[str] | None = None, ) -> _IncompatibleKeys: """ Converting full state dict into a sharded state dict @@ -634,6 +678,11 @@ def load_model_from_full_model_state_dict( # set. for target_param_name in list(custom_param_sd): full_tensor = custom_param_sd.pop(target_param_name) + if skip_param_names and target_param_name in skip_param_names: + continue + if "::" in target_param_name: + logger.warning("Skipping packed NVFP4 export key mixed into dense shards: %s", target_param_name) + continue meta_sharded_param = meta_sd.get(target_param_name) if meta_sharded_param is None: # Some checkpoints include extra entries that are not part of the @@ -701,6 +750,13 @@ def load_model_from_full_model_state_dict( model.reverse_param_names_mapping = reverse_param_names_mapping unused_keys = set(meta_sd.keys()) - set(sharded_sd.keys()) + skipped_unused = unused_keys & skip_param_names if skip_param_names else set() + for skipped_name in skipped_unused: + _drop_state_dict_parameter(model, skipped_name) + unused_keys -= skipped_unused + if skipped_unused: + logger.info("Deferred %d NVFP4 linear weights to the packed DiT export (%s)", + len(skipped_unused), _summarize_param_names(skipped_unused)) if unused_keys: # Say which of these the adapter is about to fill in. Reporting all of them as # "unloaded" was accurate when zero-init was the only outcome; with an adapter diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py new file mode 100644 index 0000000000..0c2885a9fa --- /dev/null +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -0,0 +1,315 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Comfy ``int8_tensorwise`` + ConvRot overlay for the MiniMax-H3 video VAE decoder. + +The export stores decoder transformer linears as signed int8 with per-output +channel scales and a JSON ``comfy_quant`` marker: + + {"format": "int8_tensorwise", "convrot": true, "convrot_groupsize": 256} + +Weights were rotated offline by a normalized regular Hadamard (group 256). +Inference rotates activations with the same matrix, row-quantizes them, then +runs int8 GEMM. Encoder convolutions stay dense; only the ViT decoder blocks +are quantized. + +Comfy names (``to_qkv``, ``ff.w1`` / ``ff.w2``, ``x_embedder``) are remapped +onto FastVideo's split Q/K/V and ``ff.net`` surface. +""" + +from __future__ import annotations + +import json +import math +import os +from pathlib import Path +from typing import Any + +import torch +import torch.nn as nn +import torch.nn.functional as F +from safetensors.torch import load_file as safetensors_load_file + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) + +INT8_CONVROT_FILENAME = "minimax_h3_video_vae_int8_convrot.safetensors" +_COM_FY_QUANT_SUFFIX = ".comfy_quant" +_DEFAULT_GROUP_SIZE = 256 + +_HADAMARD_CACHE: dict[tuple[int, str, torch.dtype], torch.Tensor] = {} + + +def regular_hadamard(size: int, *, device: torch.device, dtype: torch.dtype) -> torch.Tensor: + """Normalized regular Hadamard of order ``4**k`` (ConvRot Theorem 3.3).""" + cache_key = (size, str(device), dtype) + cached = _HADAMARD_CACHE.get(cache_key) + if cached is not None: + return cached + if size < 4 or (size & (size - 1)) != 0 or math.log(size, 4) % 1: + raise ValueError(f"Regular Hadamard size must be a power of 4, got {size}") + h4 = torch.tensor( + [[1, 1, 1, -1], [1, 1, -1, 1], [1, -1, 1, 1], [-1, 1, 1, 1]], + dtype=dtype, + device=device, + ) + hadamard = h4 + current = 4 + while current < size: + hadamard = torch.kron(hadamard, h4) + current *= 4 + hadamard = hadamard / math.sqrt(size) + _HADAMARD_CACHE[cache_key] = hadamard + return hadamard + + +def rotate_activation(x: torch.Tensor, group_size: int) -> torch.Tensor: + """Apply the block-diagonal ConvRot Hadamard to the last dimension of ``x``.""" + features = x.shape[-1] + if features % group_size: + raise ValueError(f"features {features} are not divisible by convrot group_size {group_size}") + groups = features // group_size + hadamard = regular_hadamard(group_size, device=x.device, dtype=x.dtype) + grouped = x.reshape(*x.shape[:-1], groups, group_size) + return torch.matmul(grouped, hadamard).reshape(x.shape) + + +def parse_comfy_quant_marker(blob: torch.Tensor) -> dict[str, Any]: + """Decode the uint8 JSON marker Comfy stores next to each quantized linear.""" + raw = bytes(blob.detach().cpu().contiguous().view(torch.uint8).numpy()) + raw = raw.split(b"\x00", 1)[0] + marker = json.loads(raw.decode("utf-8")) + if not isinstance(marker, dict): + raise ValueError("comfy_quant marker must be a JSON object") + return marker + + +class Int8ConvRotLinear(nn.Module): + """W8A8 linear matching Comfy ``int8_tensorwise`` (+ optional ConvRot).""" + + def __init__( + self, + in_features: int, + out_features: int, + *, + bias: bool, + convrot: bool, + group_size: int, + ) -> None: + super().__init__() + self.in_features = in_features + self.out_features = out_features + self.convrot = convrot + self.group_size = group_size + 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: + self.bias = nn.Parameter(torch.zeros(out_features), requires_grad=False) + else: + self.register_parameter("bias", None) + + def _dequant_weight(self, dtype: torch.dtype) -> torch.Tensor: + return self.weight.to(dtype) * self.weight_scale.to(dtype) + + @staticmethod + def _dequant_int8_gemm( + acc: torch.Tensor, + x_scale: torch.Tensor, + weight_scale: torch.Tensor, + ) -> torch.Tensor: + # 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() + 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) + else: + 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 _int8_linear_from_tensors( + weight: torch.Tensor, + scale: torch.Tensor, + bias: torch.Tensor | None, + marker: dict[str, Any], +) -> Int8ConvRotLinear: + if weight.dtype != torch.int8: + raise ValueError(f"expected int8 weight, got {weight.dtype}") + if scale.ndim == 1: + scale = scale.unsqueeze(1) + convrot = bool(marker.get("convrot", False)) + group_size = int(marker.get("convrot_groupsize", _DEFAULT_GROUP_SIZE)) if convrot else _DEFAULT_GROUP_SIZE + if marker.get("format") not in (None, "int8_tensorwise"): + raise ValueError(f"unsupported comfy_quant format {marker.get('format')!r}") + layer = Int8ConvRotLinear( + weight.shape[1], + weight.shape[0], + bias=bias is not None, + convrot=convrot and weight.shape[1] % group_size == 0, + group_size=group_size, + ) + layer.weight.copy_(weight) + layer.weight_scale.copy_(scale.to(torch.float32)) + if bias is not None and layer.bias is not None: + layer.bias.data.copy_(bias) + return layer + + +def _split_qkv( + weight: torch.Tensor, + scale: torch.Tensor, + bias: torch.Tensor | None, + marker: dict[str, Any], + *, + heads: int, + dim_head: int, +) -> tuple[Int8ConvRotLinear, Int8ConvRotLinear, Int8ConvRotLinear]: + """Undo Comfy fused QKV: per-head ``[q, k, v]`` then stacked heads.""" + expected = heads * 3 * dim_head + if weight.shape[0] != expected: + raise ValueError(f"fused to_qkv out_features {weight.shape[0]} != heads*3*dim_head {expected}") + if scale.ndim == 1: + scale = scale.unsqueeze(1) + in_features = weight.shape[1] + weight = weight.view(heads, 3, dim_head, in_features) + scale = scale.view(heads, 3, dim_head, 1) + bias_view = None if bias is None else bias.view(heads, 3, dim_head) + pieces = [] + for index in range(3): + w = weight[:, index].reshape(heads * dim_head, in_features) + s = scale[:, index].reshape(heads * dim_head, 1) + b = None if bias_view is None else bias_view[:, index].reshape(heads * dim_head) + pieces.append(_int8_linear_from_tensors(w.contiguous(), s.contiguous(), b, marker)) + return pieces[0], pieces[1], pieces[2] + + +def _swap_swiglu_halves(tensor: torch.Tensor) -> torch.Tensor: + """Comfy ``[gate; value]`` → FastVideo ``[value; gate]`` along out_features.""" + if tensor.shape[0] % 2: + raise ValueError(f"SwiGLU packed rows {tensor.shape[0]} are not even") + half = tensor.shape[0] // 2 + return torch.cat([tensor[half:], tensor[:half]], dim=0) + + +def overlay_minimax_h3_int8_convrot_decoder(vae: nn.Module, checkpoint_path: str | Path) -> int: + """Swap H3 VAE decoder transformer linears for Comfy int8-convrot weights. + + Dense decoder tensors that the export still stores in float (embed, norms, + scales, ``proj_out``) are copied onto the matching FastVideo modules. + Returns the number of quantized linears installed. + """ + path = Path(checkpoint_path) + tensors = safetensors_load_file(str(path)) + decoder = getattr(vae, "decoder", None) + if decoder is None: + raise ValueError("MiniMax-H3 VAE overlay expected a `.decoder` module") + + if "decoder.x_embedder.weight" in tensors and hasattr(decoder, "proj_in"): + decoder.proj_in.weight.data.copy_(tensors["decoder.x_embedder.weight"]) + if decoder.proj_in.bias is not None and "decoder.x_embedder.bias" in tensors: + decoder.proj_in.bias.data.copy_(tensors["decoder.x_embedder.bias"]) + for name in ( + "decoder.register_tokens", + "decoder.norm_out.weight", + "decoder.norm_out.bias", + "decoder.proj_out.weight", + "decoder.proj_out.bias", + ): + if name not in tensors: + continue + module_name, _, param_name = name.removeprefix("decoder.").rpartition(".") + target = decoder if not module_name else decoder.get_submodule(module_name) + getattr(target, param_name).data.copy_(tensors[name]) + + quantized = 0 + blocks = decoder.transformer_blocks + for index, block in enumerate(blocks): + prefix = f"decoder.transformer_blocks.{index}" + marker = parse_comfy_quant_marker(tensors[f"{prefix}.attn.to_qkv{_COM_FY_QUANT_SUFFIX}"]) + to_q, to_k, to_v = _split_qkv( + tensors[f"{prefix}.attn.to_qkv.weight"], + tensors[f"{prefix}.attn.to_qkv.weight_scale"], + tensors.get(f"{prefix}.attn.to_qkv.bias"), + marker, + heads=int(block.attn.heads), + dim_head=int(block.attn.dim_head), + ) + block.attn.to_q = to_q + block.attn.to_k = to_k + block.attn.to_v = to_v + quantized += 3 + out_marker = parse_comfy_quant_marker(tensors[f"{prefix}.attn.to_out{_COM_FY_QUANT_SUFFIX}"]) + to_out = _int8_linear_from_tensors( + tensors[f"{prefix}.attn.to_out.weight"], + tensors[f"{prefix}.attn.to_out.weight_scale"], + tensors.get(f"{prefix}.attn.to_out.bias"), + out_marker, + ) + block.attn.to_out[0] = to_out + quantized += 1 + w1_marker = parse_comfy_quant_marker(tensors[f"{prefix}.ff.w1{_COM_FY_QUANT_SUFFIX}"]) + w1_weight = _swap_swiglu_halves(tensors[f"{prefix}.ff.w1.weight"]) + w1_scale = _swap_swiglu_halves(tensors[f"{prefix}.ff.w1.weight_scale"]) + w1_bias = tensors.get(f"{prefix}.ff.w1.bias") + if w1_bias is not None: + w1_bias = _swap_swiglu_halves(w1_bias) + block.ff.net[0].proj = _int8_linear_from_tensors(w1_weight, w1_scale, w1_bias, w1_marker) + quantized += 1 + w2_marker = parse_comfy_quant_marker(tensors[f"{prefix}.ff.w2{_COM_FY_QUANT_SUFFIX}"]) + block.ff.net[2] = _int8_linear_from_tensors( + tensors[f"{prefix}.ff.w2.weight"], + tensors[f"{prefix}.ff.w2.weight_scale"], + tensors.get(f"{prefix}.ff.w2.bias"), + w2_marker, + ) + quantized += 1 + for dense_name in ("norm1.weight", "norm2.weight", "scale1", "scale2"): + key = f"{prefix}.{dense_name}" + if key not in tensors: + continue + if dense_name.endswith(".weight"): + module = block.get_submodule(dense_name.rsplit(".", 1)[0]) + module.weight.data.copy_(tensors[key]) + else: + getattr(block, dense_name).data.copy_(tensors[key]) + + device = next(vae.parameters()).device + vae.decoder.to(device) + logger.info("Overlaid MiniMax-H3 int8-convrot VAE decoder: %s linears from %s", quantized, path) + return quantized + + +def is_int8_convrot_vae_path(path: str | Path) -> bool: + """Return whether *path* is the Comfy int8-convrot overlay file.""" + raw = os.path.basename(str(path)) + real = os.path.basename(os.path.realpath(str(path))) + return INT8_CONVROT_FILENAME in (raw, real) or "int8_convrot" in raw or "int8_convrot" in real + + +def dense_vae_safetensors(paths: list[str]) -> list[str]: + """Drop the ConvRot overlay so it is not loaded as a dense VAE shard.""" + return [path for path in paths if not is_int8_convrot_vae_path(path)] + + +def find_int8_convrot_vae_path(model_path: str | Path) -> Path | None: + """Return the ConvRot export next to a VAE directory, if present.""" + candidate = Path(model_path) / INT8_CONVROT_FILENAME + return candidate if candidate.is_file() else None diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index b39f5c443c..1b78b4645f 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -11,6 +11,7 @@ from typing import Any import torch +from torch.distributed.tensor import DTensor from fastvideo.attention.selector import (_active_component_attention_backend_scope, coerce_attn_backend, get_env_variable_attn_backend) @@ -77,6 +78,15 @@ def _checkpoint_has_vsa_gates(transformer_dir: Path) -> bool: return False +def _module_has_dtensor_params(module: Any) -> bool: + """FSDP2 shards stay put; packed NVFP4 params are ordinary tensors and can ping-pong.""" + parameters = getattr(module, "parameters", None) + if not callable(parameters): + return False + first = next(parameters(), None) + return first is not None and isinstance(first, DTensor) + + @dataclass(frozen=True) class _H3VideoGeometry: spatial_compression_ratio: int @@ -390,24 +400,47 @@ def _load_denoise_modules(self, fastvideo_args: FastVideoArgs) -> None: finally: self._required_config_modules = saved + def _unified_memory_host(self) -> bool: + from fastvideo.pipelines import composed_pipeline_base + from fastvideo.platforms import current_platform + + device = composed_pipeline_base.get_local_torch_device() + device_id = 0 if device.index is None else int(device.index) + return bool(current_platform.has_unified_memory(device_id)) + def _release_text_encoder(self) -> None: - stage = self._stage_name_mapping.get("conditioning_stage") - if stage is not None: - stage.conditioner = None - encoder = self.modules.pop("text_encoder", None) + encoder = self.get_module("text_encoder") if encoder is None: return - logger.info("Released MiniMax-H3 text encoder after conditioning") - del encoder + # Unified-memory boxes cannot keep Qwen around even on "CPU". Discrete + # GPUs can: pin it in host RAM and borrow the GPU only for encode. + if self._unified_memory_host(): + stage = self._stage_name_mapping.get("conditioning_stage") + if stage is not None: + stage.conditioner = None + self.modules.pop("text_encoder", None) + logger.info("Released MiniMax-H3 text encoder after conditioning") + del encoder + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return + if not self._move_module(encoder, "cpu"): + return gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() + logger.info("Moved MiniMax-H3 text encoder to CPU after conditioning") def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: - """Reload Qwen3-VL after `_release_text_encoder` so a later request can encode.""" + """Reload or GPU-restore Qwen3-VL so a later request can encode.""" encoder = self.get_module("text_encoder") stage = self._stage_name_mapping.get("conditioning_stage") if encoder is not None: + if not self._unified_memory_host(): + from fastvideo.pipelines import composed_pipeline_base + + self._move_module(encoder, composed_pipeline_base.get_local_torch_device()) if stage is not None and getattr(stage, "conditioner", None) is None: stage.conditioner = encoder return @@ -424,11 +457,50 @@ def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: if stage is not None: stage.conditioner = self.get_module("text_encoder") + def _move_module(self, module: Any, device: str | torch.device) -> bool: + # MiniMaxH3ConditioningStage already refuses .to() on DTensor; default + # sequential Qwen is FSDP2-sharded. Packed NVFP4 skips that wrap. + if _module_has_dtensor_params(module): + return False + module.to(device) + return True + + def _park_denoise_modules(self) -> None: + """Free the GPU so the NVFP4 encoder can encode without sitting on the DiT.""" + parked = False + for name in _DENOISE_MODULE_NAMES: + module = self.get_module(name) + if module is None: + continue + if self._move_module(module, "cpu"): + parked = True + if parked: + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + logger.info("Parked MiniMax-H3 denoise modules on CPU for text encode") + + def _restore_denoise_modules(self) -> None: + from fastvideo.pipelines import composed_pipeline_base + + device = composed_pipeline_base.get_local_torch_device() + restored = False + for name in _DENOISE_MODULE_NAMES: + module = self.get_module(name) + if module is None: + continue + if self._move_module(module, device): + restored = True + if restored: + logger.info("Restored MiniMax-H3 denoise modules to %s", device) + def _run_condition_then_denoise(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: for name in ("input_preparation_stage", "conditioning_stage"): batch = self._stage_name_mapping[name](batch, fastvideo_args) self._release_text_encoder() self._load_denoise_modules(fastvideo_args) + if not self._unified_memory_host(): + self._restore_denoise_modules() if not self._denoise_stages_ready: self._add_denoise_stages(ref2va=self._ref2va) for name in ( @@ -530,6 +602,8 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward # when neither deferral flag was on. if self._defer_denoise_modules(fastvideo_args): try: + if not self._unified_memory_host(): + self._park_denoise_modules() self._ensure_text_encoder(fastvideo_args) if self._denoise_stages_ready: logger.info("Running MiniMax-H3 condition stages before denoise (subsequent request)") diff --git a/fastvideo/tests/api/test_typed_quant_flow.py b/fastvideo/tests/api/test_typed_quant_flow.py index b7820d987d..1a3c6cf9aa 100644 --- a/fastvideo/tests/api/test_typed_quant_flow.py +++ b/fastvideo/tests/api/test_typed_quant_flow.py @@ -56,6 +56,21 @@ def test_typed_transformer_quant_resolves_to_nvfp4_instance(captured_kwargs) -> assert isinstance(captured_kwargs["transformer_quant"], NVFP4Config), (f"Expected NVFP4Config instance, got " f"{type(captured_kwargs['transformer_quant']).__name__}") + assert captured_kwargs["transformer_quant"].layer_profile == "refine" + + +def test_typed_nvfp4_layer_profile_pins_h3_dit(captured_kwargs) -> None: + from fastvideo.layers.quantization.nvfp4_config import NVFP4Config as CurrentNVFP4Config + + cfg = GeneratorConfig( + model_path="FastVideo/LTX2-Distilled-Diffusers", + engine=EngineConfig(quantization=QuantizationConfig(transformer_quant="NVFP4", layer_profile="h3_dit"), ), + ) + generator_config_to_fastvideo_args(cfg) + quant = captured_kwargs["transformer_quant"] + assert type(quant).__name__ == "NVFP4Config" + assert getattr(quant, "layer_profile") == "h3_dit" + assert isinstance(quant, CurrentNVFP4Config) def test_no_typed_quant_omits_transformer_quant_kwarg(captured_kwargs) -> None: diff --git a/fastvideo/tests/entrypoints/test_openai_video_client.py b/fastvideo/tests/entrypoints/test_openai_video_client.py index 3e48b8e647..7355ce8c44 100644 --- a/fastvideo/tests/entrypoints/test_openai_video_client.py +++ b/fastvideo/tests/entrypoints/test_openai_video_client.py @@ -191,6 +191,7 @@ def test_playground_assets_and_config_do_not_generate(local_server): assert response.headers["X-Content-Type-Options"] == "nosniff" html = response.read().decode() assert "Generate video" in html + assert "Older clip" in html assert 'src="./playground.js"' in html for asset in ["playground.js", "playground.css"]: with urlopen(origin + "/playground/" + asset) as response: diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_config.py b/fastvideo/tests/ops/quantization/test_nvfp4_config.py index 977528189c..3a9fa0d4dc 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_config.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_config.py @@ -44,10 +44,18 @@ def test_nvfp4config_layer_profile_round_trips_from_dict(): from fastvideo.layers.quantization.nvfp4_config import NVFP4Config config = NVFP4Config.from_config({"layer_profile": "base"}) assert config.layer_profile == "base" + config = NVFP4Config.from_config({"layer_profile": "h3_dit"}) + assert config.layer_profile == "h3_dit" config = NVFP4Config.from_config({}) assert config.layer_profile == "refine" +def test_nvfp4config_rejects_unknown_layer_profile(): + from fastvideo.layers.quantization.nvfp4_config import NVFP4Config + with pytest.raises(ValueError, match="h3_dit"): + NVFP4Config(layer_profile="h3-dit") + + def test_nvfp4_kernel_call_raises_clear_error_without_flashinfer(monkeypatch): """A call into the NVFP4 kernels must raise an actionable ImportError when flashinfer is missing, not a confusing diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py b/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py new file mode 100644 index 0000000000..80a7f4e3df --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py @@ -0,0 +1,88 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU tests for the packed MiniMax-H3 DiT NVFP4 export overlay.""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn as nn +from safetensors.torch import save_file + +import fastvideo.layers.quantization.nvfp4_config as nv + + +def _method(prefix: str) -> nv.NVFP4QuantizeMethod: + method = object.__new__(nv.NVFP4QuantizeMethod) + method.weight_fp4 = None + method.weight_scale = None + method.x_global_sf = torch.tensor(1.0, dtype=torch.float32) + method.layer_prefix = prefix + method._is_refine_only_layer = False + method._retain_original_weights = None + return method + + +class _ExportLinear(nn.Module): + + def __init__(self, prefix: str, out_dim: int = 8, in_dim: int = 16) -> None: + super().__init__() + self.weight = nn.Parameter(torch.ones(out_dim, in_dim, dtype=torch.bfloat16), requires_grad=False) + self.quant_method = _method(prefix) + + +def _model() -> nn.Module: + root = nn.Module() + root.transformer_blocks = nn.ModuleList([nn.Module()]) + root.transformer_blocks[0].attn = nn.Module() + root.transformer_blocks[0].attn.to_q = _ExportLinear("transformer_blocks.0.attn.to_q") + return root + + +def test_dense_transformer_safetensors_drops_packed_export(tmp_path) -> None: + packed = tmp_path / nv.H3_NVFP4_DIT_EXPORT_FILENAME + shard = tmp_path / "diffusion_pytorch_model-00001-of-00008.safetensors" + packed.write_bytes(b"") + shard.write_bytes(b"") + kept = nv.dense_transformer_safetensors([str(packed), str(shard)]) + assert kept == [str(shard)] + assert nv.find_minimax_h3_nvfp4_dit_export([str(shard)]) == str(packed) + + +def test_load_minimax_h3_nvfp4_dit_export_overlays_and_purges_weight(tmp_path) -> None: + model = _model() + linear = model.transformer_blocks[0].attn.to_q + prefix = "transformer_blocks.0.attn.to_q" + tensors = { + f"{prefix}::_nvfp4_weight": torch.zeros(8, 8, dtype=torch.uint8), + f"{prefix}::_nvfp4_weight_scale": torch.zeros(8, 1, dtype=torch.uint8), + f"{prefix}::_nvfp4_alpha": torch.tensor(0.5, dtype=torch.float32), + f"{prefix}::_weight_global_sf": torch.tensor(2.0, dtype=torch.bfloat16), + } + path = tmp_path / nv.H3_NVFP4_DIT_EXPORT_FILENAME + save_file(tensors, str(path)) + + loaded = nv.load_minimax_h3_nvfp4_dit_export(model, str(path), device="cpu") + + assert loaded == 1 + assert linear.weight is None + assert linear._nvfp4_weight.dtype is torch.uint8 + assert tuple(linear._nvfp4_weight.shape) == (8, 8) + assert linear._nvfp4_alpha.item() == 0.5 + assert linear._weight_global_sf.item() == 2.0 + + +def test_load_minimax_h3_nvfp4_dit_export_rejects_untagged_linear(tmp_path) -> None: + model = _model() + model.transformer_blocks[0].attn.to_q.quant_method = object() + prefix = "transformer_blocks.0.attn.to_q" + tensors = { + f"{prefix}::_nvfp4_weight": torch.zeros(8, 8, dtype=torch.uint8), + f"{prefix}::_nvfp4_weight_scale": torch.zeros(8, 1, dtype=torch.uint8), + f"{prefix}::_nvfp4_alpha": torch.tensor(1.0, dtype=torch.float32), + f"{prefix}::_weight_global_sf": torch.tensor(1.0, dtype=torch.bfloat16), + } + path = tmp_path / nv.H3_NVFP4_DIT_EXPORT_FILENAME + save_file(tensors, str(path)) + + with pytest.raises(RuntimeError, match="layer_profile='h3_dit'"): + nv.load_minimax_h3_nvfp4_dit_export(model, str(path), device="cpu") diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py b/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py index d5529c04b5..0e03d2daff 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_minimax_h3_wiring.py @@ -47,6 +47,35 @@ def test_is_minimax_h3_nvfp4_linear_prefix_non_main_ffn_linear(prefix: str) -> N assert not nvfp4.is_minimax_h3_nvfp4_linear_prefix(prefix) +@pytest.mark.parametrize( + "prefix", + [ + "transformer_blocks.0.attn.to_q", + "transformer_blocks.41.attn.to_out", + "minimax_h3.transformer_blocks.12.ff.fc_in", + "transformer_blocks.0.attn.to_k", + "transformer_blocks.0.attn.to_v", + "transformer_blocks.3.ff.fc_out", + ], +) +def test_is_minimax_h3_nvfp4_dit_linear_prefix_attn_and_ffn(prefix: str) -> None: + assert nvfp4.is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + + +@pytest.mark.parametrize( + "prefix", + [ + "transformer_blocks.0.adaln_proj.linear", + "token_refiner.refiner_blocks.0.ff.fc_in", + "transformer_blocks.0.attn.to_gate_compress", + "proj_in", + "transformer_blocks.0.attn.to_q.weight", + ], +) +def test_is_minimax_h3_nvfp4_dit_linear_prefix_rejects_non_export_linears(prefix: str) -> None: + assert not nvfp4.is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + + def test_nvfp4config_get_quant_method_minimax_h3_feed_forward() -> None: if not torch.cuda.is_available(): pytest.skip("NVFP4QuantizeMethod construction requires CUDA") diff --git a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index 02955b73fe..c47c044b08 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -27,19 +27,30 @@ def region(self, name): def _stub_module(name: str) -> SimpleNamespace: if name in {"scheduler"}: - return SimpleNamespace(shift=12.0, name=name) - if name in {"audio_scheduler"}: - return SimpleNamespace(shift=3.0, name=name) - if name == "transformer": + module = SimpleNamespace(shift=12.0, name=name) + elif name in {"audio_scheduler"}: + module = SimpleNamespace(shift=3.0, name=name) + elif name == "transformer": # LoRAPipeline reads exclude_lora_layers off the DiT arch config. - return SimpleNamespace( + module = SimpleNamespace( name=name, config=SimpleNamespace(arch_config=SimpleNamespace(exclude_lora_layers=[])), ) - return SimpleNamespace(name=name) + else: + module = SimpleNamespace(name=name) + def to(device): + module.device = device + module.moved_to.append(device) + return module -def _patch_pipeline_construction(monkeypatch, events: list) -> None: + module.device = None + module.moved_to = [] + module.to = to + return module + + +def _patch_pipeline_construction(monkeypatch, events: list, *, unified_memory: bool = False) -> None: monkeypatch.setattr( composed_pipeline_base, "maybe_init_distributed_environment_and_model_parallel", @@ -49,7 +60,7 @@ def _patch_pipeline_construction(monkeypatch, events: list) -> None: monkeypatch.setattr(composed_pipeline_base, "get_world_group", lambda: SimpleNamespace(local_rank=0)) monkeypatch.setattr(composed_pipeline_base, "get_or_create_profiler", lambda trace_dir: _Profiler()) monkeypatch.setattr(composed_pipeline_base, "warmup_sequence_parallel_communication", lambda: None) - monkeypatch.setattr("fastvideo.platforms.current_platform.has_unified_memory", lambda device_id: False) + monkeypatch.setattr("fastvideo.platforms.current_platform.has_unified_memory", lambda device_id: unified_memory) monkeypatch.setattr("fastvideo.platforms.current_platform.is_mps", lambda: False) @@ -103,6 +114,7 @@ def fake_add_denoise(*, ref2va: bool) -> None: monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + encoder = pipeline.get_module("text_encoder") batch = ForwardBatch(data_type="video", prompt="alpine dancer") out = pipeline.forward(batch, args) @@ -111,13 +123,138 @@ def fake_add_denoise(*, ref2va: bool) -> None: assert "transformer" in loads[1] assert "vae" in loads[1] assert "text_encoder" not in loads[1] - assert pipeline.get_module("text_encoder") is None - assert condition_stage.conditioner is None + assert pipeline.get_module("text_encoder") is encoder + assert encoder.moved_to[-1] == "cpu" + assert condition_stage.conditioner is not None assert pipeline.get_module("transformer") is not None assert pipeline._denoise_stages_ready is True + moves_before_second = len(encoder.moved_to) second = pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) assert second is not None + assert len(loads) == 2 + assert pipeline.get_module("text_encoder") is encoder + assert condition_stage.conditioner is not None + assert len(encoder.moved_to) > moves_before_second + assert encoder.moved_to[moves_before_second] == torch.device("cpu") + + +def test_sequential_skips_host_offload_for_dtensor_params(monkeypatch) -> None: + class _FakeDTensor: + pass + + monkeypatch.setattr( + "fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.DTensor", + _FakeDTensor, + ) + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + loads: list[list[str]] = [] + + def _dtensor_stub(name: str) -> SimpleNamespace: + module = _stub_module(name) + + def parameters(): + yield _FakeDTensor() + + module.parameters = parameters + return module + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + if name in modules: + continue + modules[name] = _dtensor_stub(name) if name == "text_encoder" else _stub_module(name) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=True, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + + passthrough = lambda batch, _args: batch + monkeypatch.setattr(pipeline._stage_name_mapping["input_preparation_stage"], "forward", passthrough) + monkeypatch.setattr(pipeline._stage_name_mapping["conditioning_stage"], "forward", passthrough) + original_add_denoise = pipeline._add_denoise_stages + + def fake_add_denoise(*, ref2va: bool) -> None: + original_add_denoise(ref2va=ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + monkeypatch.setattr(pipeline._stage_name_mapping[name], "forward", passthrough) + + monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + encoder = pipeline.get_module("text_encoder") + pipeline.forward(ForwardBatch(data_type="video", prompt="alpine dancer"), args) + pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) + + assert len(loads) == 2 + assert encoder.moved_to == [] + transformer = pipeline.get_module("transformer") + assert transformer is not None + assert transformer.moved_to[-1] == torch.device("cpu") + + +def test_unified_memory_sequential_deletes_encoder_and_reloads(monkeypatch) -> None: + events: list = [] + _patch_pipeline_construction(monkeypatch, events, unified_memory=True) + loads: list[list[str]] = [] + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + modules.setdefault(name, _stub_module(name)) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=True, + lazy_module_load=False, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + + condition_stage = pipeline._stage_name_mapping["conditioning_stage"] + passthrough = lambda batch, _args: batch + monkeypatch.setattr(pipeline._stage_name_mapping["input_preparation_stage"], "forward", passthrough) + monkeypatch.setattr(condition_stage, "forward", passthrough) + original_add_denoise = pipeline._add_denoise_stages + + def fake_add_denoise(*, ref2va: bool) -> None: + original_add_denoise(ref2va=ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + monkeypatch.setattr(pipeline._stage_name_mapping[name], "forward", passthrough) + + monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + pipeline.forward(ForwardBatch(data_type="video", prompt="alpine dancer"), args) + + assert pipeline.get_module("text_encoder") is None + assert condition_stage.conditioner is None + assert len(loads) == 2 + + pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) assert len(loads) == 3 assert loads[2] == ["text_encoder"] assert pipeline.get_module("text_encoder") is None diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py new file mode 100644 index 0000000000..8c39d02695 --- /dev/null +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py @@ -0,0 +1,174 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contracts for the MiniMax-H3 Comfy int8-convrot VAE overlay.""" + +from __future__ import annotations + +import json + +import torch +import torch.nn as nn +from safetensors.torch import save_file + +from fastvideo.models.vaes.minimax_h3_int8_convrot import ( + Int8ConvRotLinear, + dense_vae_safetensors, + overlay_minimax_h3_int8_convrot_decoder, + parse_comfy_quant_marker, + regular_hadamard, + rotate_activation, +) + + +def test_regular_hadamard_is_orthogonal() -> None: + hadamard = regular_hadamard(16, device=torch.device("cpu"), dtype=torch.float64) + identity = torch.eye(16, dtype=torch.float64) + torch.testing.assert_close(hadamard @ hadamard.T, identity, atol=1e-12, rtol=0.0) + + +def test_convrot_rotation_is_an_involution() -> None: + torch.manual_seed(0) + x = torch.randn(4, 256) + rotated = rotate_activation(x, 256) + twice = rotate_activation(rotated, 256) + torch.testing.assert_close(twice, x, atol=1e-5, rtol=1e-5) + + +def test_int8_gemm_scales_in_float32_not_fp16() -> None: + acc = torch.tensor([[100_000]], dtype=torch.int32) + x_scale = torch.tensor([[1.0 / 127.0]], dtype=torch.float16) + weight_scale = torch.tensor([[1.0 / 127.0]], dtype=torch.float16) + overflowed = acc.to(torch.float16) * x_scale * weight_scale.t() + scaled = Int8ConvRotLinear._dequant_int8_gemm(acc, x_scale, weight_scale) + assert torch.isinf(overflowed).all() + assert torch.isfinite(scaled).all() + torch.testing.assert_close(scaled, acc.float() * x_scale.float() * weight_scale.t().float()) + + +def test_int8_linear_matches_dequantized_matmul() -> None: + torch.manual_seed(1) + layer = Int8ConvRotLinear(256, 32, bias=True, convrot=True, group_size=256) + layer.weight.copy_(torch.randint(-8, 8, (32, 256), dtype=torch.int8)) + layer.weight_scale.copy_(torch.linspace(0.01, 0.02, 32).unsqueeze(1)) + layer.bias.data.copy_(torch.randn(32)) + x = torch.randn(3, 256) + out = layer(x) + rotated = rotate_activation(x, 256) + expected = torch.nn.functional.linear(rotated, layer.weight.float() * layer.weight_scale, layer.bias) + torch.testing.assert_close(out, expected, atol=1e-5, rtol=1e-5) + + +def test_dense_vae_safetensors_drops_convrot_overlay() -> None: + kept = dense_vae_safetensors([ + "/tmp/vae/diffusion_pytorch_model-00001-of-00003.safetensors", + "/tmp/vae/minimax_h3_video_vae_int8_convrot.safetensors", + "/tmp/vae/other.safetensors", + ]) + assert kept == [ + "/tmp/vae/diffusion_pytorch_model-00001-of-00003.safetensors", + "/tmp/vae/other.safetensors", + ] + + +def test_parse_comfy_quant_marker_reads_padded_uint8() -> None: + payload = json.dumps({"format": "int8_tensorwise", "convrot": True, "convrot_groupsize": 256}).encode() + blob = torch.zeros(72, dtype=torch.uint8) + blob[:len(payload)] = torch.tensor(list(payload), dtype=torch.uint8) + marker = parse_comfy_quant_marker(blob) + assert marker["format"] == "int8_tensorwise" + assert marker["convrot"] is True + assert marker["convrot_groupsize"] == 256 + + +class _FakeH3VAE(nn.Module): + + def __init__(self) -> None: + super().__init__() + decoder = nn.Module() + decoder.proj_in = nn.Linear(24, 8, bias=True) + decoder.register_tokens = nn.Parameter(torch.zeros(1, 4, 8), requires_grad=False) + decoder.norm_out = nn.LayerNorm(8) + decoder.proj_out = nn.Linear(8, 12, bias=True) + block = nn.Module() + block.norm1 = nn.RMSNorm(8, elementwise_affine=True) + block.norm2 = nn.RMSNorm(8, elementwise_affine=True) + block.scale1 = nn.Parameter(torch.zeros(8), requires_grad=False) + block.scale2 = nn.Parameter(torch.zeros(8), requires_grad=False) + attn = nn.Module() + attn.heads = 2 + attn.dim_head = 4 + attn.to_q = nn.Linear(8, 8, bias=True) + attn.to_k = nn.Linear(8, 8, bias=True) + attn.to_v = nn.Linear(8, 8, bias=True) + attn.to_out = nn.ModuleList([nn.Linear(8, 8, bias=True)]) + ff = nn.Module() + swiglu = nn.Module() + swiglu.proj = nn.Linear(8, 16, bias=True) + ff.net = nn.ModuleList([swiglu, nn.Dropout(0.0), nn.Linear(8, 8, bias=True)]) + block.attn = attn + block.ff = ff + decoder.transformer_blocks = nn.ModuleList([block]) + self.decoder = decoder + + +def _marker_tensor() -> torch.Tensor: + payload = json.dumps({"format": "int8_tensorwise", "convrot": False}).encode() + blob = torch.zeros(64, dtype=torch.uint8) + blob[:len(payload)] = torch.tensor(list(payload), dtype=torch.uint8) + return blob + + +def test_overlay_splits_fused_qkv_and_ffn(tmp_path) -> None: + tensors = { + "decoder.x_embedder.weight": torch.ones(8, 24), + "decoder.x_embedder.bias": torch.zeros(8), + "decoder.register_tokens": torch.ones(1, 4, 8), + "decoder.norm_out.weight": torch.ones(8), + "decoder.norm_out.bias": torch.zeros(8), + "decoder.proj_out.weight": torch.ones(12, 8), + "decoder.proj_out.bias": torch.zeros(12), + "decoder.transformer_blocks.0.attn.to_qkv.weight": torch.arange(24 * 8, dtype=torch.int8).reshape(24, 8), + "decoder.transformer_blocks.0.attn.to_qkv.weight_scale": torch.ones(24, 1), + "decoder.transformer_blocks.0.attn.to_qkv.bias": torch.zeros(24), + "decoder.transformer_blocks.0.attn.to_qkv.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.attn.to_out.weight": torch.ones(8, 8, dtype=torch.int8), + "decoder.transformer_blocks.0.attn.to_out.weight_scale": torch.ones(8, 1), + "decoder.transformer_blocks.0.attn.to_out.bias": torch.zeros(8), + "decoder.transformer_blocks.0.attn.to_out.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.ff.w1.weight": torch.cat( + [torch.ones(8, 8, dtype=torch.int8), torch.full((8, 8), 2, dtype=torch.int8)], dim=0), + "decoder.transformer_blocks.0.ff.w1.weight_scale": torch.cat( + [torch.ones(8, 1), torch.full((8, 1), 0.5)], dim=0), + "decoder.transformer_blocks.0.ff.w1.bias": torch.cat([torch.ones(8), torch.full((8, ), 3.0)], dim=0), + "decoder.transformer_blocks.0.ff.w1.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.ff.w2.weight": torch.ones(8, 8, dtype=torch.int8), + "decoder.transformer_blocks.0.ff.w2.weight_scale": torch.ones(8, 1), + "decoder.transformer_blocks.0.ff.w2.bias": torch.zeros(8), + "decoder.transformer_blocks.0.ff.w2.comfy_quant": _marker_tensor(), + "decoder.transformer_blocks.0.norm1.weight": torch.ones(8), + "decoder.transformer_blocks.0.norm2.weight": torch.ones(8), + "decoder.transformer_blocks.0.scale1": torch.ones(8), + "decoder.transformer_blocks.0.scale2": torch.ones(8), + } + path = tmp_path / "minimax_h3_video_vae_int8_convrot.safetensors" + save_file(tensors, path) + vae = _FakeH3VAE() + installed = overlay_minimax_h3_int8_convrot_decoder(vae, path) + assert installed == 6 + q = vae.decoder.transformer_blocks[0].attn.to_q + k = vae.decoder.transformer_blocks[0].attn.to_k + v = vae.decoder.transformer_blocks[0].attn.to_v + fused = tensors["decoder.transformer_blocks.0.attn.to_qkv.weight"].view(2, 3, 4, 8) + assert isinstance(q, Int8ConvRotLinear) + assert q.weight.shape == (8, 8) + assert torch.equal(q.weight, fused[:, 0].reshape(8, 8)) + assert torch.equal(k.weight, fused[:, 1].reshape(8, 8)) + assert torch.equal(v.weight, fused[:, 2].reshape(8, 8)) + proj = vae.decoder.transformer_blocks[0].ff.net[0].proj + assert isinstance(proj, Int8ConvRotLinear) + assert isinstance(vae.decoder.transformer_blocks[0].ff.net[2], Int8ConvRotLinear) + assert torch.equal(proj.weight[:8], torch.full((8, 8), 2, dtype=torch.int8)) + assert torch.equal(proj.weight[8:], torch.ones(8, 8, dtype=torch.int8)) + assert torch.equal(proj.weight_scale[:8], torch.full((8, 1), 0.5)) + assert torch.equal(proj.bias[:8], torch.full((8, ), 3.0)) + assert not hasattr(vae.decoder, "mask_token") + assert torch.equal(vae.decoder.proj_in.weight, torch.ones(8, 24)) From c38b4e27ef21e7eeb40cd82a432d5bfae2178b7b Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 14:53:16 -0700 Subject: [PATCH 002/107] [feat]: add CompactH3 NVFP4 cookbook paths for Blackwell GPUs Drop playground queue experiments from this PR, reject ConvRot overlays that cannot rotate activations, and skip FSDP2 modules that mix dense params with DTensors. (cherry picked from commit f16b87eb28eb0966d8d200058ccfb30c2713ca72) --- ...vfp4.yaml => basic_compacth3_rtx5090.yaml} | 0 .../basic/basic_compacth3_rtx_pro6000.yaml | 57 ++++ ...fp4.yaml => openai_compacth3_rtx5090.yaml} | 0 ...yaml => openai_compacth3_rtx_pro6000.yaml} | 0 .../openai_fasth3_5090_nvfp4_taeh3.yaml | 57 ---- .../entrypoints/openai/static/playground.css | 3 - .../entrypoints/openai/static/playground.html | 10 +- .../entrypoints/openai/static/playground.js | 261 +++++++----------- 8 files changed, 164 insertions(+), 224 deletions(-) rename examples/inference/basic/{basic_fasth3_5090_nvfp4.yaml => basic_compacth3_rtx5090.yaml} (100%) create mode 100644 examples/inference/basic/basic_compacth3_rtx_pro6000.yaml rename examples/serving/{openai_fasth3_5090_nvfp4.yaml => openai_compacth3_rtx5090.yaml} (100%) rename examples/serving/{openai_fasth3_pro6000_nvfp4.yaml => openai_compacth3_rtx_pro6000.yaml} (100%) delete mode 100644 examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml diff --git a/examples/inference/basic/basic_fasth3_5090_nvfp4.yaml b/examples/inference/basic/basic_compacth3_rtx5090.yaml similarity index 100% rename from examples/inference/basic/basic_fasth3_5090_nvfp4.yaml rename to examples/inference/basic/basic_compacth3_rtx5090.yaml diff --git a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml new file mode 100644 index 0000000000..068cd79c89 --- /dev/null +++ b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml @@ -0,0 +1,57 @@ +# FastH3 42-block NVFP4 DiT on one RTX 5090 (sm_120), 832x480 4-step. +# +# Requires the packed ``transformer/nvfp4_weights.safetensors`` export next to +# the dense shards, and ``engine.quantization.layer_profile: h3_dit`` so the +# loader overlays those buffers instead of materializing the 32 GB GEMMs. +# Full decode uses the Comfy int8-convrot VAE when +# ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` sits beside the dense +# shards. Live preview is ``video_decode_backend: taeh3`` on the serve config. +# +# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_5090_nvfp4.yaml +generator: + model_path: /workspace/models/fasth3-42-live + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: true + vae: true + pin_cpu_memory: true + lazy_module_load: true + compile: + enabled: false + vae_enabled: false + pipeline: + experimental: + attention_backend: ATTN_QAT_INFER + h3_sequential_load: true + inference_torch_compile: false + vae_parallel_decode: false +request: + prompt: >- + A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain + peaks above a blue valley filled with thin morning mist. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_5090_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/serving/openai_fasth3_5090_nvfp4.yaml b/examples/serving/openai_compacth3_rtx5090.yaml similarity index 100% rename from examples/serving/openai_fasth3_5090_nvfp4.yaml rename to examples/serving/openai_compacth3_rtx5090.yaml diff --git a/examples/serving/openai_fasth3_pro6000_nvfp4.yaml b/examples/serving/openai_compacth3_rtx_pro6000.yaml similarity index 100% rename from examples/serving/openai_fasth3_pro6000_nvfp4.yaml rename to examples/serving/openai_compacth3_rtx_pro6000.yaml diff --git a/examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml b/examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml deleted file mode 100644 index cb9cea4228..0000000000 --- a/examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml +++ /dev/null @@ -1,57 +0,0 @@ -# FastH3 5090 live preview: packed NVFP4 DiT/encoder + TAEH3 decode. -# Quality path is openai_fasth3_5090_nvfp4.yaml (Comfy int8 VAE). This one -# skips the 10 GB video VAE so clip time can approach DiT + tiny decode. -# -# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ -# FLASHINFER_CUDA_ARCH_LIST=12.0a \ -# fastvideo serve --config examples/serving/openai_fasth3_5090_nvfp4_taeh3.yaml -generator: - model_path: /workspace/models/fasth3-42-live - engine: - num_gpus: 1 - use_fsdp_inference: false - quantization: - transformer_quant: NVFP4 - layer_profile: h3_dit - parallelism: - tp_size: 1 - sp_size: 1 - offload: - dit: false - dit_layerwise: false - text_encoder: true - image_encoder: true - vae: false - pin_cpu_memory: true - lazy_module_load: false - compile: - enabled: false - vae_enabled: false - pipeline: - workload_type: t2v - experimental: - attention_backend: ATTN_QAT_INFER - h3_sequential_load: true - inference_torch_compile: false - vae_parallel_decode: false - video_decode_backend: taeh3 - -server: - host: 127.0.0.1 - port: 8000 - output_dir: outputs/openai_fasth3_5090_nvfp4_taeh3 - served_model_name: fasth3 - -default_request: - negative_prompt: "" - sampling: - height: 480 - width: 832 - num_frames: 124 - fps: 24 - num_inference_steps: 5 - guidance_scale: 1.0 - batch_cfg: false - seed: 2026 - output: - return_frames: false diff --git a/fastvideo/entrypoints/openai/static/playground.css b/fastvideo/entrypoints/openai/static/playground.css index 9119d18252..cb7a0ee09b 100644 --- a/fastvideo/entrypoints/openai/static/playground.css +++ b/fastvideo/entrypoints/openai/static/playground.css @@ -81,9 +81,6 @@ pre { padding: 16px; font-size: 12px; background: var(--inset); border-radius: 8 video { width: 100%; max-height: 420px; } #job-status, .error { margin-block-start: 16px; font-size: 14px; } .error { color: var(--error); overflow-wrap: anywhere; } -.clip-nav { display: flex; flex-wrap: wrap; gap: 12px; margin-block-start: 16px; } -.clip-nav[hidden] { display: none !important; } -.clip-nav button { flex: 1 1 140px; } .result-actions { display: flex; flex-wrap: wrap; gap: 16px; align-items: center; margin-block-start: 16px; font-size: 14px; } .job-id { color: var(--muted); font: 12px/1.5 ui-monospace, monospace; overflow-wrap: anywhere; margin-block-start: 12px; } .history { border-block-start: 1px solid var(--line); margin-block-start: 28px; padding-block-start: 20px; } diff --git a/fastvideo/entrypoints/openai/static/playground.html b/fastvideo/entrypoints/openai/static/playground.html index 66264428eb..7dace42f86 100644 --- a/fastvideo/entrypoints/openai/static/playground.html +++ b/fastvideo/entrypoints/openai/static/playground.html @@ -18,7 +18,7 @@

    Text to video + audio

    A new prompt. The same running server.

    -

    Generate, watch, and queue the next prompt. Clips stay in Recent jobs so you can switch between them while the server keeps working.

    +

    Generate, watch, and refine. This page and your API clients share one running FastVideo server.

    Connecting to server… @@ -37,7 +37,7 @@

    Create a video

    Resolution and sampling come from the server configuration.

    -

    Queue the next prompt while a clip plays. Closing this page does not cancel a job.

    +

    Closing this page does not cancel a job.

    Use this prompt with cURL @@ -55,10 +55,6 @@

    Create a video

    No generation starts until you submit a prompt.

    -
    @@ -67,7 +63,7 @@

    Create a video

    Recent jobs

    -

    Select any clip to play it. The prompt box is left alone so you can queue the next one. History clears when the server restarts.

    +

    Shared with other clients. Job history clears when the server restarts.

      Loading recent jobs…

      diff --git a/fastvideo/entrypoints/openai/static/playground.js b/fastvideo/entrypoints/openai/static/playground.js index 2479aa957e..eeba66ddb1 100644 --- a/fastvideo/entrypoints/openai/static/playground.js +++ b/fastvideo/entrypoints/openai/static/playground.js @@ -6,11 +6,9 @@ const active = (job) => ["queued", "in_progress"].includes(job.status); let model = ""; let readyLabel = "Server ready"; - let watching = null; - let jobs = []; - let pollTimer = 0; + let busy = false; + let currentJob = null; let pollVersion = 0; - let submitting = false; const api = async (path, options = {}) => { const response = await fetch(path, { ...options, signal: AbortSignal.timeout(15000) }); @@ -18,19 +16,10 @@ if (!response.ok) throw new Error(body.error?.message || `Request failed (HTTP ${response.status}).`); return body; }; - const jobPath = (id) => `/v1/videos/${encodeURIComponent(id)}`; - const syncGenerateButton = () => { - $("generate").disabled = !model || submitting; - $("generate").textContent = jobs.some(active) ? "Queue next clip ↗" : "Generate video ↗"; - }; - const watchingIndex = () => jobs.findIndex((job) => job.id === watching?.id); - const syncClipNav = () => { - const index = watchingIndex(); - const show = jobs.length > 1 && index >= 0; - $("clip-nav").hidden = !show; - if (!show) return; - $("older-clip").disabled = index >= jobs.length - 1; - $("newer-clip").disabled = index <= 0; + const setBusy = (value) => { + busy = value; + $("generate").disabled = value || !model; + $("generate").textContent = value ? "Waiting for this job…" : "Generate video ↗"; }; const showError = (message = "") => { $("error").textContent = message; @@ -47,172 +36,137 @@ $("copy-curl").disabled = !model || !$("prompt").value.trim() || !$("seed").validity.valid; $("copy-status").textContent = ""; }; - const setVideoSource = (job) => { - const video = $("video"); - const complete = job.status === "completed"; - video.hidden = !complete; - $("empty-preview").hidden = complete; - $("download").hidden = !complete; - if (!complete) { - video.pause(); - video.removeAttribute("src"); - delete video.dataset.job; - $("empty-preview").querySelector("h3").textContent = job.status === "failed" ? "Generation failed" : "Your job is on the server"; - $("empty-preview").querySelector("p").textContent = job.status === "failed" - ? "Read the error below before trying again." - : "Queue another prompt or open a finished clip from Recent jobs. This one keeps generating."; - return; - } - $("download").href = `${jobPath(job.id)}/content`; - $("download").download = `${job.id}.mp4`; - if (video.dataset.job !== job.id) { - video.src = `${jobPath(job.id)}/content`; - video.dataset.job = job.id; + const refreshJobs = async () => { + $("refresh-jobs").disabled = true; + try { + const result = await api("/v1/videos?limit=8&order=desc"); + $("jobs").replaceChildren(); + result.data.forEach((job) => { + const item = document.createElement("li"); + const button = document.createElement("button"); + button.type = "button"; + button.setAttribute("aria-current", String(job.id === currentJob?.id)); + const prompt = document.createElement("span"); + prompt.className = "job-prompt"; + prompt.textContent = job.prompt || job.id; + const state = document.createElement("span"); + state.className = "job-state"; + state.textContent = labels[job.status] || job.status; + button.append(prompt, state); + button.addEventListener("click", () => { + if (busy) { + showError("Wait for this job, or open another playground tab to inspect a different job."); + return; + } + $("prompt").value = job.prompt || ""; + // Jobs do not report the seed; do not pair an old seed with this prompt. + $("seed").value = ""; + updateCurl(); + followJob(job.id); + }); + item.append(button); + $("jobs").append(item); + }); + $("history-status").textContent = result.data.length ? "Showing the latest jobs from this server." : "No jobs yet. Submit a prompt to start."; + } catch (error) { + $("history-status").textContent = `Could not load jobs. Check the server and select Refresh jobs. ${error.message}`; + } finally { + $("refresh-jobs").disabled = false; } }; - const renderWatching = (job) => { - watching = job; - const url = new URL(location.href); - url.searchParams.set("job", job.id); - history.replaceState({}, "", url); + const renderJob = (job) => { + currentJob = job; $("job-state").textContent = labels[job.status] || job.status; $("job-state").dataset.status = job.status; $("job-id").textContent = `Job ${job.id}`; - $("job-link").href = jobPath(job.id); + const path = `/v1/videos/${encodeURIComponent(job.id)}`; + $("job-link").href = path; $("job-link").hidden = false; - setVideoSource(job); - if (job.status === "completed") { - $("job-status").textContent = "Playing this clip. Queue the next prompt without leaving it, or pick another job from Recent jobs."; - } else if (job.status === "queued") { - $("job-status").textContent = "Queued. The server runs one generation at a time. Finished clips stay playable while you wait."; - } else if (job.status === "failed") { - $("job-status").textContent = "The server could not complete this job."; + const complete = job.status === "completed"; + $("video").hidden = !complete; + $("empty-preview").hidden = complete; + $("download").hidden = !complete; + if (complete) { + $("video").src = `${path}/content`; + $("download").href = `${path}/content`; + $("download").download = `${job.id}.mp4`; + $("job-status").textContent = "Video ready. Change the prompt and generate again without restarting the server."; } else { - $("job-status").textContent = "Generating. Open any finished clip from Recent jobs; this job keeps running."; + $("video").pause(); + $("video").removeAttribute("src"); + $("empty-preview").querySelector("h3").textContent = job.status === "failed" ? "Generation failed" : "Your job is on the server"; + $("empty-preview").querySelector("p").textContent = job.status === "failed" ? "Read the error below before trying again." : "You can keep editing the next prompt while you wait."; + $("job-status").textContent = job.status === "queued" ? "Queued. The server runs one generation at a time." : job.status === "failed" ? "The server could not complete this job." : "Generating video and audio. This page checks the job status automatically."; } if (job.status === "failed") showError(job.error?.message || "Check the server logs before submitting another job."); - syncClipNav(); - for (const button of $("jobs").querySelectorAll("button[data-job-id]")) { - if (button.dataset.jobId === job.id) button.setAttribute("aria-current", "true"); - else button.removeAttribute("aria-current"); - } - }; - const selectJob = (job) => { - showError(); - renderWatching(job); - }; - const renderHistory = () => { - $("jobs").replaceChildren(); - jobs.forEach((job) => { - const item = document.createElement("li"); - const button = document.createElement("button"); - button.type = "button"; - button.dataset.jobId = job.id; - if (job.id === watching?.id) button.setAttribute("aria-current", "true"); - const prompt = document.createElement("span"); - prompt.className = "job-prompt"; - prompt.textContent = job.prompt || job.id; - const state = document.createElement("span"); - state.className = "job-state"; - state.textContent = labels[job.status] || job.status; - button.append(prompt, state); - button.addEventListener("click", () => selectJob(job)); - item.append(button); - $("jobs").append(item); - }); - const queued = jobs.filter((job) => job.status === "queued").length; - const running = jobs.some((job) => job.status === "in_progress"); - if (!jobs.length) { - $("history-status").textContent = "No jobs yet. Submit a prompt to start."; - } else if (queued || running) { - $("history-status").textContent = running - ? `One clip generating${queued ? `, ${queued} queued` : ""}. Select any row to play or inspect it.` - : `${queued} clip${queued === 1 ? "" : "s"} queued. Select any row to play a finished clip.`; - } else { - $("history-status").textContent = "Select any clip to play it. The prompt box is left alone so you can queue the next one."; - } }; - const schedulePoll = () => { - clearTimeout(pollTimer); - if (!jobs.some(active)) return; + const followJob = async (id) => { const version = ++pollVersion; - pollTimer = setTimeout(async () => { - if (version !== pollVersion) return; - try { - await refreshJobs(); + const url = new URL(location.href); + url.searchParams.set("job", id); + history.replaceState({}, "", url); + showError(); + setBusy(true); + $("check-status").hidden = true; + const deadline = Date.now() + 30 * 60 * 1000; + let restorePrompt = !$("prompt").value.trim(); + try { + while (version === pollVersion) { + const job = await api(`/v1/videos/${encodeURIComponent(id)}`); + if (version !== pollVersion) return; $("connection").textContent = readyLabel; - $("check-status").hidden = true; - } catch (error) { - $("connection").textContent = "Job status unavailable · check the server"; - showError(`${error.message} Select Check status to reconnect. A connection error does not cancel generation.`); - $("check-status").hidden = false; + if (restorePrompt) { + $("prompt").value = job.prompt || ""; + updateCurl(); + restorePrompt = false; + } + renderJob(job); + if (!active(job)) break; + if (Date.now() >= deadline) throw new Error("Stopped checking after 30 minutes. The job may still be running."); + await new Promise((resolve) => setTimeout(resolve, 2000)); } - }, 2000); - }; - const refreshJobs = async () => { - $("refresh-jobs").disabled = true; - try { - const result = await api("/v1/videos?limit=32&order=desc"); - jobs = result.data; - const latestWatching = jobs.find((job) => job.id === watching?.id); - if (latestWatching) watching = latestWatching; - renderHistory(); - if (latestWatching) renderWatching(latestWatching); - syncGenerateButton(); - syncClipNav(); - schedulePoll(); } catch (error) { - $("history-status").textContent = `Could not load jobs. Check the server and select Refresh jobs. ${error.message}`; - throw error; + if (version !== pollVersion) return; + $("connection").textContent = "Job status unavailable · check the server"; + showError(`${error.message} Select Check status to reconnect. A connection error does not cancel generation. If the job is missing, the server may have restarted or the job was deleted.`); + $("check-status").hidden = false; } finally { - $("refresh-jobs").disabled = false; + if (version === pollVersion) { + setBusy(false); + await refreshJobs(); + } } }; $("generate-form").addEventListener("submit", async (event) => { event.preventDefault(); - if (!model || submitting) return; + if (busy || !model) return; if (!$("prompt").value.trim()) { showError("Write a prompt before generating a video."); $("prompt").focus(); return; } showError(); - submitting = true; - syncGenerateButton(); - const keepClip = watching?.status === "completed"; - $("job-status").textContent = keepClip ? "Queuing the next clip…" : "Submitting the prompt…"; + setBusy(true); + $("job-status").textContent = "Submitting the prompt…"; try { const job = await api("/v1/videos", { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify(payload()), }); - await refreshJobs(); - if (keepClip) { - $("job-status").textContent = "Queued. This clip keeps playing while the server generates the next one."; - } else { - selectJob(job); - } + renderJob(job); + await followJob(job.id); } catch (error) { showError(`${error.message} Check Recent jobs before submitting again; the server may have received the prompt. Submissions are never retried automatically.`); $("job-status").textContent = "Could not confirm the submission."; - await refreshJobs().catch(() => undefined); - } finally { - submitting = false; - syncGenerateButton(); + setBusy(false); + await refreshJobs(); } }); $("prompt").addEventListener("input", updateCurl); $("seed").addEventListener("input", updateCurl); - $("refresh-jobs").addEventListener("click", () => refreshJobs().catch(() => undefined)); - $("older-clip").addEventListener("click", () => { - const index = watchingIndex(); - if (index >= 0 && index < jobs.length - 1) selectJob(jobs[index + 1]); - }); - $("newer-clip").addEventListener("click", () => { - const index = watchingIndex(); - if (index > 0) selectJob(jobs[index - 1]); - }); + $("refresh-jobs").addEventListener("click", refreshJobs); $("check-status").addEventListener("click", () => { - refreshJobs().catch(() => undefined); + const id = new URL(location.href).searchParams.get("job"); + if (id) followJob(id); }); $("video").addEventListener("error", () => { if (!$("video").hasAttribute("src")) return; @@ -235,7 +189,7 @@ $("connection").textContent = readyLabel; $("lifetime").textContent = config.runtime === "mlx" ? "MLX keeps the server and prompt cache available, but releases model components between phases to limit unified-memory use. Closing this page does not cancel a job." - : "The model stays loaded until you stop the server. Queue the next prompt while a clip plays; closing this page does not cancel a job."; + : "The model stays loaded until you stop the server. Closing this page does not cancel a job."; $("model").textContent = model; const d = config.defaults; const facts = []; @@ -244,18 +198,11 @@ if (d.fps) facts.push(`${d.fps} fps`); if (d.seed != null) $("seed").placeholder = String(d.seed); $("settings").textContent = facts.length ? `Server defaults · ${facts.join(" · ")}` : "Resolution and sampling come from the server configuration."; - syncGenerateButton(); + setBusy(false); updateCurl(); - await refreshJobs().catch(() => undefined); const id = new URL(location.href).searchParams.get("job"); - const linked = jobs.find((job) => job.id === id); - if (linked) { - if (!$("prompt").value.trim() && linked.prompt) { - $("prompt").value = linked.prompt; - updateCurl(); - } - selectJob(linked); - } + if (id) await followJob(id); + else await refreshJobs(); } catch (error) { $("connection").textContent = "Server unavailable"; $("history-status").textContent = "Start the H3 server, then reload this page."; From a0682348f7860576b6ddbc65116f1a642d130386 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 14:53:24 -0700 Subject: [PATCH 003/107] [feat]: document CompactH3 NVFP4 on RTX 5090 and RTX PRO 6000 Add H3 cookbook recipes for the 42-block checkpoint, reject ConvRot overlays whose group size cannot rotate activations, and skip FSDP2 modules that mix dense parameters with DTensors. (cherry picked from commit 7852505a6e6c41739db670aaf619155bd5e3c664) --- docs/assets/cookbook-recipes.json | 72 ++++++++++++++++++- docs/cookbook/cosmos.md | 2 +- docs/cookbook/flux.md | 2 +- docs/cookbook/glm-image.md | 2 +- docs/cookbook/hunyuan.md | 2 +- docs/cookbook/kandinsky5.md | 2 +- docs/cookbook/longcat.md | 2 +- docs/cookbook/ltx.md | 2 +- docs/cookbook/matrix-game.md | 2 +- docs/cookbook/minimax-h3.md | 22 ++++-- docs/cookbook/mmaudio.md | 2 +- docs/cookbook/stable-audio.md | 2 +- docs/cookbook/stable-diffusion.md | 2 +- docs/cookbook/turbodiffusion.md | 2 +- docs/cookbook/wan.md | 2 +- docs/cookbook/z-image.md | 2 +- .../basic/basic_compacth3_rtx5090.yaml | 28 ++++---- .../basic/basic_compacth3_rtx_pro6000.yaml | 37 +++++----- .../serving/openai_compacth3_rtx5090.yaml | 18 ++--- .../serving/openai_compacth3_rtx_pro6000.yaml | 22 +++--- .../models/vaes/minimax_h3_int8_convrot.py | 7 +- .../basic/minimax_h3/minimax_h3_pipeline.py | 3 +- .../test_minimax_h3_sequential_start.py | 63 ++++++++++++++++ .../vaes/test_minimax_h3_int8_convrot.py | 14 ++++ 24 files changed, 242 insertions(+), 72 deletions(-) diff --git a/docs/assets/cookbook-recipes.json b/docs/assets/cookbook-recipes.json index 8875361dcc..242ad9211b 100644 --- a/docs/assets/cookbook-recipes.json +++ b/docs/assets/cookbook-recipes.json @@ -1,5 +1,5 @@ { - "version": 11, + "version": 12, "recipes": [ { "id": "fastwan21-t2v", @@ -578,6 +578,76 @@ "Height, width, frames, and steps in the YAML are examples. Edit them or pass CLI flags. See docs/getting_started/installation/spark_pair.md." ] }, + { + "id": "compacth3-rtx5090", + "group": "compacth3-rtx5090", + "group_label": "CompactH3 on RTX 5090", + "group_task": "4-step 42-block text to video + audio", + "family": "minimax_h3", + "stage": "inference", + "task": "Few-step text to video (with audio)", + "label": "CompactH3 NVFP4 on RTX 5090", + "summary": "Run the 42-block CompactH3 NVFP4 DiT with the NVFP4 Qwen3-VL encoder, Comfy int8-convrot VAE, and SageAttention3 FP4 on one 32 GB RTX 5090. Sequential load parks the encoder in pinned host RAM.", + "model": "./CompactH3", + "source": "examples/inference/basic/basic_compacth3_rtx5090.yaml", + "serving": { + "source": "examples/serving/openai_compacth3_rtx5090.yaml", + "install": "UV_TORCH_BACKEND=cu128 uv pip install -e \".[fasth3]\"", + "prepare": "hf download FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree --local-dir ./CompactH3 --include model_index.json --include \"tokenizer/**\" --include \"processor/**\" --include \"scheduler/**\" --include \"audio_scheduler/**\" --include \"audio_vae/**\" --include \"vae/**\"\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-bf16 --local-dir ./CompactH3/transformer\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-nvfp4 --local-dir ./CompactH3/transformer --include nvfp4_weights.safetensors\nhf download KyleNeverGivesUp/FastH3-text-encoder-nvfp4 --local-dir ./CompactH3/text_encoder\nhf download Comfy-Org/MiniMax-H3 --local-dir ./Comfy-MiniMax-H3 --include vae/minimax_h3_video_vae_int8_convrot.safetensors\ncp ./Comfy-MiniMax-H3/vae/minimax_h3_video_vae_int8_convrot.safetensors ./CompactH3/vae/", + "env": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a" + }, + "command": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 fastvideo generate --config examples/inference/basic/basic_compacth3_rtx5090.yaml", + "gpu_types": ["NVIDIA"], + "hardware": { + "platform": "cuda", + "gpu_count": 1, + "evidence": "source-configured" + }, + "evidence": "Source-backed", + "expected_artifact": "MP4 under outputs/compacth3_rtx5090/", + "modes": ["T2VA", "CompactH3 NVFP4", "RTX 5090"], + "limitations": [ + "Assemble ./CompactH3 before running. The DiT NVFP4 export and encoder snapshot are gated; run huggingface-cli login and accept each repo license.", + "32 GB cannot keep the NVFP4 encoder and DiT on the GPU together. Keep h3_sequential_load on and lazy_module_load off.", + "Blackwell sm_120 needs ATTN_QAT_INFER, FLASHINFER_CUDA_ARCH_LIST=12.0a, and a CUDA 12.8 PyTorch wheel.", + "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768; 1344x768 on 32 GB is unmeasured." + ] + }, + { + "id": "compacth3-rtx-pro6000", + "group": "compacth3-rtx-pro6000", + "group_label": "CompactH3 on RTX PRO 6000", + "group_task": "4-step 42-block text to video + audio", + "family": "minimax_h3", + "stage": "inference", + "task": "Few-step text to video (with audio)", + "label": "CompactH3 NVFP4 on RTX PRO 6000 Blackwell", + "summary": "Run CompactH3 NVFP4 with the encoder, DiT, and int8-convrot VAE resident on one 96 GB RTX PRO 6000 Blackwell. The checked-in example is 1344x768 and 124 frames (5.17 s).", + "model": "./CompactH3", + "source": "examples/inference/basic/basic_compacth3_rtx_pro6000.yaml", + "serving": { + "source": "examples/serving/openai_compacth3_rtx_pro6000.yaml", + "install": "UV_TORCH_BACKEND=cu128 uv pip install -e \".[fasth3]\"", + "prepare": "hf download FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree --local-dir ./CompactH3 --include model_index.json --include \"tokenizer/**\" --include \"processor/**\" --include \"scheduler/**\" --include \"audio_scheduler/**\" --include \"audio_vae/**\" --include \"vae/**\"\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-bf16 --local-dir ./CompactH3/transformer\nhf download aryan5v/FastH3-20B-42block-dmd2-ckpt1400-nvfp4 --local-dir ./CompactH3/transformer --include nvfp4_weights.safetensors\nhf download KyleNeverGivesUp/FastH3-text-encoder-nvfp4 --local-dir ./CompactH3/text_encoder\nhf download Comfy-Org/MiniMax-H3 --local-dir ./Comfy-MiniMax-H3 --include vae/minimax_h3_video_vae_int8_convrot.safetensors\ncp ./Comfy-MiniMax-H3/vae/minimax_h3_video_vae_int8_convrot.safetensors ./CompactH3/vae/", + "env": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a" + }, + "command": "FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 fastvideo generate --config examples/inference/basic/basic_compacth3_rtx_pro6000.yaml", + "gpu_types": ["NVIDIA"], + "hardware": { + "platform": "cuda", + "gpu_count": 1, + "evidence": "source-configured" + }, + "evidence": "Source-backed", + "expected_artifact": "MP4 under outputs/compacth3_rtx_pro6000/", + "modes": ["T2VA", "CompactH3 NVFP4", "RTX PRO 6000"], + "limitations": [ + "Assemble ./CompactH3 before running. The DiT NVFP4 export and encoder snapshot are gated; run huggingface-cli login and accept each repo license.", + "96 GB keeps the encoder, DiT, and VAE on GPU. Do not enable h3_sequential_load or lazy_module_load on this box.", + "Blackwell sm_120 needs ATTN_QAT_INFER, FLASHINFER_CUDA_ARCH_LIST=12.0a, and a CUDA 12.8 PyTorch wheel.", + "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768." + ] + }, { "id": "fasth3-8step-v2-cuda", "group": "fasth3-8step-v2", diff --git a/docs/cookbook/cosmos.md b/docs/cookbook/cosmos.md index 2ee864e5b7..fbb813fae6 100644 --- a/docs/cookbook/cosmos.md +++ b/docs/cookbook/cosmos.md @@ -5,7 +5,7 @@ hide: # Cosmos recipes -
      +
      All model families
      diff --git a/docs/cookbook/flux.md b/docs/cookbook/flux.md index b8bbb59052..f2f159aed7 100644 --- a/docs/cookbook/flux.md +++ b/docs/cookbook/flux.md @@ -5,7 +5,7 @@ hide: # FLUX recipes -
      +
      All model families
      diff --git a/docs/cookbook/glm-image.md b/docs/cookbook/glm-image.md index b4a2c39673..ed88351a17 100644 --- a/docs/cookbook/glm-image.md +++ b/docs/cookbook/glm-image.md @@ -5,7 +5,7 @@ hide: # GLM-Image recipes -
      +
      All model families
      diff --git a/docs/cookbook/hunyuan.md b/docs/cookbook/hunyuan.md index 8408dad626..66acf7f288 100644 --- a/docs/cookbook/hunyuan.md +++ b/docs/cookbook/hunyuan.md @@ -5,7 +5,7 @@ hide: # Hunyuan recipes -
      +
      All model families
      diff --git a/docs/cookbook/kandinsky5.md b/docs/cookbook/kandinsky5.md index e07993f24d..119296eae7 100644 --- a/docs/cookbook/kandinsky5.md +++ b/docs/cookbook/kandinsky5.md @@ -5,7 +5,7 @@ hide: # Kandinsky 5 recipes -
      +
      All model families
      diff --git a/docs/cookbook/longcat.md b/docs/cookbook/longcat.md index f898060525..ce2a81574a 100644 --- a/docs/cookbook/longcat.md +++ b/docs/cookbook/longcat.md @@ -5,7 +5,7 @@ hide: # LongCat recipes -
      +
      All model families
      diff --git a/docs/cookbook/ltx.md b/docs/cookbook/ltx.md index de1bf1c982..a9354f30da 100644 --- a/docs/cookbook/ltx.md +++ b/docs/cookbook/ltx.md @@ -5,7 +5,7 @@ hide: # LTX recipes -
      +
      All model families
      diff --git a/docs/cookbook/matrix-game.md b/docs/cookbook/matrix-game.md index 8bcb544f55..26160e8650 100644 --- a/docs/cookbook/matrix-game.md +++ b/docs/cookbook/matrix-game.md @@ -5,7 +5,7 @@ hide: # Matrix Game recipes -
      +
      All model families
      diff --git a/docs/cookbook/minimax-h3.md b/docs/cookbook/minimax-h3.md index 8f3b4f33b2..b4cc9224f8 100644 --- a/docs/cookbook/minimax-h3.md +++ b/docs/cookbook/minimax-h3.md @@ -10,8 +10,10 @@ launch. Some Hub repo names still say Preview. That name is historical. V1 is a full model, not a demo. **V2** is the eight-step checkpoint. More forwards is why V2 is the higher-quality FastH3. The V2 schedule contract is in [FastH3 distilled checkpoint schedules](../inference/fasth3-distilled.md). +**CompactH3** is the 42-block 20B NVFP4 H3 checkpoint for one Blackwell GPU +(RTX 5090 or RTX PRO 6000). The FastH3 V1 and V2 recipes are unchanged. -
      +
      All model families
      @@ -21,8 +23,8 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in

      Primary focus · Inference

      MiniMax H3 recipes

      -

      Generate video and audio with H3. Run a server on CUDA, one DGX Spark, or Apple Silicon MLX to iterate on prompts, or call the pipeline directly from Python.

      - 9 maintained recipes +

      Generate video and audio with H3. Run a server on CUDA, one Blackwell GPU, one DGX Spark, or Apple Silicon MLX to iterate on prompts, or call the pipeline directly from Python.

      + 11 maintained recipes
      @@ -47,7 +49,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in

      Supported modes

      CUDA covers T2VA, FL2VA, and Ref2VA on the full checkpoint, plus FastH3 - V1 and FastH3 V2. FastH3 V1 also has a DGX Spark runtime with + V1 and FastH3 V2, plus CompactH3 NVFP4 on one Blackwell GPU. FastH3 V1 also has a DGX Spark runtime with a 1-Spark or 2-Spark device row. MLX is T2VA only: V1 and V2. Temporal --fast, spatial --fast-spatial, and opt-in VSA are flags on the same MLX script, not extra recipes. @@ -64,7 +66,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in T2VA - Full H3, FastH3 V1, FastH3 LoRA, FastH3 V2 + Full H3, FastH3 V1, FastH3 LoRA, FastH3 V2, CompactH3 NVFP4 FastH3 V1 or FastH3 V2 after a local DiT conversion @@ -102,6 +104,11 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in FastH3 V1 on one GB10, or two Sparks with Ray sequence parallel (sp_size=2) over QSFP RoCE. Select NVIDIA DGX Spark, then 1 Spark or 2 Sparks. Not wired + + CompactH3 NVFP4 + 42-block 20B checkpoint on one RTX 5090 (32 GB, sequential encoder offload) or one RTX PRO 6000 Blackwell (96 GB, encoder+DiT+VAE resident). SageAttention3 FP4, packed NVFP4 DiT, Comfy int8-convrot VAE. + Not wired +

      @@ -110,7 +117,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in

      Pick an H3 recipe and runtime

      -

      Choose the result you want, then use a maintained CUDA, DGX Spark, or MLX path. +

      Choose the result you want, then use a maintained CUDA, Blackwell, DGX Spark, or MLX path. Device claims stay tied to checked-in sources and recorded runs.

      @@ -275,7 +282,8 @@ cd FastVideo
    • The MLX source runtime supports T2VA, optional temporal --fast, optional spatial --fast-spatial, and opt-in VSA on --include-vsa checkpoints. FastH3 V2 MLX converts with --include-vsa and runs eight forwards. FL2VA, Ref2VA, and two-pass refinement are not wired.
    • GPU count and VAE decode backend are configurable in the builder above for FastH3 CUDA recipes. Only the value shown by default has a recorded run; other supported values are unmeasured here.
    • DGX Spark is a runtime on FastH3 V1, not a separate family card. Select NVIDIA DGX Spark, then 1 Spark or 2 Sparks. The CUDA GPU-count knob does not apply to Spark.
    • -
    • GB10 has no FA4 / sm_100a VSA kernel. Keep FASTVIDEO_FA4=0 and FASTVIDEO_VSA_SM100A=0. Legal num_frames values are 17n+5, capped at 362 (15.08 s). A 345-frame request on one Spark can OOM.
    • +
    • CompactH3 NVFP4 is one Blackwell GPU. RTX 5090 (32 GB) parks the encoder in pinned host RAM. RTX PRO 6000 Blackwell (96 GB) keeps encoder, DiT, and VAE resident. Keep FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER, FASTVIDEO_FA4=0, FASTVIDEO_VSA_SM100A=0, and FLASHINFER_CUDA_ARCH_LIST=12.0a.
    • +
    • GB10 has no FA4 / sm_100a VSA kernel. Keep FASTVIDEO_FA4=0 and FASTVIDEO_VSA_SM100A=0. Legal num_frames values are 17n+5, capped at 362 (15.08 s). A 345-frame request on one Spark can OOM. Native 16:9 sizes include 832×480 and 1344×768.
    • Gated or missing checkpoints: run huggingface-cli login and confirm you accepted the model's license on Hugging Face.
    • diff --git a/docs/cookbook/mmaudio.md b/docs/cookbook/mmaudio.md index 8bdd93515a..5e127c9b2b 100644 --- a/docs/cookbook/mmaudio.md +++ b/docs/cookbook/mmaudio.md @@ -5,7 +5,7 @@ hide: # MMAudio recipes -
      +
      All model families
      diff --git a/docs/cookbook/stable-audio.md b/docs/cookbook/stable-audio.md index a2ca588ef9..20dc7bfd7a 100644 --- a/docs/cookbook/stable-audio.md +++ b/docs/cookbook/stable-audio.md @@ -5,7 +5,7 @@ hide: # Stable Audio recipes -
      +
      All model families
      diff --git a/docs/cookbook/stable-diffusion.md b/docs/cookbook/stable-diffusion.md index a3c8c62ab2..07c87fdb59 100644 --- a/docs/cookbook/stable-diffusion.md +++ b/docs/cookbook/stable-diffusion.md @@ -5,7 +5,7 @@ hide: # Stable Diffusion recipes -
      +
      All model families
      diff --git a/docs/cookbook/turbodiffusion.md b/docs/cookbook/turbodiffusion.md index 92b259fec4..14820b5635 100644 --- a/docs/cookbook/turbodiffusion.md +++ b/docs/cookbook/turbodiffusion.md @@ -5,7 +5,7 @@ hide: # TurboDiffusion recipes -
      +
      All model families
      diff --git a/docs/cookbook/wan.md b/docs/cookbook/wan.md index 4554e09660..1d72368078 100644 --- a/docs/cookbook/wan.md +++ b/docs/cookbook/wan.md @@ -5,7 +5,7 @@ hide: # Wan recipes -
      +
      All model families
      diff --git a/docs/cookbook/z-image.md b/docs/cookbook/z-image.md index 8f4301f647..b53c79ef16 100644 --- a/docs/cookbook/z-image.md +++ b/docs/cookbook/z-image.md @@ -5,7 +5,7 @@ hide: # Z-Image recipes -
      +
      All model families
      diff --git a/examples/inference/basic/basic_compacth3_rtx5090.yaml b/examples/inference/basic/basic_compacth3_rtx5090.yaml index 068cd79c89..f17114b240 100644 --- a/examples/inference/basic/basic_compacth3_rtx5090.yaml +++ b/examples/inference/basic/basic_compacth3_rtx5090.yaml @@ -1,17 +1,18 @@ -# FastH3 42-block NVFP4 DiT on one RTX 5090 (sm_120), 832x480 4-step. +# CompactH3 (42-block 20B NVFP4) on one RTX 5090 (sm_120, 32 GB). +# Packed ``transformer/nvfp4_weights.safetensors`` plus +# ``engine.quantization.layer_profile: h3_dit``. Full decode uses the Comfy +# int8-convrot VAE when ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` +# sits beside the dense VAE shards. # -# Requires the packed ``transformer/nvfp4_weights.safetensors`` export next to -# the dense shards, and ``engine.quantization.layer_profile: h3_dit`` so the -# loader overlays those buffers instead of materializing the 32 GB GEMMs. -# Full decode uses the Comfy int8-convrot VAE when -# ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` sits beside the dense -# shards. Live preview is ``video_decode_backend: taeh3`` on the serve config. +# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). +# Sequential load parks the encoder in pinned host RAM; do not enable +# lazy_module_load. Legal H3 frame counts are 17n+5, max 362 (15.08 s). # # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ -# FASTVIDEO_STAGE_LOGGING=1 \ -# fastvideo generate --config examples/inference/basic/basic_fasth3_5090_nvfp4.yaml +# FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_compacth3_rtx5090.yaml generator: - model_path: /workspace/models/fasth3-42-live + model_path: ./CompactH3 engine: num_gpus: 1 use_fsdp_inference: false @@ -25,9 +26,9 @@ generator: dit: false dit_layerwise: false text_encoder: true - vae: true + vae: false pin_cpu_memory: true - lazy_module_load: true + lazy_module_load: false compile: enabled: false vae_enabled: false @@ -37,6 +38,7 @@ generator: h3_sequential_load: true inference_torch_compile: false vae_parallel_decode: false + video_decode_backend: h3-vae request: prompt: >- A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain @@ -52,6 +54,6 @@ request: guidance_scale: 1.0 batch_cfg: false output: - output_path: outputs/fasth3_5090_nvfp4/ + output_path: outputs/compacth3_rtx5090/ save_video: true return_frames: false diff --git a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml index 068cd79c89..4cbf422a9e 100644 --- a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml +++ b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml @@ -1,17 +1,19 @@ -# FastH3 42-block NVFP4 DiT on one RTX 5090 (sm_120), 832x480 4-step. +# CompactH3 (42-block 20B NVFP4) on one RTX PRO 6000 Blackwell (sm_120, 96 GB). +# Packed ``transformer/nvfp4_weights.safetensors`` plus +# ``engine.quantization.layer_profile: h3_dit``. Full decode uses the Comfy +# int8-convrot VAE when ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` +# sits beside the dense VAE shards. # -# Requires the packed ``transformer/nvfp4_weights.safetensors`` export next to -# the dense shards, and ``engine.quantization.layer_profile: h3_dit`` so the -# loader overlays those buffers instead of materializing the 32 GB GEMMs. -# Full decode uses the Comfy int8-convrot VAE when -# ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` sits beside the dense -# shards. Live preview is ``video_decode_backend: taeh3`` on the serve config. +# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). +# 96 GB keeps the encoder, DiT, and VAE resident. request.sampling is an +# example: native 16:9 sizes include 832x480 and 1344x768. Legal H3 frame +# counts are 17n+5, max 362 (15.08 s). # # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ -# FASTVIDEO_STAGE_LOGGING=1 \ -# fastvideo generate --config examples/inference/basic/basic_fasth3_5090_nvfp4.yaml +# FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_compacth3_rtx_pro6000.yaml generator: - model_path: /workspace/models/fasth3-42-live + model_path: ./CompactH3 engine: num_gpus: 1 use_fsdp_inference: false @@ -24,19 +26,20 @@ generator: offload: dit: false dit_layerwise: false - text_encoder: true - vae: true + text_encoder: false + vae: false pin_cpu_memory: true - lazy_module_load: true + lazy_module_load: false compile: enabled: false vae_enabled: false pipeline: experimental: attention_backend: ATTN_QAT_INFER - h3_sequential_load: true + h3_sequential_load: false inference_torch_compile: false vae_parallel_decode: false + video_decode_backend: h3-vae request: prompt: >- A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain @@ -44,14 +47,14 @@ request: negative_prompt: "" sampling: seed: 2026 - height: 480 - width: 832 + height: 768 + width: 1344 num_frames: 124 fps: 24 num_inference_steps: 5 guidance_scale: 1.0 batch_cfg: false output: - output_path: outputs/fasth3_5090_nvfp4/ + output_path: outputs/compacth3_rtx_pro6000/ save_video: true return_frames: false diff --git a/examples/serving/openai_compacth3_rtx5090.yaml b/examples/serving/openai_compacth3_rtx5090.yaml index 63339f9290..252df0b522 100644 --- a/examples/serving/openai_compacth3_rtx5090.yaml +++ b/examples/serving/openai_compacth3_rtx5090.yaml @@ -1,13 +1,15 @@ -# OpenAI-compatible FastH3 server for one RTX 5090 (sm_120). -# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE (quality decode). -# lazy_module_load is off so DiT/VAE stay resident; sequential load keeps the -# encoder in host RAM. TAEH3 live preview: openai_fasth3_5090_nvfp4_taeh3.yaml +# CompactH3 (42-block 20B NVFP4) on one RTX 5090 (sm_120, 32 GB). +# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. +# 32 GB cannot co-reside the encoder and DiT; sequential load parks the +# encoder in pinned host RAM between phases. Do not enable lazy_module_load. +# +# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). # # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a \ -# fastvideo serve --config examples/serving/openai_fasth3_5090_nvfp4.yaml +# fastvideo serve --config examples/serving/openai_compacth3_rtx5090.yaml generator: - model_path: /workspace/models/fasth3-42-live + model_path: ./CompactH3 engine: num_gpus: 1 use_fsdp_inference: false @@ -40,8 +42,8 @@ generator: server: host: 127.0.0.1 port: 8000 - output_dir: outputs/openai_fasth3_5090_nvfp4 - served_model_name: fasth3 + output_dir: outputs/openai_compacth3_rtx5090 + served_model_name: compacth3 default_request: negative_prompt: "" diff --git a/examples/serving/openai_compacth3_rtx_pro6000.yaml b/examples/serving/openai_compacth3_rtx_pro6000.yaml index 698c083a7c..844b98faad 100644 --- a/examples/serving/openai_compacth3_rtx_pro6000.yaml +++ b/examples/serving/openai_compacth3_rtx_pro6000.yaml @@ -1,12 +1,16 @@ -# OpenAI-compatible FastH3 server for one RTX PRO 6000 Blackwell (sm_120, 96 GB). -# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. 96 GB holds encoder, -# DiT, and VAE together, so sequential ping-pong stays off. +# CompactH3 (42-block 20B NVFP4) on one RTX PRO 6000 Blackwell (sm_120, 96 GB). +# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. 96 GB holds the +# encoder, DiT, and VAE resident, so sequential ping-pong stays off. +# +# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). +# default_request.sampling is an example. Native 16:9 sizes include 832x480 +# and 1344x768. Legal H3 frame counts are 17n+5, max 362 (15.08 s). # # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a \ -# fastvideo serve --config examples/serving/openai_fasth3_pro6000_nvfp4.yaml +# fastvideo serve --config examples/serving/openai_compacth3_rtx_pro6000.yaml generator: - model_path: /workspace/models/fasth3-42-live + model_path: ./CompactH3 engine: num_gpus: 1 use_fsdp_inference: false @@ -39,14 +43,14 @@ generator: server: host: 127.0.0.1 port: 8000 - output_dir: outputs/openai_fasth3_pro6000_nvfp4 - served_model_name: fasth3 + output_dir: outputs/openai_compacth3_rtx_pro6000 + served_model_name: compacth3 default_request: negative_prompt: "" sampling: - height: 480 - width: 832 + height: 768 + width: 1344 num_frames: 124 fps: 24 num_inference_steps: 5 diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py index 0c2885a9fa..bded75dee8 100644 --- a/fastvideo/models/vaes/minimax_h3_int8_convrot.py +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -96,6 +96,11 @@ def __init__( group_size: int, ) -> None: super().__init__() + if convrot and in_features % group_size != 0: + raise ValueError( + f"ConvRot weights were rotated with group_size {group_size}, " + f"but in_features {in_features} is not divisible by that group" + ) self.in_features = in_features self.out_features = out_features self.convrot = convrot @@ -163,7 +168,7 @@ def _int8_linear_from_tensors( weight.shape[1], weight.shape[0], bias=bias is not None, - convrot=convrot and weight.shape[1] % group_size == 0, + convrot=convrot, group_size=group_size, ) layer.weight.copy_(weight) diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 1b78b4645f..6ab034ccf9 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -83,8 +83,7 @@ def _module_has_dtensor_params(module: Any) -> bool: parameters = getattr(module, "parameters", None) if not callable(parameters): return False - first = next(parameters(), None) - return first is not None and isinstance(first, DTensor) + return any(isinstance(parameter, DTensor) for parameter in parameters()) @dataclass(frozen=True) diff --git a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index c47c044b08..c513ff1239 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -207,6 +207,69 @@ def fake_add_denoise(*, ref2va: bool) -> None: assert transformer.moved_to[-1] == torch.device("cpu") +def test_sequential_skips_host_offload_when_dense_params_precede_dtensors(monkeypatch) -> None: + class _FakeDTensor: + pass + + monkeypatch.setattr( + "fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.DTensor", + _FakeDTensor, + ) + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + loads: list[list[str]] = [] + + def _mixed_stub(name: str) -> SimpleNamespace: + module = _stub_module(name) + + def parameters(): + yield torch.zeros(1) + yield _FakeDTensor() + + module.parameters = parameters + return module + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + if name in modules: + continue + modules[name] = _mixed_stub(name) if name == "text_encoder" else _stub_module(name) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=True, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + + passthrough = lambda batch, _args: batch + monkeypatch.setattr(pipeline._stage_name_mapping["input_preparation_stage"], "forward", passthrough) + monkeypatch.setattr(pipeline._stage_name_mapping["conditioning_stage"], "forward", passthrough) + original_add_denoise = pipeline._add_denoise_stages + + def fake_add_denoise(*, ref2va: bool) -> None: + original_add_denoise(ref2va=ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + monkeypatch.setattr(pipeline._stage_name_mapping[name], "forward", passthrough) + + monkeypatch.setattr(pipeline, "_add_denoise_stages", fake_add_denoise) + encoder = pipeline.get_module("text_encoder") + pipeline.forward(ForwardBatch(data_type="video", prompt="alpine dancer"), args) + assert encoder.moved_to == [] + + def test_unified_memory_sequential_deletes_encoder_and_reloads(monkeypatch) -> None: events: list = [] _patch_pipeline_construction(monkeypatch, events, unified_memory=True) diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py index 8c39d02695..8a4984ed7c 100644 --- a/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py @@ -11,6 +11,7 @@ from fastvideo.models.vaes.minimax_h3_int8_convrot import ( Int8ConvRotLinear, + _int8_linear_from_tensors, dense_vae_safetensors, overlay_minimax_h3_int8_convrot_decoder, parse_comfy_quant_marker, @@ -44,6 +45,19 @@ def test_int8_gemm_scales_in_float32_not_fp16() -> None: torch.testing.assert_close(scaled, acc.float() * x_scale.float() * weight_scale.t().float()) +def test_int8_linear_rejects_incompatible_convrot_group() -> None: + marker = {"format": "int8_tensorwise", "convrot": True, "convrot_groupsize": 256} + weight = torch.ones(8, 8, dtype=torch.int8) + scale = torch.ones(8, 1) + try: + _int8_linear_from_tensors(weight, scale, None, marker) + except ValueError as error: + assert "group_size 256" in str(error) + assert "in_features 8" in str(error) + else: + raise AssertionError("incompatible ConvRot overlay must be rejected") + + def test_int8_linear_matches_dequantized_matmul() -> None: torch.manual_seed(1) layer = Int8ConvRotLinear(256, 32, bias=True, convrot=True, group_size=256) From d224e527f245ecd447d900848e4763e444290c04 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 15:03:34 -0700 Subject: [PATCH 004/107] [bugfix]: refuse FSDP NVFP4 retain and drop restating comments (cherry picked from commit 19a408faa5d7323d4407c80e4d3ab566406a6146) --- fastvideo/layers/quantization/nvfp4_config.py | 12 ++--------- .../models/vaes/minimax_h3_int8_convrot.py | 3 --- .../basic/minimax_h3/minimax_h3_pipeline.py | 1 - .../ops/quantization/test_nvfp4_purge.py | 21 +++++++++++++++++++ 4 files changed, 23 insertions(+), 14 deletions(-) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 60c81c72a0..3c84f1be37 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -111,12 +111,10 @@ def is_minimax_h3_nvfp4_dit_linear_prefix(prefix: str) -> bool: def is_minimax_h3_nvfp4_dit_export_path(path: str) -> bool: - """Return whether *path* is the packed MiniMax-H3 DiT NVFP4 export file.""" return os.path.basename(path) == H3_NVFP4_DIT_EXPORT_FILENAME def find_minimax_h3_nvfp4_dit_export(weight_paths: list[str]) -> str | None: - """Locate ``nvfp4_weights.safetensors`` next to a transformer shard list.""" seen: list[str] = [] for path in weight_paths: if is_minimax_h3_nvfp4_dit_export_path(path) and os.path.isfile(path): @@ -490,8 +488,6 @@ def __init__(self, layer_profile: str = "refine", retain_original_weights: bool if layer_profile not in ("base", "refine", "h3_dit"): raise ValueError("NVFP4Config.layer_profile must be one of 'base', 'refine', or 'h3_dit', " f"got {layer_profile!r}") - # ``base`` / ``refine``: LTX-2 stage-1 vs stage-2 layer sets. - # ``h3_dit``: packed MiniMax-H3 attention+FFN export. self.layer_profile = layer_profile # Original bf16 ``layer.weight`` retention after FP4 conversion. # Default (None/False): purge the purgeable originals -- every @@ -528,9 +524,6 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str): if not isinstance(layer, LinearBase): return None - # LTX-2 switches its active subset by stage at runtime. MiniMax-H3 - # uses the fixed main-transformer FFN set unless ``h3_dit`` selects - # the packed attention+FFN export. if self.layer_profile == "h3_dit": tagged = is_minimax_h3_nvfp4_dit_linear_prefix(prefix) else: @@ -593,9 +586,8 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: if retain: retained += 1 elif isinstance(weight, DTensor): - # ponytail: purging FSDP-sharded originals needs per-shard - # resharding bookkeeping; skip until a sharded deploy needs it. - retained += 1 + raise RuntimeError("NVFP4 cannot purge FSDP-sharded bf16 weights. Use a packed NVFP4 " + "export, or convert without FSDP sharding.") else: purged_bytes += weight.numel() * weight.element_size() purged += 1 diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py index bded75dee8..649541fd19 100644 --- a/fastvideo/models/vaes/minimax_h3_int8_convrot.py +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -63,7 +63,6 @@ def regular_hadamard(size: int, *, device: torch.device, dtype: torch.dtype) -> def rotate_activation(x: torch.Tensor, group_size: int) -> torch.Tensor: - """Apply the block-diagonal ConvRot Hadamard to the last dimension of ``x``.""" features = x.shape[-1] if features % group_size: raise ValueError(f"features {features} are not divisible by convrot group_size {group_size}") @@ -303,7 +302,6 @@ def overlay_minimax_h3_int8_convrot_decoder(vae: nn.Module, checkpoint_path: str def is_int8_convrot_vae_path(path: str | Path) -> bool: - """Return whether *path* is the Comfy int8-convrot overlay file.""" raw = os.path.basename(str(path)) real = os.path.basename(os.path.realpath(str(path))) return INT8_CONVROT_FILENAME in (raw, real) or "int8_convrot" in raw or "int8_convrot" in real @@ -315,6 +313,5 @@ def dense_vae_safetensors(paths: list[str]) -> list[str]: def find_int8_convrot_vae_path(model_path: str | Path) -> Path | None: - """Return the ConvRot export next to a VAE directory, if present.""" candidate = Path(model_path) / INT8_CONVROT_FILENAME return candidate if candidate.is_file() else None diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 6ab034ccf9..f916e6982b 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -432,7 +432,6 @@ def _release_text_encoder(self) -> None: logger.info("Moved MiniMax-H3 text encoder to CPU after conditioning") def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: - """Reload or GPU-restore Qwen3-VL so a later request can encode.""" encoder = self.get_module("text_encoder") stage = self._stage_name_mapping.get("conditioning_stage") if encoder is not None: diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_purge.py b/fastvideo/tests/ops/quantization/test_nvfp4_purge.py index 385cc47566..18b8566349 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_purge.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_purge.py @@ -114,6 +114,27 @@ def fake_mm_fp4(x_fp4, w_t, x_scale, w_scale_t, alpha, out_dtype, out, backend): assert out.shape == (2, 3, 8) +def test_convert_refuses_fsdp_sharded_bf16_purge(monkeypatch) -> None: + import torch.distributed.tensor as tdt + + class DummyDTensor: + def __init__(self, data: torch.Tensor) -> None: + self._data = data + + def to_local(self) -> torch.Tensor: + return self._data + + def float(self) -> torch.Tensor: + return self._data.float() + + monkeypatch.setattr(tdt, "DTensor", DummyDTensor) + model = _model(retain=False) + del model.always_fp4._parameters["weight"] + object.__setattr__(model.always_fp4, "weight", DummyDTensor(torch.randn(8, 16, dtype=torch.bfloat16))) + with pytest.raises(RuntimeError, match="FSDP-sharded"): + nv.convert_model_to_nvfp4(model) + + def test_dense_path_after_purge_raises_with_flag_named(monkeypatch) -> None: """Defensive guard: convert never purges dense-capable layers anymore, but a hand-purged module hitting the dense path must fail loudly.""" From d07828340afcd9738ad8ea593c2274f7d687ea21 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 15:07:08 -0700 Subject: [PATCH 005/107] [feat]: enable CompactH3 VAE compile on RTX PRO 6000 (cherry picked from commit 7d2a059427665bcb66ffd4d9385356ac0a9bc63e) --- docs/assets/cookbook-recipes.json | 7 ++++--- docs/cookbook/cosmos.md | 2 +- docs/cookbook/flux.md | 2 +- docs/cookbook/glm-image.md | 2 +- docs/cookbook/hunyuan.md | 2 +- docs/cookbook/kandinsky5.md | 2 +- docs/cookbook/longcat.md | 2 +- docs/cookbook/ltx.md | 2 +- docs/cookbook/matrix-game.md | 2 +- docs/cookbook/minimax-h3.md | 4 ++-- docs/cookbook/mmaudio.md | 2 +- docs/cookbook/stable-audio.md | 2 +- docs/cookbook/stable-diffusion.md | 2 +- docs/cookbook/turbodiffusion.md | 2 +- docs/cookbook/wan.md | 2 +- docs/cookbook/z-image.md | 2 +- .../inference/basic/basic_compacth3_rtx_pro6000.yaml | 11 +++++++---- examples/serving/openai_compacth3_rtx_pro6000.yaml | 9 ++++++--- 18 files changed, 33 insertions(+), 26 deletions(-) diff --git a/docs/assets/cookbook-recipes.json b/docs/assets/cookbook-recipes.json index 242ad9211b..399cd4a345 100644 --- a/docs/assets/cookbook-recipes.json +++ b/docs/assets/cookbook-recipes.json @@ -1,5 +1,5 @@ { - "version": 12, + "version": 13, "recipes": [ { "id": "fastwan21-t2v", @@ -622,7 +622,7 @@ "stage": "inference", "task": "Few-step text to video (with audio)", "label": "CompactH3 NVFP4 on RTX PRO 6000 Blackwell", - "summary": "Run CompactH3 NVFP4 with the encoder, DiT, and int8-convrot VAE resident on one 96 GB RTX PRO 6000 Blackwell. The checked-in example is 1344x768 and 124 frames (5.17 s).", + "summary": "Run CompactH3 NVFP4 with the encoder, DiT, and int8-convrot VAE resident on one 96 GB RTX PRO 6000 Blackwell. The checked-in example is 1344x768 and 124 frames (5.17 s) with VAE torch.compile. Use 832x480 for clip-queue playground traffic.", "model": "./CompactH3", "source": "examples/inference/basic/basic_compacth3_rtx_pro6000.yaml", "serving": { @@ -644,8 +644,9 @@ "limitations": [ "Assemble ./CompactH3 before running. The DiT NVFP4 export and encoder snapshot are gated; run huggingface-cli login and accept each repo license.", "96 GB keeps the encoder, DiT, and VAE on GPU. Do not enable h3_sequential_load or lazy_module_load on this box.", + "Enable compile.vae_enabled. Leave inference_torch_compile off: FlashInfer and Sage3 custom ops cannot be compiled.", "Blackwell sm_120 needs ATTN_QAT_INFER, FLASHINFER_CUDA_ARCH_LIST=12.0a, and a CUDA 12.8 PyTorch wheel.", - "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768." + "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768. CompactH3 is dense; do not turn on VIDEO_SPARSE_ATTN_H3 until a VSA-trained student exists." ] }, { diff --git a/docs/cookbook/cosmos.md b/docs/cookbook/cosmos.md index fbb813fae6..89e39e8457 100644 --- a/docs/cookbook/cosmos.md +++ b/docs/cookbook/cosmos.md @@ -5,7 +5,7 @@ hide: # Cosmos recipes -
      +
      All model families
      diff --git a/docs/cookbook/flux.md b/docs/cookbook/flux.md index f2f159aed7..82508bd81d 100644 --- a/docs/cookbook/flux.md +++ b/docs/cookbook/flux.md @@ -5,7 +5,7 @@ hide: # FLUX recipes -
      +
      All model families
      diff --git a/docs/cookbook/glm-image.md b/docs/cookbook/glm-image.md index ed88351a17..51f4a7f28e 100644 --- a/docs/cookbook/glm-image.md +++ b/docs/cookbook/glm-image.md @@ -5,7 +5,7 @@ hide: # GLM-Image recipes -
      +
      All model families
      diff --git a/docs/cookbook/hunyuan.md b/docs/cookbook/hunyuan.md index 66acf7f288..75caf3144d 100644 --- a/docs/cookbook/hunyuan.md +++ b/docs/cookbook/hunyuan.md @@ -5,7 +5,7 @@ hide: # Hunyuan recipes -
      +
      All model families
      diff --git a/docs/cookbook/kandinsky5.md b/docs/cookbook/kandinsky5.md index 119296eae7..bbd96547b5 100644 --- a/docs/cookbook/kandinsky5.md +++ b/docs/cookbook/kandinsky5.md @@ -5,7 +5,7 @@ hide: # Kandinsky 5 recipes -
      +
      All model families
      diff --git a/docs/cookbook/longcat.md b/docs/cookbook/longcat.md index ce2a81574a..3b43788d43 100644 --- a/docs/cookbook/longcat.md +++ b/docs/cookbook/longcat.md @@ -5,7 +5,7 @@ hide: # LongCat recipes -
      +
      All model families
      diff --git a/docs/cookbook/ltx.md b/docs/cookbook/ltx.md index a9354f30da..53bf0fa14f 100644 --- a/docs/cookbook/ltx.md +++ b/docs/cookbook/ltx.md @@ -5,7 +5,7 @@ hide: # LTX recipes -
      +
      All model families
      diff --git a/docs/cookbook/matrix-game.md b/docs/cookbook/matrix-game.md index 26160e8650..15b62ae34a 100644 --- a/docs/cookbook/matrix-game.md +++ b/docs/cookbook/matrix-game.md @@ -5,7 +5,7 @@ hide: # Matrix Game recipes -
      +
      All model families
      diff --git a/docs/cookbook/minimax-h3.md b/docs/cookbook/minimax-h3.md index b4cc9224f8..e664054eac 100644 --- a/docs/cookbook/minimax-h3.md +++ b/docs/cookbook/minimax-h3.md @@ -13,7 +13,7 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in **CompactH3** is the 42-block 20B NVFP4 H3 checkpoint for one Blackwell GPU (RTX 5090 or RTX PRO 6000). The FastH3 V1 and V2 recipes are unchanged. -
      +
      All model families
      @@ -282,7 +282,7 @@ cd FastVideo
    • The MLX source runtime supports T2VA, optional temporal --fast, optional spatial --fast-spatial, and opt-in VSA on --include-vsa checkpoints. FastH3 V2 MLX converts with --include-vsa and runs eight forwards. FL2VA, Ref2VA, and two-pass refinement are not wired.
    • GPU count and VAE decode backend are configurable in the builder above for FastH3 CUDA recipes. Only the value shown by default has a recorded run; other supported values are unmeasured here.
    • DGX Spark is a runtime on FastH3 V1, not a separate family card. Select NVIDIA DGX Spark, then 1 Spark or 2 Sparks. The CUDA GPU-count knob does not apply to Spark.
    • -
    • CompactH3 NVFP4 is one Blackwell GPU. RTX 5090 (32 GB) parks the encoder in pinned host RAM. RTX PRO 6000 Blackwell (96 GB) keeps encoder, DiT, and VAE resident. Keep FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER, FASTVIDEO_FA4=0, FASTVIDEO_VSA_SM100A=0, and FLASHINFER_CUDA_ARCH_LIST=12.0a.
    • +
    • CompactH3 NVFP4 is one Blackwell GPU. RTX 5090 (32 GB) parks the encoder in pinned host RAM. RTX PRO 6000 Blackwell (96 GB) keeps encoder, DiT, and VAE resident. Keep FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER, FASTVIDEO_FA4=0, FASTVIDEO_VSA_SM100A=0, and FLASHINFER_CUDA_ARCH_LIST=12.0a. On PRO 6000 enable VAE compile and leave DiT inference_torch_compile off. CompactH3 is a dense prune; do not enable VIDEO_SPARSE_ATTN_H3 until a VSA-trained student exists.
    • GB10 has no FA4 / sm_100a VSA kernel. Keep FASTVIDEO_FA4=0 and FASTVIDEO_VSA_SM100A=0. Legal num_frames values are 17n+5, capped at 362 (15.08 s). A 345-frame request on one Spark can OOM. Native 16:9 sizes include 832×480 and 1344×768.
    • Gated or missing checkpoints: run huggingface-cli login and confirm you accepted the model's license on Hugging Face.
    • diff --git a/docs/cookbook/mmaudio.md b/docs/cookbook/mmaudio.md index 5e127c9b2b..a6e5654e77 100644 --- a/docs/cookbook/mmaudio.md +++ b/docs/cookbook/mmaudio.md @@ -5,7 +5,7 @@ hide: # MMAudio recipes -
      +
      All model families
      diff --git a/docs/cookbook/stable-audio.md b/docs/cookbook/stable-audio.md index 20dc7bfd7a..b2ccbc533b 100644 --- a/docs/cookbook/stable-audio.md +++ b/docs/cookbook/stable-audio.md @@ -5,7 +5,7 @@ hide: # Stable Audio recipes -
      +
      All model families
      diff --git a/docs/cookbook/stable-diffusion.md b/docs/cookbook/stable-diffusion.md index 07c87fdb59..653b636779 100644 --- a/docs/cookbook/stable-diffusion.md +++ b/docs/cookbook/stable-diffusion.md @@ -5,7 +5,7 @@ hide: # Stable Diffusion recipes -
      +
      All model families
      diff --git a/docs/cookbook/turbodiffusion.md b/docs/cookbook/turbodiffusion.md index 14820b5635..860abee1cf 100644 --- a/docs/cookbook/turbodiffusion.md +++ b/docs/cookbook/turbodiffusion.md @@ -5,7 +5,7 @@ hide: # TurboDiffusion recipes -
      +
      All model families
      diff --git a/docs/cookbook/wan.md b/docs/cookbook/wan.md index 1d72368078..a7a55d0009 100644 --- a/docs/cookbook/wan.md +++ b/docs/cookbook/wan.md @@ -5,7 +5,7 @@ hide: # Wan recipes -
      +
      All model families
      diff --git a/docs/cookbook/z-image.md b/docs/cookbook/z-image.md index b53c79ef16..7f39034dc8 100644 --- a/docs/cookbook/z-image.md +++ b/docs/cookbook/z-image.md @@ -5,7 +5,7 @@ hide: # Z-Image recipes -
      +
      All model families
      diff --git a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml index 4cbf422a9e..a85fbb36c2 100644 --- a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml +++ b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml @@ -5,9 +5,12 @@ # sits beside the dense VAE shards. # # Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). -# 96 GB keeps the encoder, DiT, and VAE resident. request.sampling is an -# example: native 16:9 sizes include 832x480 and 1344x768. Legal H3 frame -# counts are 17n+5, max 362 (15.08 s). +# 96 GB keeps the encoder, DiT, and VAE resident. VAE compile is the measured +# decode win on this card; do not wrap the DiT with inference_torch_compile +# (FlashInfer / Sage3 custom ops). request.sampling is an example: native +# 16:9 sizes include 832x480 (faster clip-queue) and 1344x768. Legal H3 +# frame counts are 17n+5, max 362 (15.08 s). CompactH3 is dense; do not +# enable VIDEO_SPARSE_ATTN_H3 until a VSA-trained student exists. # # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ @@ -32,7 +35,7 @@ generator: lazy_module_load: false compile: enabled: false - vae_enabled: false + vae_enabled: true pipeline: experimental: attention_backend: ATTN_QAT_INFER diff --git a/examples/serving/openai_compacth3_rtx_pro6000.yaml b/examples/serving/openai_compacth3_rtx_pro6000.yaml index 844b98faad..99d1a64019 100644 --- a/examples/serving/openai_compacth3_rtx_pro6000.yaml +++ b/examples/serving/openai_compacth3_rtx_pro6000.yaml @@ -1,10 +1,13 @@ # CompactH3 (42-block 20B NVFP4) on one RTX PRO 6000 Blackwell (sm_120, 96 GB). # Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. 96 GB holds the # encoder, DiT, and VAE resident, so sequential ping-pong stays off. +# Enable VAE compile; leave inference_torch_compile off (custom ops). +# Playground /v1/videos uses default_request. 832x480 124 frames is the +# clip-queue path; 1344x768 is the quality default below. CompactH3 is dense. # # Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). -# default_request.sampling is an example. Native 16:9 sizes include 832x480 -# and 1344x768. Legal H3 frame counts are 17n+5, max 362 (15.08 s). +# Native 16:9 sizes include 832x480 and 1344x768. Legal H3 frame counts are +# 17n+5, max 362 (15.08 s). # # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a \ @@ -30,7 +33,7 @@ generator: lazy_module_load: false compile: enabled: false - vae_enabled: false + vae_enabled: true pipeline: workload_type: t2v experimental: From a766348ae4a717ade2c2239b10460eacae09ba54 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 15:39:39 -0700 Subject: [PATCH 006/107] [bugfix]: refuse zero-gate CompactH3 VSA Dense CompactH3 constructed VIDEO_SPARSE_ATTN_H3 tiles with all-zero gates. Fail at load and first denoise, and drop restating CompactH3 banners. (cherry picked from commit cc59773d005220ab443d5e2ac7de10d937dcff67) --- docs/assets/cookbook-recipes.json | 2 +- .../basic/basic_compacth3_rtx5090.yaml | 10 --- .../basic/basic_compacth3_rtx_pro6000.yaml | 14 ---- .../serving/openai_compacth3_rtx5090.yaml | 7 -- .../serving/openai_compacth3_rtx_pro6000.yaml | 11 --- fastvideo/layers/quantization/nvfp4_config.py | 6 +- fastvideo/models/loader/fsdp_load.py | 10 --- .../basic/minimax_h3/minimax_h3_pipeline.py | 48 ++---------- .../minimax_h3/stages/minimax_h3_denoising.py | 6 ++ .../pipelines/basic/minimax_h3/vsa_guard.py | 32 ++++++++ .../quantization/test_nvfp4_h3_dit_export.py | 1 - .../tests/stages/test_minimax_h3_vsa_guard.py | 74 +++++++++++++++++++ .../vaes/test_minimax_h3_int8_convrot.py | 1 - 13 files changed, 121 insertions(+), 101 deletions(-) create mode 100644 fastvideo/pipelines/basic/minimax_h3/vsa_guard.py create mode 100644 fastvideo/tests/stages/test_minimax_h3_vsa_guard.py diff --git a/docs/assets/cookbook-recipes.json b/docs/assets/cookbook-recipes.json index 399cd4a345..19839b8b7d 100644 --- a/docs/assets/cookbook-recipes.json +++ b/docs/assets/cookbook-recipes.json @@ -646,7 +646,7 @@ "96 GB keeps the encoder, DiT, and VAE on GPU. Do not enable h3_sequential_load or lazy_module_load on this box.", "Enable compile.vae_enabled. Leave inference_torch_compile off: FlashInfer and Sage3 custom ops cannot be compiled.", "Blackwell sm_120 needs ATTN_QAT_INFER, FLASHINFER_CUDA_ARCH_LIST=12.0a, and a CUDA 12.8 PyTorch wheel.", - "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768. CompactH3 is dense; do not turn on VIDEO_SPARSE_ATTN_H3 until a VSA-trained student exists." + "Legal num_frames values are 17n+5, capped at 362 (15.08 s). Native 16:9 sizes include 832x480 and 1344x768. Dense CompactH3 has zero VSA gates; the pipeline raises if VIDEO_SPARSE_ATTN_H3 is loaded with all-zero to_gate_compress weights." ] }, { diff --git a/examples/inference/basic/basic_compacth3_rtx5090.yaml b/examples/inference/basic/basic_compacth3_rtx5090.yaml index f17114b240..46ce251b8f 100644 --- a/examples/inference/basic/basic_compacth3_rtx5090.yaml +++ b/examples/inference/basic/basic_compacth3_rtx5090.yaml @@ -1,13 +1,3 @@ -# CompactH3 (42-block 20B NVFP4) on one RTX 5090 (sm_120, 32 GB). -# Packed ``transformer/nvfp4_weights.safetensors`` plus -# ``engine.quantization.layer_profile: h3_dit``. Full decode uses the Comfy -# int8-convrot VAE when ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` -# sits beside the dense VAE shards. -# -# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). -# Sequential load parks the encoder in pinned host RAM; do not enable -# lazy_module_load. Legal H3 frame counts are 17n+5, max 362 (15.08 s). -# # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ # fastvideo generate --config examples/inference/basic/basic_compacth3_rtx5090.yaml diff --git a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml index a85fbb36c2..6c92ae7e8e 100644 --- a/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml +++ b/examples/inference/basic/basic_compacth3_rtx_pro6000.yaml @@ -1,17 +1,3 @@ -# CompactH3 (42-block 20B NVFP4) on one RTX PRO 6000 Blackwell (sm_120, 96 GB). -# Packed ``transformer/nvfp4_weights.safetensors`` plus -# ``engine.quantization.layer_profile: h3_dit``. Full decode uses the Comfy -# int8-convrot VAE when ``vae/minimax_h3_video_vae_int8_convrot.safetensors`` -# sits beside the dense VAE shards. -# -# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). -# 96 GB keeps the encoder, DiT, and VAE resident. VAE compile is the measured -# decode win on this card; do not wrap the DiT with inference_torch_compile -# (FlashInfer / Sage3 custom ops). request.sampling is an example: native -# 16:9 sizes include 832x480 (faster clip-queue) and 1344x768. Legal H3 -# frame counts are 17n+5, max 362 (15.08 s). CompactH3 is dense; do not -# enable VIDEO_SPARSE_ATTN_H3 until a VSA-trained student exists. -# # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a FASTVIDEO_STAGE_LOGGING=1 \ # fastvideo generate --config examples/inference/basic/basic_compacth3_rtx_pro6000.yaml diff --git a/examples/serving/openai_compacth3_rtx5090.yaml b/examples/serving/openai_compacth3_rtx5090.yaml index 252df0b522..136fc0af9a 100644 --- a/examples/serving/openai_compacth3_rtx5090.yaml +++ b/examples/serving/openai_compacth3_rtx5090.yaml @@ -1,10 +1,3 @@ -# CompactH3 (42-block 20B NVFP4) on one RTX 5090 (sm_120, 32 GB). -# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. -# 32 GB cannot co-reside the encoder and DiT; sequential load parks the -# encoder in pinned host RAM between phases. Do not enable lazy_module_load. -# -# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). -# # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a \ # fastvideo serve --config examples/serving/openai_compacth3_rtx5090.yaml diff --git a/examples/serving/openai_compacth3_rtx_pro6000.yaml b/examples/serving/openai_compacth3_rtx_pro6000.yaml index 99d1a64019..0f947feec0 100644 --- a/examples/serving/openai_compacth3_rtx_pro6000.yaml +++ b/examples/serving/openai_compacth3_rtx_pro6000.yaml @@ -1,14 +1,3 @@ -# CompactH3 (42-block 20B NVFP4) on one RTX PRO 6000 Blackwell (sm_120, 96 GB). -# Packed NVFP4 DiT + NVFP4 encoder + Comfy int8-convrot VAE. 96 GB holds the -# encoder, DiT, and VAE resident, so sequential ping-pong stays off. -# Enable VAE compile; leave inference_torch_compile off (custom ops). -# Playground /v1/videos uses default_request. 832x480 124 frames is the -# clip-queue path; 1344x768 is the quality default below. CompactH3 is dense. -# -# Assemble ./CompactH3 first (see the MiniMax H3 cookbook CompactH3 recipe). -# Native 16:9 sizes include 832x480 and 1344x768. Legal H3 frame counts are -# 17n+5, max 362 (15.08 s). -# # FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FLASHINFER_CUDA_ARCH_LIST=12.0a \ # fastvideo serve --config examples/serving/openai_compacth3_rtx_pro6000.yaml diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 3c84f1be37..04558b7fa1 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -96,7 +96,6 @@ def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: def is_minimax_h3_nvfp4_linear_prefix(prefix: str) -> bool: - """Return whether *prefix* is a main MiniMax-H3 transformer-block FFN linear.""" return _MINIMAX_H3_NVFP4_FF_PREFIX.search(prefix) is not None @@ -630,9 +629,8 @@ def load_minimax_h3_nvfp4_dit_export( """Load a packed NVFP4H3 DiT export onto already-tagged NVFP4 linears. Keys are ``::`` with the four buffers - ``convert_model_to_nvfp4`` registers. The bf16 ``weight`` is dropped so a - 32 GB card never materializes the dense GEMMs. Every export prefix must - match an NVFP4 linear, and every NVFP4 linear must appear in the export. + ``convert_model_to_nvfp4`` registers. Every export prefix must match an + NVFP4 linear, and every NVFP4 linear must appear in the export. """ from safetensors import safe_open diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index ad1867a4b3..020a8f970a 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -225,8 +225,6 @@ def maybe_load_fsdp_model( load_minimax_h3_nvfp4_dit_export, nvfp4_linear_weight_param_names, ) - # Always drop the packed export from the dense shard list so ``::`` keys are - # never treated as bf16 parameters. Overlay it only for the H3 DiT profile. packed_candidate = find_minimax_h3_nvfp4_dit_export(weight_dir_list) weight_dir_list = dense_transformer_safetensors(weight_dir_list) quant_config = getattr(init_params.get("config"), "quant_config", None) @@ -296,11 +294,6 @@ def maybe_load_fsdp_model( fsdp_shard_conditions=model._fsdp_shard_conditions, pin_cpu_memory=pin_cpu_memory) - # Host offload is already disabled on unified memory (GB10). Staging the - # 35B FastH3 DiT on CPU and then copying to CUDA doubled that working set - # and took minutes. Follow cpu_offload: read onto the accelerator. - # Packed NVFP4H3 DiT exports still read the remaining dense shards on CPU - # so the 32 GB GEMMs never land on a 32 GB card before they are skipped. nvfp4_skip_param_names: set[str] = set() if packed_nvfp4_export is not None: nvfp4_skip_param_names = nvfp4_linear_weight_param_names(model) @@ -321,9 +314,6 @@ def maybe_load_fsdp_model( strength=lora_strength, ) if dense_lora_patch is not None: - # H3's compression gate is created only by the VSA attention backend. Loading a - # VSA student under dense attention would otherwise warn about 50 unmatched - # replacements and continue with a silently incomplete model. model_parameter_names = {name for name, _ in model.named_parameters()} missing_vsa_gates = sorted(name for name in dense_lora_patch.replacement_parameters if "gate_compress" in name and name not in model_parameter_names) diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index f916e6982b..1310f53fc9 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -"""FastVideo composed pipelines for MiniMax H3.""" from __future__ import annotations @@ -30,6 +29,7 @@ MiniMaxH3LatentPreparationStage, MiniMaxH3VideoDecodingStage, ) +from fastvideo.pipelines.basic.minimax_h3.vsa_guard import refuse_zero_initialized_h3_vsa from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase from fastvideo.pipelines.lora_pipeline import LoRAPipeline from fastvideo.pipelines.pipeline_batch_info import ForwardBatch @@ -37,9 +37,6 @@ logger = init_logger(__name__) -# Same split as the MLX runtime: condition, release the ~66 GB Qwen3-VL stack, -# then load DiT + VAEs. Keeping them resident together OOMs unified-memory -# boxes (GB10 / Spark) even though host offload is correctly disabled there. _DENOISE_MODULE_NAMES = ("vae", "audio_vae", "transformer") @@ -79,7 +76,6 @@ def _checkpoint_has_vsa_gates(transformer_dir: Path) -> bool: def _module_has_dtensor_params(module: Any) -> bool: - """FSDP2 shards stay put; packed NVFP4 params are ordinary tensors and can ping-pong.""" parameters = getattr(module, "parameters", None) if not callable(parameters): return False @@ -111,7 +107,6 @@ def _default_audio_geometry() -> _H3AudioGeometry: def _apply_h3_checkpoint_arch_configs(model_path: str, fastvideo_args: FastVideoArgs, extra_config_module_map: dict[str, str]) -> None: - """Overlay checkpoint config.json onto pipeline configs without loading weights.""" root = Path(model_path) vae_dir = root / extra_config_module_map.get("vae", "vae") if (vae_dir / "config.json").is_file(): @@ -140,16 +135,7 @@ def _use_taeh3_t2va(fastvideo_args: FastVideoArgs | None, *, ref2va: bool) -> bo class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): - """Shared loading and target-generation path for MiniMax H3. - Inherits ``LoRAPipeline`` so acceleration and distillation adapters can be merged - in; without it every adapter is rejected with "pipeline is not a LoRAPipeline". - """ - - # The linears every published H3 adapter targets. Left unset, ``LoRAPipeline`` - # wraps *every* linear in the DiT -- including ``proj_in``, whose ``.weight`` the - # forward pass reads directly. ``BaseLayerWithLoRA`` exposes no ``.weight``, so - # that wrapping turns generation into an AttributeError before the first step. lora_target_modules = [ "attn.to_q", "attn.to_k", @@ -158,8 +144,6 @@ class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): "ff.fc_in", "ff.fc_out", "adaln_proj.linear", - # The final AdaLN projection. Published community adapters (larryvrh's Turbo) - # target it as `final_layer.adaln_proj.linear`. "norm_out.linear", ] @@ -175,10 +159,6 @@ class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): "scheduler", "audio_scheduler", ] - # Deferral is safe here: geometry scalars come from checkpoint config.json - # (applied in initialize_pipeline without loading weights), no stage - # constructor reads a deferred component, and initialize_pipeline only - # inspects the schedulers, which are never deferred. _lazy_module_names = ("text_encoder", "transformer", "vae", "audio_vae") def __init__(self, *args: Any, **kwargs: Any) -> None: @@ -200,14 +180,14 @@ def get_hf_download_allow_patterns(cls) -> list[str]: def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None: _apply_h3_checkpoint_arch_configs(self.model_path, fastvideo_args, self._extra_config_module_map) - # Each modality's scheduler_config.json owns its shift. Base H3 keeps - # 12/3; a distilled checkpoint can serialize a different trained pair - # (for example 10/3) without being silently rewritten to base defaults. for module_name, modality in (("scheduler", "video"), ("audio_scheduler", "audio")): shift = getattr(self.get_module(module_name), "shift", None) if shift is None or not math.isfinite(float(shift)) or float(shift) <= 0: raise ValueError(f"MiniMax-H3 {modality} scheduler must expose a positive finite shift, got {shift}.") self._load_checkpoint_schedule(fastvideo_args) + transformer = self.get_module("transformer") + if transformer is not None: + refuse_zero_initialized_h3_vsa(transformer) def _checkpoint_facts(self) -> tuple[Any, bool]: """The checkpoint's parsed ``fastvideo_inference.json`` (None without one) and whether its transformer @@ -319,13 +299,6 @@ def _load_checkpoint_schedule(self, fastvideo_args: FastVideoArgs) -> None: def _defer_denoise_modules(self, fastvideo_args: FastVideoArgs) -> bool: if not fastvideo_args.inference_mode or bool(getattr(fastvideo_args, "training_mode", False)): return False - # Both mechanisms defer the same four modules and both decide when to - # free them. Running them together strips DiT/VAEs from the first load - # (sequential) while the base wraps the encoder in a proxy (lazy), so - # post_init's VAE compile transform has nothing to attach to. Lazy is - # the more general owner — including auto-on for unified memory — so it - # wins whenever it is on. Sequential remains the H3-only fallback when - # lazy is off. if bool(getattr(fastvideo_args, "lazy_module_load", False)): logger.info("MiniMax-H3 sequential module load off: lazy_module_load owns deferral") return False @@ -372,8 +345,6 @@ def load_modules(self, return super().load_modules(fastvideo_args, loaded_modules) saved = list(self.required_config_modules) - # Always defer the full denoise set on the first load. TAEH3 T2VA then - # omits the video VAE from the second load via `_denoise_module_names`. self._required_config_modules = [name for name in saved if name not in _DENOISE_MODULE_NAMES] try: logger.info("Loading MiniMax-H3 condition modules first: %s", self._required_config_modules) @@ -456,15 +427,12 @@ 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: - # MiniMaxH3ConditioningStage already refuses .to() on DTensor; default - # sequential Qwen is FSDP2-sharded. Packed NVFP4 skips that wrap. if _module_has_dtensor_params(module): return False module.to(device) return True def _park_denoise_modules(self) -> None: - """Free the GPU so the NVFP4 encoder can encode without sitting on the DiT.""" parked = False for name in _DENOISE_MODULE_NAMES: module = self.get_module(name) @@ -511,7 +479,6 @@ def _run_condition_then_denoise(self, batch: ForwardBatch, fastvideo_args: FastV return batch def _input_video_geometry(self, fastvideo_args: FastVideoArgs) -> Any: - """Read canvas scalars from checkpoint JSON, not a live VAE proxy.""" arch = getattr(getattr(fastvideo_args.pipeline_config, "vae_config", None), "arch_config", None) if arch is not None: return arch @@ -552,6 +519,8 @@ def _add_condition_stages(self, fastvideo_args: FastVideoArgs, *, ref2va: bool) def _add_denoise_stages(self, *, ref2va: bool) -> None: transformer = self.get_module("transformer") + if transformer is not None: + refuse_zero_initialized_h3_vsa(transformer) vae = self.get_module("vae") audio_vae = self.get_module("audio_vae") scheduler = self.get_module("scheduler") @@ -593,11 +562,6 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward if not self.post_init_called: self.post_init() - # Sequential encode-then-release is the H3-only fallback. Lazy and the - # fully-resident discrete-GPU path both keep a complete stage list and - # must use the base forward so abort cleanup and text_encoder_cpu_offload - # still apply. Releasing Qwen on every request was re-reading it from disk - # when neither deferral flag was on. if self._defer_denoise_modules(fastvideo_args): try: if not self._unified_memory_host(): diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py index c59ac6f113..83d9a004a0 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py @@ -22,6 +22,8 @@ build_row_timesteps, ) from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_latent_preparation import MINIMAX_H3_LAYOUT_KEY +from fastvideo.pipelines.basic.minimax_h3.vsa_guard import refuse_zero_initialized_h3_vsa +from fastvideo.pipelines.lazy_module import is_lazy_module from fastvideo.pipelines.pipeline_batch_info import ForwardBatch from fastvideo.pipelines.stages.base import PipelineStage from fastvideo.pipelines.stages.validators import StageValidators as V @@ -191,6 +193,10 @@ def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> V @torch.no_grad() def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: """Denoise the packed H3 video and audio streams over one shared schedule.""" + transformer = self.transformer + if is_lazy_module(transformer): + transformer.materialize() + refuse_zero_initialized_h3_vsa(transformer) layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY) if not isinstance(layout, MiniMaxH3PackedLayout): raise ValueError("MiniMax-H3 packed layout is missing before denoising.") diff --git a/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py new file mode 100644 index 0000000000..4bed84a162 --- /dev/null +++ b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py @@ -0,0 +1,32 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import Any + +from fastvideo.pipelines.lazy_module import is_lazy_module + + +def refuse_zero_initialized_h3_vsa(transformer: Any) -> None: + if is_lazy_module(transformer) and not transformer.is_materialized: + return + blocks = getattr(transformer, "transformer_blocks", None) + if blocks is None: + return + saw_weight = False + any_trained = False + for block in blocks: + attention = getattr(block, "attn", None) + gate = getattr(attention, "to_gate_compress", None) if attention is not None else None + if gate is None: + continue + weight = getattr(gate, "weight", None) + if weight is None: + continue + saw_weight = True + if bool((weight != 0).any()): + any_trained = True + break + if saw_weight and not any_trained: + raise RuntimeError("VIDEO_SPARSE_ATTN_H3 is loaded but every to_gate_compress weight is zero. " + "Load a VSA-trained student or a LoRA that carries gates; dense CompactH3 cannot run VSA.") diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py b/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py index 80a7f4e3df..32b9f84d14 100644 --- a/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py +++ b/fastvideo/tests/ops/quantization/test_nvfp4_h3_dit_export.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -"""CPU tests for the packed MiniMax-H3 DiT NVFP4 export overlay.""" from __future__ import annotations diff --git a/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py new file mode 100644 index 0000000000..c4c44a4931 --- /dev/null +++ b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py @@ -0,0 +1,74 @@ +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch +import torch.nn as nn + +from fastvideo.pipelines.basic.minimax_h3.vsa_guard import refuse_zero_initialized_h3_vsa +from fastvideo.pipelines.lazy_module import LazyModule + + +class _GateAttention(nn.Module): + + def __init__(self, trained: bool, *, has_weight: bool = True) -> None: + super().__init__() + if has_weight: + self.to_gate_compress = nn.Linear(4, 4, bias=False) + with torch.no_grad(): + if trained: + self.to_gate_compress.weight.fill_(0.1) + else: + self.to_gate_compress.weight.zero_() + else: + self.to_gate_compress = SimpleNamespace(weight=None) + + +class _Block(nn.Module): + + def __init__(self, trained: bool, *, has_weight: bool = True) -> None: + super().__init__() + self.attn = _GateAttention(trained, has_weight=has_weight) + + +class _Transformer(nn.Module): + + def __init__(self, trained: bool, *, empty: bool = False, has_weight: bool = True) -> None: + super().__init__() + blocks = [] if empty else [_Block(trained, has_weight=has_weight)] + self.transformer_blocks = nn.ModuleList(blocks) + + +def test_refuse_zero_initialized_h3_vsa_raises() -> None: + with pytest.raises(RuntimeError, match="to_gate_compress"): + refuse_zero_initialized_h3_vsa(_Transformer(trained=False)) + + +def test_refuse_zero_initialized_h3_vsa_allows_trained_gates() -> None: + refuse_zero_initialized_h3_vsa(_Transformer(trained=True)) + + +def test_refuse_zero_initialized_h3_vsa_skips_dense_transformer() -> None: + refuse_zero_initialized_h3_vsa(SimpleNamespace()) + + +def test_refuse_zero_initialized_h3_vsa_skips_empty_blocks() -> None: + refuse_zero_initialized_h3_vsa(_Transformer(trained=False, empty=True)) + + +def test_refuse_zero_initialized_h3_vsa_skips_missing_gate_weight() -> None: + refuse_zero_initialized_h3_vsa(_Transformer(trained=False, has_weight=False)) + + +def test_refuse_zero_initialized_h3_vsa_skips_unmaterialized_lazy_module() -> None: + loads: list[int] = [] + + def loader() -> nn.Module: + loads.append(1) + return _Transformer(trained=False) + + refuse_zero_initialized_h3_vsa(LazyModule("transformer", loader)) + assert loads == [] diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py index 8a4984ed7c..165ccc41e5 100644 --- a/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_convrot.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -"""CPU contracts for the MiniMax-H3 Comfy int8-convrot VAE overlay.""" from __future__ import annotations From 1fb455d7b0632707a6c595fdab14136f635f4be3 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Thu, 1 Oct 2026 13:43:13 -0700 Subject: [PATCH 007/107] [wip]: sm_120 H3 FP4 experiments: SP FP8 exchange path and Modal drivers Snapshot of every RTX PRO 6000 experiment, including the Ulysses FP8 q/k/v exchange path that has not executed yet (8-GPU capacity was unavailable) and the Modal drivers used for every measurement. The ship-ready subset is on h3-sm120-sparse-fp4. (cherry picked from commit da835b48c7b0ac635fd38636f657e3a81a74e9be) --- fastvideo-kernel/attn_qat_infer/api.py | 185 +++++++- .../attn_qat_infer/blackwell/api.cu | 87 +++- .../attn_qat_infer/blackwell/kernel_ws.h | 6 +- .../attn_qat_infer/blackwell/launch.h | 4 +- .../blackwell/mainloop_tma_ws.h | 149 ++++++- .../attn_qat_infer/blackwell/params.h | 17 + fastvideo/api/schema.py | 2 + fastvideo/layers/quantization/nvfp4_config.py | 29 +- fastvideo/models/dits/minimax_h3.py | 41 +- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 282 ++++++++++++ fastvideo/models/loader/fsdp_load.py | 4 +- fastvideo/models/vaes/minimax_h3_video.py | 26 +- scripts/benchmarks/minimax_h3_pro6000/a2a8.py | 72 ++++ scripts/benchmarks/minimax_h3_pro6000/app.py | 408 ++++++++++++++++++ .../minimax_h3_pro6000/bench_code.py | 338 +++++++++++++++ .../minimax_h3_pro6000/download_personal.py | 40 ++ .../convert_minimax_h3_modelopt_nvfp4_dit.py | 217 ++++++++++ 17 files changed, 1879 insertions(+), 28 deletions(-) create mode 100644 fastvideo/models/dits/minimax_h3_vsa_fp4.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/a2a8.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/app.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/bench_code.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/download_personal.py create mode 100644 scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py diff --git a/fastvideo-kernel/attn_qat_infer/api.py b/fastvideo-kernel/attn_qat_infer/api.py index f506d692da..03ea3a6565 100644 --- a/fastvideo-kernel/attn_qat_infer/api.py +++ b/fastvideo-kernel/attn_qat_infer/api.py @@ -106,12 +106,32 @@ def pad_to_block_size(x): if enable_smoothing_q: delta_s = torch.matmul(qm, k.transpose(-2, -1)).to(torch.float32).contiguous() else: # used to disable q smoothing - B, H, L, D = q.shape - delta_s = torch.zeros((B, H, L // BLOCK_M, k.shape[2]), device=q.device, dtype=torch.float32) + delta_s = _zero_delta_s(q.shape[0], q.shape[1], k.shape[2], q.device) return q, k, v, delta_s +_ZERO_DELTA_S: dict = {} + + +def _zero_delta_s(batch: int, heads: int, kv_len: int, device: torch.device) -> torch.Tensor: + """Cached all-zero delta_s for unsmoothed Q, read with per_block_mean=False. + + The kernel reads delta_s as a contiguous [B, H, rows, KL] tensor with one + row per query block when per_block_mean is set and a single row otherwise. + With Q smoothing off every row is zero, so one shared row replaces the + [B, H, L/128, KL] tensor that was allocated and zero-filled on every call + (9.5 GB at 73k tokens, whose int32 batch stride also broke the TMA + descriptor). The kernel only reads it. + """ + key = (batch, heads, kv_len, device) + zeros = _ZERO_DELTA_S.get(key) + if zeros is None: + zeros = torch.zeros((batch, heads, 1, kv_len), device=device, dtype=torch.float32) + _ZERO_DELTA_S[key] = zeros + return zeros + + def scale_and_quant_fp4(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: assert x.ndim == 4 B, H, N, D = x.shape @@ -154,6 +174,166 @@ def blockscaled_fp4_attn(qlist: Tuple, softmax_scale, is_causal, per_block_mean, is_bf16, single_level_p_quant) +def blockscaled_fp4_attn_sparse(qlist: Tuple, + klist: Tuple, + vlist: Tuple, + delta_s: torch.Tensor, + KL: int, + q2k_idx: torch.Tensor, + q2k_num: torch.Tensor, + kv_valid: torch.Tensor | None = None, + q2k_quad: torch.Tensor | None = None, + per_block_mean: bool = True, + is_bf16: bool = True, + single_level_p_quant: bool = False, + sm_scale: float | None = None): + softmax_scale = sm_scale if sm_scale is not None else (qlist[0].shape[-1] * 2)**(-0.5) + return fp4attn_cuda.fwd_sparse(qlist[0], klist[0], vlist[0], qlist[1], klist[1], vlist[1], delta_s, KL, None, + softmax_scale, per_block_mean, is_bf16, single_level_p_quant, q2k_idx, q2k_num, + kv_valid, q2k_quad) + + +HALF_N = BLOCK_N // 2 + + +def vsa_tile_mask_to_fp4_blocks( + tile_mask: torch.Tensor, + tile_tokens: int, + tile_valid: torch.Tensor | None = None, + validate: bool = False, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: + """Convert a VSA tile mask into the FP4 kernel's block lists. + + ``tile_mask`` is a ``[B, H, T, T]`` bool mask over VSA tiles of + ``tile_tokens`` tokens (query tile, key tile); ``tile_tokens`` is 64 or a + multiple of 128. ``tile_valid`` holds each tile's valid token count (valid + tokens first, padding last, as VSA tiling lays them out). The kernel works + on 128x128 blocks, so 64-token tiles are paired: a block is listed when any + of its four 64x64 quadrants is selected, and ``q2k_quad`` carries which. + + Lists run in descending block order because the kernel visits them last + entry first: block 0 (the first prefix tile, which VSA-H3's exempt mode + gives every query) is then visited first, so every query row starts from a + finite running max. ``validate=True`` checks that (host sync). + + Returns ``(q2k_idx [B, H, M, N], q2k_num [B, H, M], kv_valid [2N], + q2k_quad [B, H, M, N] | None)``; the queries cover ``M * 128`` padded rows. + """ + if tile_tokens % HALF_N or (tile_tokens > HALF_N and tile_tokens % BLOCK_N): + raise ValueError(f"tile_tokens={tile_tokens} must be {HALF_N} or a multiple of {BLOCK_N}") + batch, heads, n_tiles, _ = tile_mask.shape + device = tile_mask.device + halves_per_tile = tile_tokens // HALF_N + half_mask = tile_mask + if halves_per_tile > 1: + half_mask = half_mask.repeat_interleave(halves_per_tile, dim=2).repeat_interleave(halves_per_tile, dim=3) + n_halves = n_tiles * halves_per_tile + if tile_valid is None: + half_valid = torch.full((n_halves, ), HALF_N, device=device, dtype=torch.int32) + else: + offsets = torch.arange(halves_per_tile, device=device, dtype=torch.int32) * HALF_N + half_valid = (tile_valid.to(torch.int32)[:, None] - offsets[None, :]).clamp(0, HALF_N).reshape(-1) + if n_halves % 2: + # Pad to whole 128-blocks: the extra key half is empty, the extra query + # half (discarded output) attends block 0 only. + half_mask = F.pad(half_mask, (0, 1, 0, 1), value=False) + half_mask[:, :, -1, 0] = True + half_valid = F.pad(half_valid, (0, 1), value=0) + n_halves += 1 + half_mask = half_mask & (half_valid > 0)[None, None, None, :] + n_blocks = n_halves // 2 + quads = half_mask.view(batch, heads, n_blocks, 2, n_blocks, 2) + weights = torch.tensor([[1, 2], [4, 8]], device=device, dtype=torch.uint8) # [row_half, col_half] + quad = (quads.to(torch.uint8) * weights[None, None, None, :, None, :]).sum(dim=(3, 5), dtype=torch.uint8) + block_mask = quad != 0 + # Descending compaction without a sort: position = running count from the right. + rev = block_mask.flip(-1) + pos = rev.cumsum(-1, dtype=torch.int32) - 1 + q2k_num = (pos[..., -1] + 1).contiguous() + cols = torch.arange(n_blocks - 1, -1, -1, device=device, dtype=torch.int32).expand_as(pos) + slot = torch.where(rev, pos, torch.full_like(pos, n_blocks)).long() + q2k_idx = torch.zeros((batch, heads, n_blocks, n_blocks + 1), device=device, dtype=torch.int32) + q2k_idx.scatter_(-1, slot, cols) + q2k_idx = q2k_idx[..., :n_blocks].contiguous() + kv_valid = half_valid.contiguous() + q2k_quad = None + if halves_per_tile == 1: + q2k_quad = torch.zeros((batch, heads, n_blocks, n_blocks + 1), device=device, dtype=torch.uint8) + q2k_quad.scatter_(-1, slot, quad.flip(-1)) + q2k_quad = q2k_quad[..., :n_blocks].contiguous() + if validate: + if int(q2k_num.min()) < 1: + raise ValueError("every query block must attend to at least one non-empty KV block") + last = (q2k_num - 1).long().unsqueeze(-1) + first_block = q2k_idx.gather(-1, last.int().long()).squeeze(-1).long() + first_quad = (q2k_quad.gather(-1, last).squeeze(-1).int() if q2k_quad is not None else + torch.full_like(first_block, 15, dtype=torch.int32)) + v0 = (kv_valid[2 * first_block] > 0).int() + v1 = (kv_valid[2 * first_block + 1] > 0).int() + row0 = ((first_quad & 1).bool() & v0.bool()) | ((first_quad & 2).bool() & v1.bool()) + row1 = ((first_quad & 4).bool() & v0.bool()) | ((first_quad & 8).bool() & v1.bool()) + if not bool((row0 & row1).all()): + raise ValueError("the first block each query block visits must give both 64-row halves a valid key") + return q2k_idx, q2k_num, kv_valid, q2k_quad + + +def sageattn_blackwell_sparse(q, + k, + v, + q2k_idx: torch.Tensor, + q2k_num: torch.Tensor, + kv_valid: torch.Tensor | None = None, + q2k_quad: torch.Tensor | None = None, + per_block_mean=True, + single_level_p_quant=True, + sm_scale: float | None = None): + """Block-sparse SageAttention3 FP4 forward (non-causal). + + Query block ``m`` (``BLOCK_M`` rows) of each (batch, head) attends only to + the ``BLOCK_N``-token KV blocks in ``q2k_idx[b, h, m, :q2k_num[b, h, m]]``, + restricted to the quadrants in ``q2k_quad`` when given; see + :func:`vsa_tile_mask_to_fp4_blocks`. Q/K/V are ``[B, H, L, D]``. + """ + QL = q.size(2) + KL = k.size(2) + is_bf16 = q.dtype == torch.bfloat16 + q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean) + per_block_mean = delta_s.shape[2] > 1 + qlist_from_cuda = scale_and_quant_fp4(q) + klist_from_cuda = scale_and_quant_fp4_permute(k) + vlist_from_cuda = scale_and_quant_fp4_transpose(v) + o_fp4 = blockscaled_fp4_attn_sparse(qlist_from_cuda, klist_from_cuda, vlist_from_cuda, delta_s, KL, q2k_idx, + q2k_num, kv_valid, q2k_quad, per_block_mean, is_bf16, single_level_p_quant, + sm_scale)[0][:, :, :QL, :].contiguous() + return o_fp4 + + +def sageattn_blackwell_sparse_bshd(q, + k, + v, + q2k_idx: torch.Tensor, + q2k_num: torch.Tensor, + kv_valid: torch.Tensor | None = None, + q2k_quad: torch.Tensor | None = None, + single_level_p_quant=True, + sm_scale: float | None = None) -> torch.Tensor: + """:func:`sageattn_blackwell_sparse` for ``[B, L, H, D]`` inputs, without copies. + + The FP4 quantizers read strided input, so the sequence-major tensors a + linear produces are quantized in place of a transpose + pad. ``L`` must be + a multiple of ``BLOCK_M`` (callers allocate the padding) and Q is + unsmoothed. Returns ``[B, H, L, D]``. + """ + batch, seq_len, heads, _ = q.shape + if seq_len % BLOCK_M: + raise ValueError(f"sequence length {seq_len} must be a multiple of {BLOCK_M}") + qh, kh, vh = (x.transpose(1, 2) for x in (q, k, v)) + delta_s = _zero_delta_s(batch, heads, seq_len, q.device) + return blockscaled_fp4_attn_sparse(scale_and_quant_fp4(qh), scale_and_quant_fp4_permute(kh), + scale_and_quant_fp4_transpose(vh), delta_s, seq_len, q2k_idx, q2k_num, kv_valid, + q2k_quad, False, q.dtype == torch.bfloat16, single_level_p_quant, sm_scale)[0] + + def sageattn_blackwell(q, k, v, @@ -191,6 +371,7 @@ def sageattn_blackwell(q, KL = k.size(2) is_bf16 = q.dtype == torch.bfloat16 q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean) + per_block_mean = delta_s.shape[2] > 1 qlist_from_cuda = scale_and_quant_fp4(q) klist_from_cuda = scale_and_quant_fp4_permute(k) vlist_from_cuda = scale_and_quant_fp4_transpose(v) diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/api.cu b/fastvideo-kernel/attn_qat_infer/blackwell/api.cu index 243d876dbf..c7a3624f0b 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/api.cu +++ b/fastvideo-kernel/attn_qat_infer/blackwell/api.cu @@ -204,8 +204,17 @@ void run_mha_fwd(Flash_fwd_params ¶ms, cudaStream_t stream, bool force_split })); } -std::vector -mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // 2) +struct SparseKvLists { + int const *q2k_idx = nullptr; + int const *q2k_num = nullptr; + int q2k_max = 0; + int num_m_blocks = 0; + int const *kv_valid = nullptr; + uint8_t const *q2k_quad = nullptr; +}; + +static std::vector +mha_fwd_impl(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // 2) const at::Tensor &k, // batch_size x seqlen_k x num_heads_k x (head_size // 2) const at::Tensor &v, // batch_size x seqlen_k x num_heads_k x (head_size // 2) const at::Tensor &sfq, @@ -218,7 +227,8 @@ mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size bool is_causal, bool per_block_mean, bool is_bf16, - bool single_level_p_quant=false // If true, use only per-row scale s_P2 (no per-block s_P1) + bool single_level_p_quant, // If true, use only per-row scale s_P2 (no per-block s_P1) + SparseKvLists const &sparse ) { auto dprops = at::cuda::getCurrentDeviceProperties(); @@ -316,6 +326,12 @@ mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // stack-local tensor whose data pointer would dangle after mha_fwd returns // while the async kernel may still be running. params.tile_count_semaphore = nullptr; + params.q2k_idx = sparse.q2k_idx; + params.q2k_num = sparse.q2k_num; + params.q2k_max = sparse.q2k_max; + params.num_m_blocks = sparse.num_m_blocks; + params.kv_valid = sparse.kv_valid; + params.q2k_quad = sparse.q2k_quad; if (seqlen_k > 0) { auto stream = at::cuda::getCurrentCUDAStream().stream(); @@ -341,7 +357,72 @@ mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size +std::vector +mha_fwd(at::Tensor &q, const at::Tensor &k, const at::Tensor &v, + const at::Tensor &sfq, const at::Tensor &sfk, const at::Tensor &sfv, + const at::Tensor &delta_s, int unpadded_k, c10::optional &out_, + const float softmax_scale, bool is_causal, bool per_block_mean, bool is_bf16, + bool single_level_p_quant=false) { + return mha_fwd_impl(q, k, v, sfq, sfk, sfv, delta_s, unpadded_k, out_, softmax_scale, + is_causal, per_block_mean, is_bf16, single_level_p_quant, SparseKvLists{}); +} + +// Block-sparse forward: query block m of (b, h) attends only to the KV blocks +// listed in q2k_idx[b, h, m, :q2k_num[b, h, m]] (BLOCK_M x BLOCK_N granularity). +// kv_valid, if given, holds the valid token count of each 64-column half of +// every KV block ([2 * num_kv_blocks], valid tokens first within a half). +// q2k_quad, if given (uint8, same shape as q2k_idx), restricts each listed +// block to the 64x64 quadrants whose bit (2 * row_half + col_half) is set, so +// 64-token VSA tiles run on 128x128 blocks. The block a list visits first +// (its last entry) must leave every query row at least one valid key. +// Non-causal only. +std::vector +mha_fwd_sparse(at::Tensor &q, const at::Tensor &k, const at::Tensor &v, + const at::Tensor &sfq, const at::Tensor &sfk, const at::Tensor &sfv, + const at::Tensor &delta_s, int unpadded_k, c10::optional &out_, + const float softmax_scale, bool per_block_mean, bool is_bf16, + bool single_level_p_quant, + const at::Tensor &q2k_idx, const at::Tensor &q2k_num, + c10::optional &kv_valid_, + c10::optional &q2k_quad_) { + const int batch_size = q.size(0); + const int num_heads = q.size(1); + const int num_m_blocks = (q.size(2) + flash::BLOCK_M - 1) / flash::BLOCK_M; + const int num_n_blocks = (k.size(2) + flash::BLOCK_N - 1) / flash::BLOCK_N; + for (auto const *t : {&q2k_idx, &q2k_num}) { + TORCH_CHECK(t->scalar_type() == torch::kInt32, "q2k_idx / q2k_num must be int32"); + CHECK_DEVICE((*t)); CHECK_CONTIGUOUS((*t)); + } + TORCH_CHECK(q2k_idx.dim() == 4, "q2k_idx must be [batch, heads, num_m_blocks, max_kv_blocks]"); + TORCH_CHECK(q2k_idx.size(0) == batch_size && q2k_idx.size(1) == num_heads && q2k_idx.size(2) == num_m_blocks, + "q2k_idx leading dims must be [batch, heads, ceil(seqlen_q / BLOCK_M)]"); + TORCH_CHECK(q2k_idx.size(3) >= 1 && q2k_idx.size(3) <= num_n_blocks, "q2k_idx last dim must be in [1, num_kv_blocks]"); + CHECK_SHAPE(q2k_num, batch_size, num_heads, num_m_blocks); + SparseKvLists sparse; + sparse.q2k_idx = q2k_idx.data_ptr(); + sparse.q2k_num = q2k_num.data_ptr(); + sparse.q2k_max = q2k_idx.size(3); + sparse.num_m_blocks = num_m_blocks; + if (kv_valid_.has_value()) { + auto const &kv_valid = kv_valid_.value(); + TORCH_CHECK(kv_valid.scalar_type() == torch::kInt32, "kv_valid must be int32"); + CHECK_DEVICE(kv_valid); CHECK_CONTIGUOUS(kv_valid); + CHECK_SHAPE(kv_valid, 2 * num_n_blocks); + sparse.kv_valid = kv_valid.data_ptr(); + } + if (q2k_quad_.has_value()) { + auto const &q2k_quad = q2k_quad_.value(); + TORCH_CHECK(q2k_quad.scalar_type() == torch::kUInt8, "q2k_quad must be uint8"); + CHECK_DEVICE(q2k_quad); CHECK_CONTIGUOUS(q2k_quad); + TORCH_CHECK(q2k_quad.sizes() == q2k_idx.sizes(), "q2k_quad must match q2k_idx's shape"); + sparse.q2k_quad = q2k_quad.data_ptr(); + } + return mha_fwd_impl(q, k, v, sfq, sfk, sfv, delta_s, unpadded_k, out_, softmax_scale, + /*is_causal=*/false, per_block_mean, is_bf16, single_level_p_quant, sparse); +} + PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.doc() = "FlashAttention"; m.def("fwd", &mha_fwd, "Forward pass"); + m.def("fwd_sparse", &mha_fwd_sparse, "Block-sparse forward pass (non-causal)"); } \ No newline at end of file diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h b/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h index 675608726a..d237499c34 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/kernel_ws.h @@ -185,14 +185,14 @@ __global__ void __launch_bounds__(Ktraits::kNWarps * cutlass::NumThreadsPerWarp, auto block_coord = work_tile_info.get_block_coord(scheduler_params); auto [m_block, bidh, bidb] = block_coord; - int n_block_max = collective_mainloop.get_n_block_max(mainloop_params, m_block); - if (Is_causal && n_block_max <= 0) { // We exit early and write 0 to gO and -inf to gLSE. + int n_block_count = collective_mainloop.get_n_block_count(mainloop_params, m_block, bidh, bidb); + if (Is_causal && n_block_count <= 0) { // We exit early and write 0 to gO and -inf to gLSE. collective_epilogue.store_zero(epilogue_params, threadIdx.x - NumCopyThreads, block_coord); continue; } collective_mainloop.mma(mainloop_params, pipeline_q, pipeline_k, pipeline_v, smem_pipe_read_q, smem_pipe_read_k, smem_pipe_read_v, - tOrO, softmax_fused, n_block_max, threadIdx.x - NumCopyThreads, work_idx, m_block, shared_storage); + tOrO, softmax_fused, n_block_count, threadIdx.x - NumCopyThreads, work_idx, m_block, bidh, bidb, shared_storage); barrier_o.wait(); collective_epilogue.mma_store(shared_storage, tiled_mma_pv, tOrO, threadIdx.x - NumCopyThreads); barrier_o.arrive(); diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/launch.h b/fastvideo-kernel/attn_qat_infer/blackwell/launch.h index 91398c5f44..edc167afca 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/launch.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/launch.h @@ -62,7 +62,9 @@ void run_flash_fwd(Flash_fwd_params ¶ms, cudaStream_t stream) { static_cast(params.delta_s_ptr), {params.seqlen_s, params.seqlen_k, params.h_k, params.b}, {params.ds_row_stride, _1{}, params.ds_head_stride, params.ds_batch_stride}, - params.scale_softmax_log2 + params.scale_softmax_log2, + params.q2k_idx, params.q2k_num, params.q2k_max, + params.num_m_blocks, params.h, params.kv_valid, params.q2k_quad }); typename CollectiveEpilogue::Params epilogue_params = CollectiveEpilogue::to_underlying_arguments({ diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h b/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h index fa2240422c..51d9ff23a3 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/mainloop_tma_ws.h @@ -174,6 +174,13 @@ struct CollectiveMainloopFwd { ShapeQKV const shape_ds; StrideQKV const stride_ds; float const softmax_scale_log2; + int const* ptr_q2k_idx{nullptr}; + int const* ptr_q2k_num{nullptr}; + int q2k_max{0}; + int num_m_blocks{0}; + int num_heads{0}; + int const* ptr_kv_valid{nullptr}; + uint8_t const* ptr_q2k_quad{nullptr}; }; // Device side kernel params @@ -194,6 +201,13 @@ struct CollectiveMainloopFwd { TMA_SFVt tma_load_SFVt; TMA_DS tma_load_DS; float const softmax_scale_log2; + int const* ptr_q2k_idx; + int const* ptr_q2k_num; + int q2k_max; + int num_m_blocks; + int num_heads; + int const* ptr_kv_valid; + uint8_t const* ptr_q2k_quad; }; @@ -261,7 +275,9 @@ struct CollectiveMainloopFwd { tma_load_K, tma_load_sfk, tma_load_Vt, tma_load_sfvt, tma_load_ds, - args.softmax_scale_log2}; + args.softmax_scale_log2, + args.ptr_q2k_idx, args.ptr_q2k_num, args.q2k_max, + args.num_m_blocks, args.num_heads, args.ptr_kv_valid, args.ptr_q2k_quad}; } /// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance @@ -290,6 +306,51 @@ struct CollectiveMainloopFwd { return n_block_max; } + // Number of KV blocks query block m_block visits: all of them when dense, + // its index-list length when block-sparse. + CUTLASS_DEVICE + int get_n_block_count(Params const& mainloop_params, int m_block, int bidh, int bidb) { + if (mainloop_params.ptr_q2k_idx == nullptr) { + return get_n_block_max(mainloop_params, m_block); + } + return mainloop_params.ptr_q2k_num[(bidb * mainloop_params.num_heads + bidh) * mainloop_params.num_m_blocks + m_block]; + } + + // KV block visited at iteration i (iterations run from count-1 down to 0). + CUTLASS_DEVICE + int get_kv_block(Params const& mainloop_params, int m_block, int bidh, int bidb, int i) { + if (mainloop_params.ptr_q2k_idx == nullptr) { + return i; + } + int64_t const row = (int64_t(bidb) * mainloop_params.num_heads + bidh) * mainloop_params.num_m_blocks + m_block; + return mainloop_params.ptr_q2k_idx[row * mainloop_params.q2k_max + i]; + } + + // Valid key columns of each 64-column half of KV block n_block (valid tokens + // first within a half). kv_valid, when given, stores two counts per block so + // 64-token VSA tiles can pad mid-block; otherwise only the sequence tail is + // masked. + CUTLASS_DEVICE + int get_kv_valid_half(Params const& mainloop_params, int n_block, int half, int unpadded_seqlen_k) { + static constexpr int kBlockN = get<1>(TileShape_MNK{}); + static constexpr int kHalfN = kBlockN / 2; + if (mainloop_params.ptr_kv_valid != nullptr) { + return mainloop_params.ptr_kv_valid[2 * n_block + half]; + } + return max(0, min(kHalfN, unpadded_seqlen_k - n_block * kBlockN - half * kHalfN)); + } + + // Quadrant mask of list entry i: bit (2 * row_half + col_half) is set when the + // 64-row query half attends the 64-column key half. 0xF when absent. + CUTLASS_DEVICE + int get_quad(Params const& mainloop_params, int m_block, int bidh, int bidb, int i) { + if (mainloop_params.ptr_q2k_quad == nullptr) { + return 0xF; + } + int64_t const row = (int64_t(bidb) * mainloop_params.num_heads + bidh) * mainloop_params.num_m_blocks + m_block; + return mainloop_params.ptr_q2k_quad[row * mainloop_params.q2k_max + i]; + } + template CUTE_HOST_DEVICE constexpr auto @@ -450,7 +511,7 @@ struct CollectiveMainloopFwd { auto [m_block, bidh, bidb] = work_tile_info.get_block_coord(scheduler_params); - int n_block_max = get_n_block_max(mainloop_params, m_block); + int n_block_count = get_n_block_count(mainloop_params, m_block, bidh, bidb); Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.begin()), SmemLayoutQ{}); Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.begin()), SmemLayoutK{}); @@ -506,7 +567,8 @@ struct CollectiveMainloopFwd { Tensor tDSsDS = group_modes<0, 3>(block_tma_ds.partition_D(sDS)); uint16_t mcast_mask_kv = 0; - int n_block = n_block_max - 1; + int n_iter = n_block_count - 1; + int n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); int lane_predicate = cute::elect_one_sync(); if (lane_predicate) { pipeline_q.producer_acquire(smem_pipe_write_q); @@ -529,11 +591,12 @@ struct CollectiveMainloopFwd { ++smem_pipe_write_v; } - n_block--; + --n_iter; if (lane_predicate) { // CUTLASS_PRAGMA_NO_UNROLL #pragma unroll 2 - for (; n_block >= 0; --n_block) { + for (; n_iter >= 0; --n_iter) { + n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); pipeline_k.producer_acquire(smem_pipe_write_k); copy(mainloop_params.tma_load_K.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv), tKgK(_, n_block), tKsK(_, smem_pipe_write_k.index())); @@ -585,6 +648,8 @@ struct CollectiveMainloopFwd { int thread_idx, int work_idx, int m_block, + int bidh, + int bidb, SharedStorage& shared_storage ) { @@ -667,7 +732,21 @@ struct CollectiveMainloopFwd { int const seqlen_q = get<0>(mainloop_params.shape_Q); int const seqlen_k = get<0>(mainloop_params.shape_K); int const unpadded_seqlen_k = get<0>(mainloop_params.unpadded_shape_K); - int n_block = n_block_count - 1; + int n_iter = n_block_count - 1; + int n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); + bool const per_block_masking = mainloop_params.ptr_q2k_idx != nullptr || mainloop_params.ptr_kv_valid != nullptr; + static_assert(kBlockM == 128 && kBlockN == 128, "quadrant masking assumes 128x128 blocks"); + // Each MMA warp owns 16 consecutive query rows, so a warp lies in one + // 64-row half. With quadrant lists a warp skips blocks its half did not + // select and the P.V chunk of a key half it did not select: masked + // scores contribute exactly zero, and the warp sharing its tensor-core + // partition runs faster meanwhile. + int const my_row_half = __shfl_sync(0xffffffff, [&] { + Tensor cS0 = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{})); + Tensor tScS0 = thread_mma_qk.partition_C(cS0); + return int(get<0>(tScS0(0))) >= kBlockM / 2 ? 1 : 0; + }(), 0); + bool const quad_skip = mainloop_params.ptr_q2k_quad != nullptr; auto copy_k_block = [&](auto block_id) { auto tSsK_stage = tSsK(_, _, _, smem_pipe_read_k.index()); @@ -744,6 +823,26 @@ struct CollectiveMainloopFwd { int local = c & 31; return (c & ~31) | (local & 1) | ((local & 24) >> 2) | ((local & 6) << 2); }; + // Sparse lists may visit a block whose 64-token halves are partially valid + // (VSA tile tails, short prefix chunks) or that only one query half + // selected (64-token tiles), at any iteration. + auto apply_sparse_mask = [&](auto& acc, int n_blk, int it) { + int const quad = get_quad(mainloop_params, m_block, bidh, bidb, it); + int const valid0 = get_kv_valid_half(mainloop_params, n_blk, 0, unpadded_seqlen_k); + int const valid1 = get_kv_valid_half(mainloop_params, n_blk, 1, unpadded_seqlen_k); + if (quad == 0xF && valid0 == kBlockN / 2 && valid1 == kBlockN / 2) { return; } + Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{})); + Tensor tScS = thread_mma_qk.partition_C(cS); + CUTLASS_PRAGMA_UNROLL + for (int i = 0; i < size(acc); ++i) { + int const col = actual_col(int(get<1>(tScS(i)))); + int const col_half = col >= kBlockN / 2; + int const row_half = int(get<0>(tScS(i))) >= kBlockM / 2; + bool const keep = ((quad >> (2 * row_half + col_half)) & 1) + && (col & (kBlockN / 2 - 1)) < (col_half ? valid1 : valid0); + if (!keep) { acc(i) = -INFINITY; } + } + }; { Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{})); Tensor tScS = thread_mma_qk.partition_C(cS); @@ -751,7 +850,9 @@ struct CollectiveMainloopFwd { for (int i = 0; i < size(tSrS); ++i) { int col = actual_col(int(get<1>(tScS(i)))); if constexpr (!Is_causal) { // Just masking based on col - if (col >= int(unpadded_seqlen_k - n_block * kBlockN)) { tSrS(i) = -INFINITY; } + if (!per_block_masking) { + if (col >= int(unpadded_seqlen_k - n_block * kBlockN)) { tSrS(i) = -INFINITY; } + } } else { if (col >= std::min(seqlen_k - n_block * kBlockN, col_limit_causal(int(get<0>(tScS(i))), n_block))) { @@ -760,6 +861,9 @@ struct CollectiveMainloopFwd { } } } + if constexpr (!Is_causal) { + if (per_block_masking) { apply_sparse_mask(tSrS, n_block, n_iter); } + } auto quantize = [&](auto mma_k, auto acc_conversion_view) { Tensor AbsMaxP_stagek = AbsMaxP(_, make_coord(_, _, mma_k)); Tensor acc_conversion_stagek = acc_conversion_view(_, _, mma_k); @@ -827,11 +931,12 @@ struct CollectiveMainloopFwd { } } - n_block--; + --n_iter; constexpr int n_masking_steps = !Is_causal ? 1 : cute::ceil_div(kBlockM, kBlockN) + 1; // // Only go through these if Is_causal, since n_masking_steps = 1 when !Is_causal CUTLASS_PRAGMA_UNROLL - for (int masking_step = 0; masking_step < n_masking_steps - 1 && n_block >= 0; ++masking_step, --n_block) { + for (int masking_step = 0; masking_step < n_masking_steps - 1 && n_iter >= 0; ++masking_step, --n_iter) { + n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{})); Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout())); consumer_wait(pipeline_k, smem_pipe_read_k); @@ -878,7 +983,22 @@ struct CollectiveMainloopFwd { } #pragma unroll 1 - for (; n_block >= 0; --n_block) { + for (; n_iter >= 0; --n_iter) { + n_block = get_kv_block(mainloop_params, m_block, bidh, bidb, n_iter); + int const row_bits = quad_skip + ? (get_quad(mainloop_params, m_block, bidh, bidb, n_iter) >> (2 * my_row_half)) & 3 + : 3; + if (row_bits == 0) { + // Neither key half is selected for this warp's rows: its softmax + // state and output are unchanged. Keep the pipeline in step. + consumer_wait(pipeline_k, smem_pipe_read_k); + pipeline_k.consumer_release(smem_pipe_read_k); + ++smem_pipe_read_k; + consumer_wait(pipeline_v, smem_pipe_read_v); + pipeline_v.consumer_release(smem_pipe_read_v); + ++smem_pipe_read_v; + continue; + } Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{})); Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout())); consumer_wait(pipeline_k, smem_pipe_read_k); @@ -896,16 +1016,21 @@ struct CollectiveMainloopFwd { } } + if (per_block_masking) { apply_sparse_mask(tSrS, n_block, n_iter); } softmax_fused.template online_softmax_with_quant(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2); Tensor tOrO = make_fragment_like(tOrO_store); + clear(tOrO); consumer_wait(pipeline_v, smem_pipe_read_v); copy_v_block(_0{}); quantize(_0{}, tSrS_converion_view); CUTLASS_PRAGMA_UNROLL for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) { - cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)), - make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO); + // v_block spans one 64-column key half (P's K mode is 2 x 64). + if ((row_bits >> v_block) & 1) { + cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)), + make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO); + } if (v_block < size<2>(tOrP) - 1) { copy_v_block(v_block + 1); quantize(v_block + 1, tSrS_converion_view); diff --git a/fastvideo-kernel/attn_qat_infer/blackwell/params.h b/fastvideo-kernel/attn_qat_infer/blackwell/params.h index fcf9bf8981..867aa7b257 100644 --- a/fastvideo-kernel/attn_qat_infer/blackwell/params.h +++ b/fastvideo-kernel/attn_qat_infer/blackwell/params.h @@ -114,6 +114,23 @@ int * __restrict__ seqused_k; int *__restrict__ blockmask; + + // Block-sparse KV iteration (non-causal only). When q2k_idx is null the + // kernel is dense. Otherwise query block m of (batch b, head h) visits + // the q2k_num[(b*h + h)*num_m_blocks + m] KV blocks listed at + // q2k_idx[((b*h + h)*num_m_blocks + m)*q2k_max + i], in any order. + int const *__restrict__ q2k_idx; + int const *__restrict__ q2k_num; + int q2k_max; + int num_m_blocks; + // Optional valid token count of each 64-column half of every KV block, + // kv_valid[2*n + half] (valid tokens first within a half); null means only + // the sequence tail beyond the unpadded key length is masked. + int const *__restrict__ kv_valid; + // Optional quadrant mask per list entry (same layout as q2k_idx): bit + // (2*row_half + col_half) set when that 64-row query half attends that + // 64-column key half. Lets 64-token VSA tiles run on 128x128 blocks. + uint8_t const *__restrict__ q2k_quad; // The K_new and V_new matrices. void * __restrict__ knew_ptr; diff --git a/fastvideo/api/schema.py b/fastvideo/api/schema.py index 2fda230ebb..521039785f 100644 --- a/fastvideo/api/schema.py +++ b/fastvideo/api/schema.py @@ -78,6 +78,8 @@ class QuantizationConfig: transformer_quant: str | None = None # Forwarded onto ``NVFP4Config`` when ``transformer_quant`` is ``NVFP4``. # ``h3_dit`` selects the packed MiniMax-H3 attention+FFN export. + # ``h3_dit_ffn`` selects a packed FFN-only export (attention stays dense). + # ``h3_dit_vsa`` is ``h3_dit`` plus the VSA compression gates. layer_profile: str | None = None diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 04558b7fa1..3c30f3b921 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -80,6 +80,7 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: _MINIMAX_H3_NVFP4_FF_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.ff\.(?:fc_in|fc_out)$") _MINIMAX_H3_NVFP4_DIT_PREFIX = re.compile( r"(?:^|\.)transformer_blocks\.\d+\.(?:attn\.to_(?:q|k|v|out)|ff\.(?:fc_in|fc_out))$") +_MINIMAX_H3_NVFP4_VSA_GATE_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.attn\.to_gate_compress$") H3_NVFP4_DIT_EXPORT_FILENAME = "nvfp4_weights.safetensors" H3_NVFP4_DIT_KEY_SEP = "::" H3_NVFP4_DIT_BUFFER_NAMES = ( @@ -335,6 +336,16 @@ def _mm_fp4( ) +def _mm_fp4_backend() -> str: + """FlashInfer ``mm_fp4`` backend (``FASTVIDEO_NVFP4_MM_BACKEND``, default ``auto``). + + On sm_120 ``auto`` picks a kernel about 2x slower than ``cutlass`` or + ``cudnn`` once activations reach tens of thousands of rows (measured at + 73k rows on an RTX PRO 6000); short sequences are unaffected. + """ + return os.environ.get("FASTVIDEO_NVFP4_MM_BACKEND", "auto") + + def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: """Coerce an activation to a dtype the FP4 linear accepts. @@ -462,7 +473,7 @@ def apply( alpha, torch.bfloat16, None, - backend='auto', + backend=_mm_fp4_backend(), ) if bias is not None: @@ -480,13 +491,18 @@ class NVFP4Config(QuantizationConfig): ``fc_in`` and ``fc_out`` in each main transformer-block FFN. ``layer_profile="h3_dit"`` expands that to the packed NVFP4H3 DiT set (attention ``to_{q,k,v,out}`` plus those FFN linears). + ``layer_profile="h3_dit_ffn"`` loads a packed export of the FFN linears + only, keeping attention projections dense (e.g. calibrated FFN-only + checkpoints such as FastH3 V2 NVFP4). + ``layer_profile="h3_dit_vsa"`` is ``h3_dit`` plus each block's VSA + compression gate ``attn.to_gate_compress`` (VSA-distilled students). """ def __init__(self, layer_profile: str = "refine", retain_original_weights: bool | None = None): super().__init__() - if layer_profile not in ("base", "refine", "h3_dit"): - raise ValueError("NVFP4Config.layer_profile must be one of 'base', 'refine', or 'h3_dit', " - f"got {layer_profile!r}") + if layer_profile not in ("base", "refine", "h3_dit", "h3_dit_ffn", "h3_dit_vsa"): + raise ValueError("NVFP4Config.layer_profile must be one of 'base', 'refine', 'h3_dit', " + f"'h3_dit_ffn', or 'h3_dit_vsa', got {layer_profile!r}") self.layer_profile = layer_profile # Original bf16 ``layer.weight`` retention after FP4 conversion. # Default (None/False): purge the purgeable originals -- every @@ -525,6 +541,11 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str): return None if self.layer_profile == "h3_dit": tagged = is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + elif self.layer_profile == "h3_dit_vsa": + tagged = (is_minimax_h3_nvfp4_dit_linear_prefix(prefix) + or _MINIMAX_H3_NVFP4_VSA_GATE_PREFIX.search(prefix) is not None) + elif self.layer_profile == "h3_dit_ffn": + tagged = is_minimax_h3_nvfp4_linear_prefix(prefix) else: tagged = is_ltx2_nvfp4_linear_prefix(prefix) or is_minimax_h3_nvfp4_linear_prefix(prefix) if tagged: diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index e25650466d..a15d233fa0 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -14,6 +14,8 @@ from fastvideo import envs from fastvideo.attention import DistributedAttention +from fastvideo.attention.backends.abstract import layer_idx_from_prefix +from fastvideo.attention.backends.video_sparse_attn_h3 import MiniMaxH3VSAMetadata from fastvideo.attention.layer import DistributedAttention_VSA from fastvideo.attention.selector import get_attn_backend from fastvideo.configs.models.dits.minimax_h3 import MiniMaxH3Config @@ -21,7 +23,9 @@ sequence_model_parallel_all_gather_with_unpad, sequence_model_parallel_shard, ) -from fastvideo.distributed.parallel_state import get_sp_world_size, model_parallel_is_initialized +from fastvideo.distributed.parallel_state import (get_sp_group, get_sp_world_size, + model_parallel_is_initialized) +from fastvideo.forward_context import get_forward_context from fastvideo.layers.linear import ReplicatedLinear from fastvideo.layers.quantization.mxfp8_config import MXFP8QuantizeMethod from fastvideo.layers.mlp import MLP @@ -30,6 +34,8 @@ from fastvideo.layers.visual_embedding import Timesteps from fastvideo.logger import init_logger from fastvideo.models.dits.base import BaseDiT +from fastvideo.models.dits.minimax_h3_vsa_fp4 import (vsa_fp4_attention, vsa_fp4_attention_sp, + vsa_fp4_requested) from fastvideo.models.dits.minimax_h3_fusions import ( HAVE_TRITON, fused_qknorm_rope, @@ -74,6 +80,13 @@ def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool: return HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() +@torch.compile(dynamic=True, fullgraph=True) +def _gated_residual(hidden_states: torch.Tensor, gate_table: torch.Tensor, indices: torch.Tensor, + branch: torch.Tensor) -> torch.Tensor: + """``hidden + gate[indices] * branch`` in one pass (eager materializes the gathered gate).""" + return hidden_states + gate_table.index_select(0, indices) * branch + + class MiniMaxH3RotaryPosEmbed(nn.Module): """Three-axis rotary frequencies over packed `(t, h, w)` coordinates.""" @@ -209,6 +222,10 @@ def __init__( prefix=prefix, fa4_packed_varlen=fa4_packed_varlen, ) + # Opt-in inference route: VSA-H3 selection on the block-sparse FP4 + # 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.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. @@ -255,6 +272,12 @@ def _resolve_gate_compress_for_compile(self) -> None: return if self._gate_compress_active is None: weight = self.to_gate_compress.weight + if weight is None: + # Packed NVFP4 gate: any nonzero E2M1 magnitude (bits 0x7 of + # either nibble) makes the branch live. + packed = self.to_gate_compress._nvfp4_weight + self._gate_compress_active = bool((packed & 0x77).any()) + return # bool() on a DTensor reduction resolves collectively, so every # rank caches the same answer. self._gate_compress_active = bool((weight != 0).any()) @@ -282,6 +305,20 @@ def forward( rotary_emb: tuple[torch.Tensor, torch.Tensor] | None, original_seq_len: int, ) -> torch.Tensor: + if (self._vsa_fp4 and rotary_emb is not None and not torch.is_grad_enabled() + and not torch.compiler.is_compiling()): + meta = get_forward_context().attn_metadata + # Exempt mode lists the first prefix tile for every query, which the + # FP4 kernel relies on to start each row from a finite running max. + if isinstance(meta, MiniMaxH3VSAMetadata) and meta.exempt: + use_fused_rope = self.fuse_qknorm_rope and _can_run_minimax_h3_fusion(hidden_states) + if not model_parallel_is_initialized() or get_sp_world_size() == 1: + hidden_states = vsa_fp4_attention(self, hidden_states, rotary_emb, meta, use_fused_rope) + else: + hidden_states = vsa_fp4_attention_sp(self, hidden_states, rotary_emb, meta, use_fused_rope, + get_sp_group()) + 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) @@ -566,6 +603,8 @@ def forward( 1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices) with nvtx_range("minimax_h3.transformer_block.feed_forward"): feed_forward_output = self.ff(norm_hidden_states) + if use_modulate_fusion and not torch.compiler.is_compiling(): + return _gated_residual(hidden_states, gate_mlp, adaln_indices, feed_forward_output) return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py new file mode 100644 index 0000000000..c83e6f7357 --- /dev/null +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -0,0 +1,282 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Inference fast path: VSA-H3 attention on the block-sparse FP4 kernel. + +Opt-in with ``FASTVIDEO_H3_VSA_FP4=1`` (no-grad, single sequence-parallel +rank, ``fastvideo-kernel`` built with ``attn_qat_infer``). The selection is +VSA-H3's own: tile pooling, top-k block mask and the gated compression branch +are unchanged; only the block-sparse attention itself runs on SageAttention3's +FP4 kernel (BF16 Triton otherwise), with 64-token tiles carried by quadrant +masks on the kernel's 128x128 blocks. + +The attention input is gathered into tile order once per block (one +``hidden_size``-wide pass; pad rows stay zero, so q/k/v pad rows are exactly +zero through the bias-free projections, RMSNorm and RoPE). That replaces the +generic path's concat, four tile scatters and three transposes, and lets q, k +and v share one activation quantization. The output returns to packed order +with one gather before ``to_out``. +""" + +from __future__ import annotations + +import math +import os +from typing import Any + +import torch + +from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAMetadata, _build_block_mask, _pool_tiles) + +VSA_FP4_ENV = "FASTVIDEO_H3_VSA_FP4" +_BLOCK = 128 + +_fp4_api: Any = None + + +def vsa_fp4_requested() -> bool: + return os.environ.get(VSA_FP4_ENV, "0") == "1" + + +def _api() -> Any: + global _fp4_api + if _fp4_api is None: + import attn_qat_infer.api as api + _fp4_api = api + return _fp4_api + + +class _TileLayout: + """Per-step tile-order state shared by every block of one forward.""" + + def __init__(self, meta: MiniMaxH3VSAMetadata, rotary_emb: tuple[torch.Tensor, torch.Tensor]) -> None: + self.tile = int(meta.tile_elems) + self.n_tiles = int(meta.variable_block_sizes.numel()) + self.rows = math.ceil(self.n_tiles * self.tile / _BLOCK) * _BLOCK + self.untile = meta.untile_combined_index + self.row_tile = self.untile // self.tile + self.rotary_src = rotary_emb + cos, sin = rotary_emb + self.cos = cos.new_zeros((self.rows, cos.shape[-1])).index_copy_(0, self.untile, cos) + self.sin = sin.new_zeros((self.rows, sin.shape[-1])).index_copy_(0, self.untile, sin) + self._buf: torch.Tensor | None = None + + def gather_in(self, x: torch.Tensor) -> torch.Tensor: + """Packed ``[B, L, C]`` -> tile-ordered ``[B, rows, C]``; pad rows stay zero. + + The buffer is reused across blocks: pad rows are never written and + every valid row is overwritten, and each block consumes it (q/k/v + projections) before the next block refills it. + """ + shape = (x.shape[0], self.rows, x.shape[-1]) + if self._buf is None or self._buf.shape != shape or self._buf.dtype != x.dtype: + self._buf = x.new_zeros(shape) + return self._buf.index_copy_(1, self.untile, x) + + +def _layout_for(meta: MiniMaxH3VSAMetadata, rotary_emb: tuple[torch.Tensor, torch.Tensor]) -> _TileLayout: + layout = getattr(meta, "_h3_fp4_layout", None) + if layout is None or layout.rotary_src[0] is not rotary_emb[0]: + layout = _TileLayout(meta, rotary_emb) + meta._h3_fp4_layout = layout # type: ignore[attr-defined] + return layout + + +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. + + NVFP4 activations use a unit global scale for every layer, so one + quantized copy is exactly what each layer would have produced. + """ + from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod + + methods = [linear.quant_method for linear in linears] + if not all(type(m) is NVFP4QuantizeMethod and m.wants_prequantized_input() for m in methods): + 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_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]``.""" + api = _api() + layout = _layout_for(meta, rotary_emb) + heads, dim = attn.num_attention_heads, attn.attention_head_dim + x_tiles = layout.gather_in(hidden_states) + 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)) + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + cos, sin = layout.cos.to(query.dtype), layout.sin.to(query.dtype) + query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) + else: + rope = (layout.cos, layout.sin) + query = attn._apply_rotary_emb(attn.norm_q(query), rope) + key = attn._apply_rotary_emb(attn.norm_k(key), rope) + + vbs = meta.variable_block_sizes + logical = layout.n_tiles * layout.tile + q_pooled = _pool_tiles(query[:, :logical], vbs, layout.tile) + k_pooled = _pool_tiles(key[:, :logical], vbs, layout.tile) + 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) + q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) + out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) + out = out.transpose(1, 2).index_select(1, layout.untile) # [B, L, H, D], packed order + + if attn.to_gate_compress is not None and attn._gate_active(): + gate, _ = attn.to_gate_compress(hidden_states) + v_pooled = _pool_tiles(value[:, :logical], vbs, layout.tile) + out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled).permute(0, 2, 1, 3).to(out.dtype) + out = out.addcmul_(out_c.index_select(1, layout.row_tile), gate.unflatten(-1, (heads, dim))) + return out.flatten(2, 3) + + +# --------------------------------------------------------------------------- +# Ulysses sequence parallelism: FP8 head/sequence exchange around the FP4 core +# --------------------------------------------------------------------------- +# +# PCIe-only boxes (e.g. 8x RTX PRO 6000) move ~21 GB/s per GPU in an +# all-to-all, so the Ulysses exchange, not compute, bounds multi-GPU latency. +# q/k/v travel as FP8 (one scale per token and head): the attention kernel +# re-quantizes them to FP4 on arrival, so FP8 transport adds error far below +# that floor while halving the bytes. The VSA gate never travels: each rank +# applies it to its own sequence rows after an all-gather of the tiny per-tile +# compression output. The attention output returns as FP8 too; ``to_out`` +# re-quantizes it to FP4. + +_FP8 = torch.float8_e4m3fn +_FP8_MAX = 448.0 + + +@torch.compile(dynamic=False, fullgraph=True) +def _pack_heads_fp8(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, + world: int) -> tuple[torch.Tensor, torch.Tensor]: + """``[rows, H, D]`` x3 -> per-destination ``[W, 3, rows, H/W, D]`` FP8 payload and scales.""" + x = torch.stack((query, key, value)).float() + scale = (x.abs().amax(dim=-1) / _FP8_MAX).clamp_min(1e-12) + payload = (x / scale[..., None]).to(_FP8) + _, rows, heads, dim = x.shape + payload = payload.view(3, rows, world, heads // world, dim).permute(2, 0, 1, 3, 4).contiguous() + scale = scale.view(3, rows, world, heads // world).permute(2, 0, 1, 3).contiguous() + return payload, scale + + +@torch.compile(dynamic=False, fullgraph=True) +def _unpack_seq_fp8(payload: torch.Tensor, scale: torch.Tensor, seq_len: int) -> torch.Tensor: + """``[W, 3, rows, Hs, D]`` from every source rank -> ``[3, seq_len, Hs, D]`` BF16 in packed order.""" + x = payload.float() * scale[..., None] + world, _, rows, heads, dim = x.shape + return x.permute(1, 0, 2, 3, 4).reshape(3, world * rows, heads, dim)[:, :seq_len].to(torch.bfloat16) + + +@torch.compile(dynamic=False, fullgraph=True) +def _pack_seq_fp8(out_bhsd: torch.Tensor, untile: torch.Tensor, world: int, + rows: int) -> tuple[torch.Tensor, torch.Tensor]: + """Tile-ordered ``[1, Hs, R, D]`` -> packed, padded, per-destination ``[W, rows, Hs, D]`` FP8 + scales.""" + x = out_bhsd[0].transpose(0, 1).index_select(0, untile).float() # [L, Hs, D] + x = torch.nn.functional.pad(x, (0, 0, 0, 0, 0, world * rows - x.shape[0])) + scale = (x.abs().amax(dim=-1) / _FP8_MAX).clamp_min(1e-12) + payload = (x / scale[..., None]).to(_FP8) + return payload.view(world, rows, *payload.shape[1:]), scale.view(world, rows, scale.shape[-1]) + + +@torch.compile(dynamic=False, fullgraph=True) +def _unpack_heads_fp8(payload: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + """``[W, rows, Hs, D]`` from every head group -> ``[rows, W*Hs, D]`` BF16.""" + x = payload.float() * scale[..., None] + world, rows, heads, dim = x.shape + return x.permute(1, 0, 2, 3).reshape(rows, world * heads, dim).to(torch.bfloat16) + + +@torch.compile(dynamic=False, fullgraph=True) +def _apply_gate(out: torch.Tensor, out_c: torch.Tensor, row_tile: torch.Tensor, gate: torch.Tensor) -> torch.Tensor: + """``out + out_c[row_tile] * gate`` for one rank's rows; out/gate ``[rows, H, D]``, out_c ``[n_tiles, H, D]``.""" + return out + out_c.index_select(0, row_tile) * gate + + +def _all_to_all(payload: torch.Tensor, scale: torch.Tensor, group: Any) -> tuple[torch.Tensor, torch.Tensor]: + import torch.distributed as dist + recv = torch.empty_like(payload) + recv_scale = torch.empty_like(scale) + dist.all_to_all_single(recv.view(torch.uint8), payload.view(torch.uint8), group=group) + dist.all_to_all_single(recv_scale, scale, group=group) + return recv, recv_scale + + +class _SPTileLayout: + """Per-step tile-order state for one rank's head subset.""" + + def __init__(self, meta: MiniMaxH3VSAMetadata, rank: int, local_rows: int) -> None: + self.tile = int(meta.tile_elems) + self.n_tiles = int(meta.variable_block_sizes.numel()) + self.rows = math.ceil(self.n_tiles * self.tile / _BLOCK) * _BLOCK + self.seq_len = int(meta.total_seq_length) + self.untile = meta.untile_combined_index + row_tile = self.untile // self.tile + local = torch.arange(rank * local_rows, (rank + 1) * local_rows, device=row_tile.device) + # Rows past the sequence are SP padding; their outputs are discarded. + self.local_row_tile = row_tile[local.clamp_max(self.seq_len - 1)] + self._buf: torch.Tensor | None = None + + def tiles_from(self, qkv: torch.Tensor) -> torch.Tensor: + """``[3, L, Hs, D]`` packed -> ``[3, rows, Hs, D]`` tile order with zero padding (reused buffer).""" + shape = (3, self.rows, *qkv.shape[2:]) + if self._buf is None or self._buf.shape != shape: + self._buf = qkv.new_zeros(shape) + return self._buf.index_copy_(1, self.untile, qkv) + + +def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], + meta: MiniMaxH3VSAMetadata, use_fused_rope: bool, sp_group: Any) -> torch.Tensor: + """Ulysses-SP attention core on local sequence rows ``[1, rows, C]``; returns pre-``to_out`` ``[1, rows, H*D]``.""" + import torch.distributed as dist + + api = _api() + world, rank = sp_group.world_size, sp_group.rank_in_group + heads, dim = attn.num_attention_heads, attn.attention_head_dim + local_rows = hidden_states.shape[1] + layout = getattr(meta, "_h3_fp4_sp_layout", None) + if layout is None: + layout = _SPTileLayout(meta, rank, local_rows) + meta._h3_fp4_sp_layout = layout # type: ignore[attr-defined] + + query, key, value = (t.unflatten(-1, (heads, dim)) + for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), hidden_states)) + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + cos, sin = rotary_emb[0].to(query.dtype), rotary_emb[1].to(query.dtype) + query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) + else: + query = attn._apply_rotary_emb(attn.norm_q(query), rotary_emb) + key = attn._apply_rotary_emb(attn.norm_k(key), rotary_emb) + + payload, scale = _pack_heads_fp8(query[0], key[0], value[0], world) + payload, scale = _all_to_all(payload, scale, sp_group.device_group) + qkv = layout.tiles_from(_unpack_seq_fp8(payload, scale, layout.seq_len)) # [3, R, Hs, D] + q_t, k_t, v_t = qkv[0:1], qkv[1:2], qkv[2:3] + + vbs = meta.variable_block_sizes + logical = layout.n_tiles * layout.tile + scores = torch.matmul(_pool_tiles(q_t[:, :logical], vbs, layout.tile), + _pool_tiles(k_t[:, :logical], vbs, layout.tile).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) + q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) + out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad) + + payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows) + payload, scale = _all_to_all(payload, scale, sp_group.device_group) + out = _unpack_heads_fp8(payload, scale) # [rows, H, D] + + if attn.to_gate_compress is not None and attn._gate_active(): + v_pooled = _pool_tiles(v_t[:, :logical], vbs, layout.tile) + out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled)[0].to(out.dtype) # [Hs, n_tiles, D] + gathered = torch.empty((world, *out_c.shape), dtype=out_c.dtype, device=out_c.device) + dist.all_gather_into_tensor(gathered, out_c.contiguous(), group=sp_group.device_group) + out_c_all = gathered.flatten(0, 1).transpose(0, 1) # [n_tiles, H, D] + gate, _ = attn.to_gate_compress(hidden_states) + out = _apply_gate(out, out_c_all, layout.local_row_tile, gate[0].unflatten(-1, (heads, dim))) + return out.flatten(1, 2).unsqueeze(0) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 020a8f970a..911f878963 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -228,7 +228,9 @@ def maybe_load_fsdp_model( packed_candidate = find_minimax_h3_nvfp4_dit_export(weight_dir_list) weight_dir_list = dense_transformer_safetensors(weight_dir_list) quant_config = getattr(init_params.get("config"), "quant_config", None) - packed_nvfp4_export = packed_candidate if getattr(quant_config, "layer_profile", None) == "h3_dit" else None + packed_profiles = ("h3_dit", "h3_dit_ffn", "h3_dit_vsa") + packed_nvfp4_export = (packed_candidate + if getattr(quant_config, "layer_profile", None) in packed_profiles else None) if packed_nvfp4_export is not None and lora_path is not None: raise ValueError("Packed MiniMax-H3 NVFP4 DiT export cannot be combined with lora_path; " "merge the adapter before exporting, or load without the packed file.") diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index e116defb7f..21be0864a9 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -7,6 +7,7 @@ """ import math +import os from collections.abc import Iterator from dataclasses import dataclass @@ -522,6 +523,11 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: ) +def _tile_batch_size() -> int: + """Spatial tiles decoded per decoder call (``FASTVIDEO_H3_VAE_TILE_BATCH``, default 1 = per tile).""" + return max(1, int(os.environ.get("FASTVIDEO_H3_VAE_TILE_BATCH", "1"))) + + def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool: """Select the video decoder that serves the H3 VAE ``decode`` path.""" return name == "decoder" and isinstance(submodule, MiniMaxH3VideoViTDecoder3d) @@ -817,10 +823,28 @@ def _decode_clip(self, z: torch.Tensor) -> torch.Tensor: ratio = self.spatial_compression_ratio rows = [] + if _tile_batch_size() > 1 and len(set(y_lengths)) == 1 and len(set(x_lengths)) == 1: + # Every tile of the grid has one shape, and the ViT decoder + # treats batch entries independently: decode the grid in a few + # large calls instead of one small call per tile. + with nvtx_range("minimax_h3.vae.decode_clip.decode_tile_batches"): + latent_tiles = [ + z[..., y_position // ratio:(y_position + y_length) // ratio, + x_position // ratio:(x_position + x_length) // ratio] + for y_position, y_length in zip(y_indices, y_lengths) + for x_position, x_length in zip(x_indices, x_lengths) + ] + decoded: list[torch.Tensor] = [] + per_call = _tile_batch_size() + for start in range(0, len(latent_tiles), per_call): + batch = torch.cat(latent_tiles[start:start + per_call], dim=0) + decoded.extend(self.decoder(self._project_decoder_tile(batch)).split(z.shape[0], dim=0)) + columns = len(x_indices) + rows = [decoded[index:index + columns] for index in range(0, len(decoded), columns)] # The eager tile driver owns NVTX so each marker remains outside # the compiled decoder graph. with nvtx_range("minimax_h3.vae.decode_clip.decode_tiles"): - for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths)): + for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths) if not rows else ()): row = [] for column_index, (x_position, x_length) in enumerate(zip(x_indices, x_lengths)): with nvtx_range(f"minimax_h3.vae.decode_clip.tile.{row_index}.{column_index}"): diff --git a/scripts/benchmarks/minimax_h3_pro6000/a2a8.py b/scripts/benchmarks/minimax_h3_pro6000/a2a8.py new file mode 100644 index 0000000000..f2659517eb --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/a2a8.py @@ -0,0 +1,72 @@ +"""All-to-all bandwidth on 8x RTX PRO 6000 at FastH3 Ulysses SP8 payload sizes (10 s, 1344x768: 73.6k tokens).""" +import modal + +image = (modal.Image.debian_slim(python_version="3.12") + .pip_install("torch==2.8.0", index_url="https://download.pytorch.org/whl/cu128")) +app = modal.App("h3-a2a8", image=image) + +TOKENS, WIDTH = 73642, 7168 +WORLD = 8 + + +def _worker(rank: int, env: dict, out_q) -> None: + import os + import time + import torch + import torch.distributed as dist + os.environ.update(env) + os.environ.update({"MASTER_ADDR": "127.0.0.1", "MASTER_PORT": "29511"}) + torch.cuda.set_device(rank) + dist.init_process_group("nccl", rank=rank, world_size=WORLD) + rows = -(-TOKENS // WORLD) + payloads = { + "qkvg_bf16": 4 * rows * WIDTH * 2, # q, k, v, gate shards before attention + "qkv_fp4_plus_gate_bf16": 3 * rows * WIDTH * 9 // 16 + rows * WIDTH * 2, + "qkvg_fp4": 4 * rows * WIDTH * 9 // 16, # 4-bit values + one e4m3 scale per 16 + "out_bf16": rows * WIDTH * 2, # attention output back to sequence shards + } + res = {} + for name, nbytes in payloads.items(): + n = (nbytes // 2 // WORLD) * WORLD + send = torch.empty(n, dtype=torch.bfloat16, device="cuda") + recv = torch.empty_like(send) + for _ in range(3): + dist.all_to_all_single(recv, send) + torch.cuda.synchronize() + dist.barrier() + iters = 10 + t = time.perf_counter() + for _ in range(iters): + dist.all_to_all_single(recv, send) + torch.cuda.synchronize() + ms = (time.perf_counter() - t) / iters * 1e3 + res[name] = {"mb_per_rank": round(n * 2 / 1e6, 1), "ms": round(ms, 2), + "algbw_gbps": round(n * 2 / ms / 1e6, 1)} + if rank == 0: + out_q.put(res) + dist.destroy_process_group() + + +@app.function(gpu="RTX-PRO-6000:8", cpu=16, memory=65536, timeout=1800) +def bench() -> dict: + import subprocess + import torch.multiprocessing as mp + report = {"topo": subprocess.run(["nvidia-smi", "topo", "-m"], capture_output=True, text=True).stdout[-3000:]} + for label, env in (("p2p_default", {}), ("p2p_disabled", {"NCCL_P2P_DISABLE": "1"})): + ctx = mp.get_context("spawn") + q = ctx.Queue() + procs = [ctx.Process(target=_worker, args=(r, env, q)) for r in range(WORLD)] + for p in procs: + p.start() + report[label] = q.get(timeout=900) + for p in procs: + p.join() + return report + + +@app.local_entrypoint() +def main(): + import json + r = bench.remote() + print(r.pop("topo")) + print("A2A", json.dumps(r, indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py new file mode 100644 index 0000000000..543af88054 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -0,0 +1,408 @@ +"""FastH3 V2 NVFP4 on one RTX PRO 6000 (sm_120): build, convert, time generations.""" +import json +import os +import pathlib +import subprocess +import time + +import modal + +WORKTREE = pathlib.Path(__file__).resolve().parents[3] +CUTLASS_COMMIT = "e67e63c331d6e4b729047c95cf6b92c8454cba89" +volume = modal.Volume.from_name("h3-pro6000-weights", create_if_missing=True) + +image = ( + modal.Image.from_registry("nvidia/cuda:13.0.1-devel-ubuntu24.04", add_python="3.12") + .apt_install("git", "build-essential", "ffmpeg", "libgl1", "libglib2.0-0") + .pip_install("uv") + .add_local_dir(WORKTREE, "/src/fastvideo", copy=True, + ignore=[".git", "**/__pycache__", "fastvideo-kernel/include/cutlass/**", + "fastvideo-kernel/include/tk/**", "fastvideo/third_party/eval/**", "docs/**", + "assets/**", "comfyui/**", "apps/**", "**/*.mp4"]) + .run_commands("cd /src/fastvideo && UV_TORCH_BACKEND=cu130 uv pip install --system -e . --no-sources") + .run_commands("uv pip install --system 'cmake==3.31.6' ninja 'scikit-build-core>=0.10' pybind11 hf_transfer") + .run_commands(f"git clone --filter=blob:none https://github.com/NVIDIA/cutlass.git /cutlass && " + f"git -C /cutlass checkout {CUTLASS_COMMIT}") + .env({"FLASHINFER_CUDA_ARCH_LIST": "12.0a", "FLASHINFER_WORKSPACE_BASE": "/vol/cache/flashinfer", + "TORCHINDUCTOR_CACHE_DIR": "/vol/cache/inductor", "TRITON_CACHE_DIR": "/vol/cache/triton", + "FASTVIDEO_VSA_SM100A": "0", "FASTVIDEO_FA4": "0", "FASTVIDEO_STAGE_LOGGING": "1", + "HF_HUB_ENABLE_HF_TRANSFER": "1"}) +) +app = modal.App("h3-pro6000-fastest", image=image) + + +def _sh(cmd: str, **kw) -> str: + proc = subprocess.run(cmd, shell=True, capture_output=True, text=True, **kw) + out = (proc.stdout + proc.stderr)[-8000:] + if proc.returncode != 0: + raise RuntimeError(f"command failed ({proc.returncode}): {cmd}\n{out}") + return out + + +@app.function(cpu=16, memory=32768, timeout=3600, volumes={"/vol": volume}) +def build_kernel() -> str: + _sh("rm -rf /src/fastvideo/fastvideo-kernel/include/cutlass && " + "ln -s /cutlass /src/fastvideo/fastvideo-kernel/include/cutlass && mkdir -p /vol/wheels/cu130") + _sh("rm -f /vol/wheels/cu130/*.whl") + env = dict(os.environ, TORCH_CUDA_ARCH_LIST="12.0a", MAX_JOBS="16", CC="gcc", CXX="g++", CUDAHOSTCXX="g++", + CMAKE_ARGS="-DFASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER=ON -DFASTVIDEO_KERNEL_BUILD_TK=OFF " + "-DGPU_BACKEND=CUDA -DCMAKE_CUDA_ARCHITECTURES=120a") + _sh("cd /src/fastvideo/fastvideo-kernel && pip wheel . --no-build-isolation --no-deps -w /vol/wheels/cu130", + env=env) + volume.commit() + return _sh("ls -la /vol/wheels/cu130") + + +def _install_kernel(): + _sh("pip install --no-deps --force-reinstall /vol/wheels/cu130/*.whl") + + +def _build_light_int8_vae() -> str: + """26-block LynnReal light decoder: dense fp16 decoder + official encoder, Kijai int8-convrot overlay.""" + from huggingface_hub import hf_hub_download + from safetensors import safe_open + from safetensors.torch import save_file + + target = pathlib.Path("/vol/fv/vae_light_int8") + if (target / "config.json").exists(): + return "exists" + target.mkdir(parents=True, exist_ok=True) + overlay = hf_hub_download("Kijai/MiniMax-H3-experimental", "minimax_h3_lynnreal_light_vae_int8_convrot.safetensors", + local_dir="/vol/kijai") + official = pathlib.Path("/vol/official/vae") + weight_map = json.loads((official / "diffusion_pytorch_model.safetensors.index.json").read_text())["weight_map"] + tensors = {} + for shard in sorted(set(weight_map.values())): + with safe_open(str(official / shard), framework="pt") as reader: + for key in reader.keys(): + if not key.startswith("decoder.transformer_blocks."): + tensors[key] = reader.get_tensor(key) + with safe_open("/vol/light-vae/lynnreal_light_vae_decoder_fp16.safetensors", framework="pt") as reader: + light_keys = list(reader.keys()) + for key in light_keys: + tensors[key] = reader.get_tensor(key) + blocks = {int(k.split(".")[2]) for k in tensors if k.startswith("decoder.transformer_blocks.")} + assert blocks == set(range(26)), sorted(blocks) + save_file(tensors, str(target / "diffusion_pytorch_model.safetensors")) + config = json.loads((official / "config.json").read_text()) + config["decoder_num_layers"] = 26 + (target / "config.json").write_text(json.dumps(config, indent=2) + "\n") + (target / "minimax_h3_video_vae_int8_convrot.safetensors").symlink_to(overlay) + return f"light decoder keys={len(light_keys)} total={len(tensors)} blocks={len(blocks)}" + + +@app.function(gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=10800, volumes={"/vol": volume}) +def convert(minimal: bool = True) -> dict: + _install_kernel() + report = {} + os.makedirs("/vol/fv", exist_ok=True) + if not os.path.exists("/vol/fv/text_encoder_nvfp4/config.json"): + report["text_encoder"] = _sh( + "cd /src/fastvideo && python scripts/checkpoint_conversion/convert_minimax_h3_text_encoder_nvfp4.py " + "--src /vol/v2-nvfp4/text_encoder --dst /vol/fv/text_encoder_nvfp4")[-1500:] + volume.commit() + for name, src, flag in (("transformer_full", "v2-nvfp4", "--quantize-attention"), + ("transformer_ffn", "v2-nvfp4", ""), + ("transformer_vsa", "v2-nvfp4", "--quantize-attention --quantize-gate"), + ("v4_transformer_vsa", "v4-nvfp4", "--quantize-attention --quantize-gate")): + if not os.path.exists(f"/vol/{src}/transformer"): + continue + if minimal and name not in ("transformer_vsa", "v4_transformer_vsa"): + continue + if not os.path.exists(f"/vol/fv/{name}/config.json"): + report[name] = _sh( + "cd /src/fastvideo && python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py " + f"--src /vol/{src}/transformer --dst /vol/fv/{name} {flag}")[-1500:] + volume.commit() + # VAE folders: official dense shards, plus Comfy's int8-convrot overlay variant. + for vae_name in (() if minimal else ("vae_dense", "vae_int8")): + target = pathlib.Path(f"/vol/fv/{vae_name}") + target.mkdir(parents=True, exist_ok=True) + for item in pathlib.Path("/vol/official/vae").iterdir(): + link = target / item.name + if not link.exists(): + link.symlink_to(item) + int8 = pathlib.Path("/vol/fv/vae_int8/minimax_h3_video_vae_int8_convrot.safetensors") + if not minimal and not int8.exists(): + int8.symlink_to("/vol/comfy/vae/minimax_h3_video_vae_int8_convrot.safetensors") + report["light_vae"] = _build_light_int8_vae() + # Model folders: V2 small components + chosen transformer / text encoder / VAE. + # 4-step VSA-0.9 model: its own manifest/schedulers, shared text encoder, VAEs and audio VAE. + root = pathlib.Path("/vol/fv/v4_vsa_light") + if pathlib.Path("/vol/v4-nvfp4/fastvideo_inference.json").exists(): + root.mkdir(parents=True, exist_ok=True) + for item in pathlib.Path("/vol/v4-nvfp4").iterdir(): + if item.name in ("transformer", "text_encoder", "vae", "audio_vae") or item.name.startswith("."): + continue + if not (root / item.name).exists(): + (root / item.name).symlink_to(item) + for comp, src in (("transformer", "/vol/fv/v4_transformer_vsa"), ("text_encoder", "/vol/fv/text_encoder_nvfp4"), + ("vae", "/vol/fv/vae_light_int8"), ("audio_vae", "/vol/v2-nvfp4/audio_vae")): + if not (root / comp).exists(): + (root / comp).symlink_to(src) + model_sets = (("v2_vsa_light", "transformer_vsa", "vae_light_int8"), + ("v2_full_light", "transformer_full", "vae_light_int8"), + ("v2_full_int8", "transformer_full", "vae_int8"), + ("v2_full_dense", "transformer_full", "vae_dense"), + ("v2_ffn_int8", "transformer_ffn", "vae_int8")) + for model, transformer, vae in (model_sets[:1] if minimal else model_sets): + root = pathlib.Path(f"/vol/fv/{model}") + root.mkdir(parents=True, exist_ok=True) + for item in pathlib.Path("/vol/v2-nvfp4").iterdir(): + if item.name in ("transformer", "text_encoder", "vae") or item.name.startswith("."): + continue + link = root / item.name + if not link.exists(): + link.symlink_to(item) + for comp, src in (("transformer", transformer), ("text_encoder", "text_encoder_nvfp4"), ("vae", vae)): + link = root / comp + if not link.exists(): + link.symlink_to(f"/vol/fv/{src}") + volume.commit() + report["models"] = _sh("ls -la /vol/fv /vol/fv/v2_vsa_light") + # Are V2's VSA compression gates trained (nonzero)? + import torch + from safetensors import safe_open + index = json.loads(pathlib.Path("/vol/v2-nvfp4/transformer/diffusion_pytorch_model.safetensors.index.json").read_text()) + gate_stats = {} + for blk in (0, 25, 49): + key = f"transformer_blocks.{blk}.attn.to_gate_compress.weight" + with safe_open(f"/vol/v2-nvfp4/transformer/{index['weight_map'][key]}", framework="pt") as reader: + w = reader.get_tensor(key).float() + gate_stats[blk] = {"abs_mean": w.abs().mean().item(), "nonzero_frac": (w != 0).float().mean().item()} + report["gate_stats"] = gate_stats + return report + + +PROMPTS = { + "kitesurf": ("A kite surfer carves hard across choppy bay water while the camera dives alongside; spray hisses off " + "the board edge, the sail flaps and snaps in the wind, and gulls cry overhead."), + "chef": ("(S1) In a bright home kitchen, a chef looks straight at the camera and says [English] Fold the eggs " + "gently and taste before you salt. A pot simmers behind her with soft bubbling and no music."), +} + + +@app.function(gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=5400, volumes={"/vol": volume}) +def run_variant(name: str, model: str, profile: str, attention: str, decode: str, vae_compile: bool, + height: int = 480, width: int = 832, warmups: int = 1, num_frames: int = 124, + env: dict | None = None, prompts: tuple = ("kitesurf", "chef"), sparsity: float = 0.8, + steps: int = 9, num_gpus: int = 1, parallel_decode: bool = False, pre_runs: tuple = ()) -> dict: + os.environ.update(env or {}) + _install_kernel() + import torch + from fastvideo import VideoGenerator + + if attention == "VIDEO_SPARSE_ATTN_H3": + os.environ["FASTVIDEO_VSA_TRITON"] = "1" + experimental = {"attention_backend": attention, "h3_sequential_load": False, "inference_torch_compile": False, + "vae_parallel_decode": parallel_decode, "video_decode_backend": decode} + if attention == "VIDEO_SPARSE_ATTN_H3": + experimental.update({"VSA_sparsity": sparsity, "VSA_tile_size": 64}) + config = { + "model_path": f"/vol/fv/{model}", + "engine": {"num_gpus": num_gpus, "use_fsdp_inference": False, + "quantization": {"transformer_quant": "NVFP4", "layer_profile": profile}, + "parallelism": {"tp_size": 1, "sp_size": num_gpus}, + "offload": {"dit": False, "dit_layerwise": False, "text_encoder": False, "vae": False, + "pin_cpu_memory": num_gpus == 1, "lazy_module_load": False}, + "compile": {"enabled": False, "vae_enabled": vae_compile}}, + "pipeline": {"experimental": experimental}, + } + t0 = time.perf_counter() + generator = VideoGenerator.from_config(config) + load_s = time.perf_counter() - t0 + out_dir = pathlib.Path(f"/vol/outputs/{name}") + out_dir.mkdir(parents=True, exist_ok=True) + results = {"name": name, "load_s": round(load_s, 1), "env": env or {}, "shape": [height, width, num_frames], + "num_gpus": num_gpus, + "runs": []} + try: + # Optional runs at other shapes first (e.g. a 480p correctness clip), same loaded model. + for j, (ph, pw, pf, pid) in enumerate(pre_runs): + request = {"prompt": PROMPTS[pid], "negative_prompt": "", + "sampling": {"seed": 20260929, "height": ph, "width": pw, "num_frames": pf, "fps": 24, + "num_inference_steps": steps, "guidance_scale": 1.0, "batch_cfg": False}, + "output": {"output_path": str(out_dir / f"pre{j:02d}_{pid}_{ph}p.mp4"), "save_video": True, + "return_frames": False}} + t = time.perf_counter() + result = generator.generate(request) + results.setdefault("pre_runs", []).append({"shape": [ph, pw, pf], "prompt": pid, + "wall_s": round(time.perf_counter() - t, 2), + "video": getattr(result, "video_path", None)}) + order = [prompts[0]] * warmups + list(prompts) + for i, pid in enumerate(order): + prompt = PROMPTS[pid] + request = {"prompt": prompt, "negative_prompt": "", + "sampling": {"seed": 20260929, "height": height, "width": width, "num_frames": num_frames, "fps": 24, + "num_inference_steps": steps, "guidance_scale": 1.0, "batch_cfg": False}, + "output": {"output_path": str(out_dir / f"{i:02d}_{pid}.mp4"), "save_video": True, + "return_frames": False}} + torch.cuda.synchronize() + t = time.perf_counter() + result = generator.generate(request) + torch.cuda.synchronize() + wall = time.perf_counter() - t + results["runs"].append({"prompt": pid, "warmup": i < warmups, "wall_s": round(wall, 2), + "generation_time_s": getattr(result, "generation_time", None), + "video": getattr(result, "video_path", None)}) + results["peak_mem_gb_device"] = _sh("nvidia-smi --query-gpu=memory.used --format=csv,noheader").strip() + finally: + generator.shutdown() + volume.commit() + return results + + +bench_image = image.add_local_file(pathlib.Path(__file__).parent / "bench_code.py", "/root/bench_code.py") + +SHAPES = [ + # name, prefix segments (text, audio rows), video latent tokens (t, h, w) after 1x2x2 patching + ("480p_124f", [256, 414], [37, 15, 26]), + ("768p_243f", [256, 810], [72, 24, 42]), +] + + +@app.function(gpu="RTX-PRO-6000", memory=65536, cpu=8, timeout=3600, volumes={"/vol": volume}, image=bench_image) +def bench_block() -> dict: + _install_kernel() + import sys + sys.path.insert(0, "/root") + import bench_code + report = {"check": bench_code.check_tile64()} + report.update(bench_code.run(SHAPES)) + return json.loads(json.dumps(report, default=str)) + + +@app.function(gpu="RTX-PRO-6000", memory=65536, cpu=8, timeout=1800, volumes={"/vol": volume}, image=bench_image) +def density_fn() -> dict: + _install_kernel() + import sys + sys.path.insert(0, "/root") + import bench_code + return json.loads(json.dumps(bench_code.density_study(), default=str)) + + +@app.function(gpu="RTX-PRO-6000", memory=65536, cpu=8, timeout=1800, volumes={"/vol": volume}, image=bench_image) +def kcheck_fn() -> dict: + _install_kernel() + import sys + sys.path.insert(0, "/root") + import bench_code + out = {"check": bench_code.check_tile64()} + out["density"] = bench_code.density_study() + return json.loads(json.dumps(out, default=str)) + + +@app.function(gpu="RTX-PRO-6000:2", memory=229376, cpu=16, timeout=5400, volumes={"/vol": volume}) +def run_variant2(*args, **kwargs) -> dict: + return run_variant.local(*args, **kwargs) + + +@app.function(gpu="RTX-PRO-6000:8", memory=65536, cpu=16, timeout=5400, volumes={"/vol": volume}) +def run_variant8(*args, **kwargs) -> dict: + return run_variant.local(*args, **kwargs) + + +@app.function(cpu=8, memory=32768, timeout=1800, volumes={"/vol": volume}) +def compare_videos(a: str, b: str) -> dict: + """Frame PSNR between two MP4s on the volume (same seed and prompt).""" + import imageio.v3 as iio + import numpy as np + fa = iio.imread(a, plugin="pyav").astype(np.float32) + fb = iio.imread(b, plugin="pyav").astype(np.float32) + n = min(len(fa), len(fb)) + mse = ((fa[:n] - fb[:n]) ** 2).reshape(n, -1).mean(axis=1) + psnr = 10 * np.log10(255.0**2 / np.maximum(mse, 1e-9)) + return {"frames": [len(fa), len(fb)], "psnr_mean": float(psnr.mean()), "psnr_min": float(psnr.min())} + + +FAST_ENV = {"FASTVIDEO_H3_VSA_FP4": "1", "FASTVIDEO_MINIMAX_H3_FUSIONS": "all", "FASTVIDEO_NVFP4_MM_BACKEND": "cutlass", + "FASTVIDEO_H3_VAE_TILE_BATCH": "28"} + + +@app.local_entrypoint() +def main(step: str = "all"): + if step == "prep_personal": + print("KERNEL", build_kernel.remote()) + print("CONVERT", json.dumps(convert.remote(minimal=True), indent=1)[:6000]) + return + if step == "sp2check": + common = dict(attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=False, height=480, width=832, + num_frames=124, warmups=0, env=FAST_ENV, prompts=("kitesurf",), sparsity=0.8, steps=9) + one = run_variant.spawn("sp1_480p", "v2_vsa_light", "h3_dit_vsa", **common) + two = run_variant2.spawn("sp2_480p", "v2_vsa_light", "h3_dit_vsa", num_gpus=2, **common) + r1, r2 = one.get(), two.get() + print("RESULT", json.dumps(r1)) + print("RESULT", json.dumps(r2)) + print("COMPARE", json.dumps(compare_videos.remote(r1["runs"][0]["video"], r2["runs"][0]["video"]))) + return + if step == "sp8": + r = run_variant8.remote( + "sp8_v2_8step", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", + vae_compile=True, height=768, width=1344, num_frames=243, warmups=1, env=FAST_ENV, + prompts=("kitesurf", "chef", "kitesurf", "chef"), sparsity=0.8, steps=9, num_gpus=8, parallel_decode=True, + pre_runs=((480, 832, 124, "kitesurf"), )) + print("RESULT", json.dumps(r)) + if r.get("pre_runs"): + print("COMPARE", json.dumps(compare_videos.remote("/vol/outputs/sp1_480p/00_kitesurf.mp4", + r["pre_runs"][0]["video"]))) + return + if step in ("bench8_v2", "bench8_v4"): + v2 = step == "bench8_v2" + r = run_variant8.remote( + f"sp8_{'v2_8step' if v2 else 'v4_4step'}_768p10s", "v2_vsa_light" if v2 else "v4_vsa_light", "h3_dit_vsa", + attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=True, height=768, width=1344, + num_frames=243, warmups=1, env=FAST_ENV, prompts=("kitesurf", "chef", "kitesurf", "chef"), + sparsity=0.8 if v2 else 0.9, steps=9 if v2 else 5, num_gpus=8, parallel_decode=True) + print("RESULT", json.dumps(r)) + return + if step == "build_bench": + print("KERNEL", build_kernel.remote()) + step = "bench" + if step in ("prep768", "e2e768"): + if step == "prep768": + print("KERNEL", build_kernel.remote()) + print("CONVERT", json.dumps(convert.remote(), indent=1)[:6000]) + fast = {"FASTVIDEO_H3_VSA_FP4": "1", "FASTVIDEO_MINIMAX_H3_FUSIONS": "all", + "FASTVIDEO_NVFP4_MM_BACKEND": "cutlass", "FASTVIDEO_H3_VAE_TILE_BATCH": "28"} + common = dict(attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=True, height=768, width=1344, + num_frames=243, warmups=1, env=fast, prompts=("kitesurf", "chef")) + calls = { + "v4_4step_vsa90_768p10s": run_variant.spawn("v4_4step_vsa90_768p10s", "v4_vsa_light", "h3_dit_vsa", + sparsity=0.9, steps=5, **common), + "v2_8step_vsa80_768p10s": run_variant.spawn("v2_8step_vsa80_768p10s", "v2_vsa_light", "h3_dit_vsa", + sparsity=0.8, steps=9, **common), + } + for name, call in calls.items(): + try: + print("RESULT", json.dumps(call.get())) + except Exception as exc: # noqa: BLE001 + print("FAILED", name, repr(exc)[:3000]) + return + if step == "kcheck": + print("KERNEL", build_kernel.remote()) + print("KCHECK", json.dumps(kcheck_fn.remote(), indent=1)) + return + if step == "density": + print("DENSITYRESULT", json.dumps(density_fn.remote(), indent=1)) + return + if step == "bench": + print("BENCHRESULT", json.dumps(bench_block.remote(), indent=1)) + return + if step in ("all", "build"): + print("KERNEL", build_kernel.remote()) + if step in ("all", "convert"): + print("CONVERT", json.dumps(convert.remote(), indent=1)[:6000]) + if step in ("all", "run"): + variants = [ + ("full_qatinfer_light", "v2_full_light", "h3_dit", "ATTN_QAT_INFER", "h3-vae", True), + ("full_vsa_light", "v2_full_light", "h3_dit", "VIDEO_SPARSE_ATTN_H3", "h3-vae", True), + ("full_qatinfer_int8", "v2_full_int8", "h3_dit", "ATTN_QAT_INFER", "h3-vae", True), + ("full_qatinfer_taeh3", "v2_full_light", "h3_dit", "ATTN_QAT_INFER", "taeh3", False), + ("ffn_vsa_int8", "v2_ffn_int8", "h3_dit_ffn", "VIDEO_SPARSE_ATTN_H3", "h3-vae", True), + ] + # One RTX PRO 6000 per variant, all at once; failures are returned, not raised. + for v, result in zip(variants, run_variant.starmap(variants, return_exceptions=True)): + if isinstance(result, Exception): + print("FAILED", v[0], repr(result)[:3000]) + else: + print("RESULT", json.dumps(result)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_code.py b/scripts/benchmarks/minimax_h3_pro6000/bench_code.py new file mode 100644 index 0000000000..3f3cc590b9 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_code.py @@ -0,0 +1,338 @@ +"""Per-component timing of one MiniMax-H3 block on the current GPU. Imported inside the Modal container.""" +import json +import math +import time + +import torch +import torch.nn.functional as F + +HID, HEADS, HD, FFN = 5376, 56, 128, 14336 + + +def timed(fn, iters=5, warmup=2): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + return round(start.elapsed_time(end) / iters, 3) + + +def smooth_tokens(prefix_len, shape, dev, gen, noise=0.5): + """Token features with spatial-temporal correlation (proxy for real q/k), packed [prefix | video].""" + t, h, w = shape + low = torch.randn(1, HEADS * HD, max(2, t // 6), max(2, h // 6), max(2, w // 6), device=dev, generator=gen) + vid = F.interpolate(low, size=(t, h, w), mode="trilinear", align_corners=False) + vid = vid.reshape(HEADS, HD, -1).permute(2, 0, 1) + vid = vid + noise * torch.randn(vid.shape, device=dev, generator=gen) + pre = torch.randn(prefix_len, HEADS, HD, device=dev, generator=gen) + return torch.cat([pre, vid]).to(torch.bfloat16)[None] # [1, L, H, D] + + +def run(shapes): + from flashinfer import SfLayout, mm_fp4, nvfp4_quantize + from fastvideo.attention.backends.video_sparse_attn_h3 import (_build_block_mask, _h3_tile_geometry, _pool_tiles) + from fastvideo_kernel.block_sparse_attn import block_sparse_attn as bsa64 + from fastvideo.models.dits.minimax_h3_fusions import (fused_qknorm_rope, fused_residual_gate_rmsnorm_modulate, + fused_rmsnorm_modulate, minimax_h3_swiglu) + import attn_qat_infer.api as fa + from attn_qat_infer.api import BLOCK_M, sageattn_blackwell, sageattn_blackwell_sparse + + dev = torch.device("cuda") + gen = torch.Generator(device=dev).manual_seed(0) + report = {"device": torch.cuda.get_device_name(0), "torch": torch.__version__} + unit = torch.tensor(1.0, device=dev) + + def fp4_weight(n, k): + w = torch.randn(n, k, device=dev, dtype=torch.bfloat16) * 0.02 + gsf = (448 * 6) / w.float().abs().max() + wq, ws = nvfp4_quantize(w, gsf, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + return w, wq, ws, (1.0 / gsf).float() + + for name, prefix_segments, video_shape in shapes: + prefix_segments = tuple(prefix_segments) + video_shape = tuple(video_shape) + L = sum(prefix_segments) + math.prod(video_shape) + r = {"tokens": L} + torch.cuda.empty_cache() + + # ---------------- linears ---------------- + lin = {} + for lname, n, k in (("qkv_one", 7168, HID), ("out", HID, 7168), ("fc_in", 2 * FFN, HID), + ("fc_out", HID, FFN)): + w, wq, ws, alpha = fp4_weight(n, k) + x = torch.randn(L, k, device=dev, dtype=torch.bfloat16) + xq, xs = nvfp4_quantize(x, unit, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + entry = {"quant_ms": timed(lambda: nvfp4_quantize(x, unit, sfLayout=SfLayout.layout_128x4, + do_shuffle=False))} + for backend in ("cudnn", "cutlass", "auto"): + try: + entry[f"mm_{backend}_ms"] = timed( + lambda: mm_fp4(xq, wq.T, xs, ws.T, alpha, torch.bfloat16, None, backend=backend)) + except Exception as exc: # noqa: BLE001 + entry[f"mm_{backend}_ms"] = f"ERR {type(exc).__name__}: {str(exc)[:120]}" + entry["bf16_ms"] = timed(lambda: F.linear(x, w)) + entry["tflops_fp4_auto"] = (round(2 * L * n * k / entry["mm_auto_ms"] / 1e9, 1) + if isinstance(entry["mm_auto_ms"], float) else None) + lin[lname] = entry + del w, wq, ws, x, xq, xs + r["linear"] = lin + + # ---------------- attention ---------------- + geom = _h3_tile_geometry(prefix_segments, video_shape, dev, (4, 4, 4)) + _, vbs, untile, n_prefix, n_video = geom + n_tiles = vbs.numel() + Lpad = n_tiles * 64 + att = {"tiles64": n_tiles, "prefix_tiles": n_prefix} + prefix_len = sum(prefix_segments) + q = smooth_tokens(prefix_len, video_shape, dev, gen) + k = (q + 0.5 * torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16)).to(torch.bfloat16) + v = torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16) + g = torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16) * 0.1 + + def tile(x): + buf = torch.zeros((1, Lpad, HEADS, HD), device=dev, dtype=x.dtype) + buf[:, untile] = x + return buf + + qt, kt, vt, gt = tile(q), tile(k), tile(v), tile(g) + att["tile_scatter4_gather1_ms"] = timed(lambda: (tile(q), tile(k), tile(v), tile(g), qt[:, untile])) + + def vsa_full(sparsity=0.8, with_gate=True): + qp = _pool_tiles(qt, vbs, 64) + kp = _pool_tiles(kt, vbs, 64) + scores = torch.matmul(qp, kp.transpose(-2, -1)) / HD**0.5 + mask = _build_block_mask(scores, n_prefix, n_video, sparsity, True) + out, _ = bsa64(qt.transpose(1, 2).contiguous(), kt.transpose(1, 2).contiguous(), + vt.transpose(1, 2).contiguous(), mask, vbs) + out = out.transpose(1, 2).contiguous() + if with_gate: + vp = _pool_tiles(vt, vbs, 64) + oc = torch.matmul(torch.softmax(scores, dim=-1), vp).permute(0, 2, 1, 3).to(out.dtype) + out = (out.view(1, n_tiles, 64, HEADS, HD) + oc.unsqueeze(2) * gt.view(1, n_tiles, 64, HEADS, HD)).view( + 1, Lpad, HEADS, HD) + return out, mask + + _, mask = vsa_full() + att["tile_density"] = round(mask.float().mean().item(), 4) + att["vsa_triton_total_ms"] = timed(vsa_full) + qb, kb, vb = (x.transpose(1, 2).contiguous() for x in (qt, kt, vt)) + att["vsa_triton_kernel_only_ms"] = timed(lambda: bsa64(qb, kb, vb, mask, vbs)) + + # dense FP4 on the unpadded packed sequence + qd, kd, vd = (x.transpose(1, 2).contiguous() for x in (q, k, v)) + att["sage3_dense_fp4_ms"] = timed(lambda: sageattn_blackwell(qd, kd, vd), iters=3) + # same with delta_s cached + per_block_mean=False (no 9.5 GB memset) + zero_ds = {} + + def sage_nods(qx, kx, vx): + QL, KL = qx.size(2), kx.size(2) + qx, kx, vx = (F.pad(x, (0, 0, 0, (BLOCK_M - x.size(2) % BLOCK_M) % BLOCK_M)).contiguous() for x in (qx, kx, vx)) + key = (qx.shape[0], qx.shape[1], kx.shape[2]) + if key not in zero_ds: + zero_ds[key] = torch.zeros((qx.shape[0], qx.shape[1], 1, kx.shape[2]), device=dev, dtype=torch.float32) + ql, kl, vl = fa.scale_and_quant_fp4(qx), fa.scale_and_quant_fp4_permute(kx), fa.scale_and_quant_fp4_transpose(vx) + return fa.blockscaled_fp4_attn(ql, kl, vl, zero_ds[key], KL, False, False, True, True, None)[0][:, :, :QL] + + att["sage3_dense_fp4_no_deltas_ms"] = timed(lambda: sage_nods(qd, kd, vd), iters=3) + out_a = sageattn_blackwell(qd, kd, vd) + out_b = sage_nods(qd, kd, vd) + att["no_deltas_max_abs_diff"] = (out_a.float() - out_b.float()).abs().max().item() + Lq = qd.shape[2] + att["dense_fp4_tflops"] = round(4 * Lq * Lq * HD * HEADS / att["sage3_dense_fp4_no_deltas_ms"] / 1e9, 1) + if L < 30000: + att["sdpa_bf16_dense_ms"] = timed(lambda: F.scaled_dot_product_attention(qd, kd, vd), iters=3) + + # sparse FP4 on 128x128 blocks with 64x64 quadrant masks (exact VSA tile-64 semantics) + nt2 = n_tiles + (n_tiles % 2) + qs, ks, vs = (F.pad(x.transpose(1, 2), (0, 0, 0, nt2 * 64 - Lpad)).contiguous() for x in (qt, kt, vt)) + q2k_idx, q2k_num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, 64, vbs, validate=True) + att["block128_density"] = round(q2k_num.float().mean().item() / (nt2 // 2), 4) + att["sparse_fp4_quad_ms"] = timed(lambda: sageattn_blackwell_sparse(qs, ks, vs, q2k_idx, q2k_num, kvv, quad)) + att["mask_to_blocks_ms"] = timed(lambda: fa.vsa_tile_mask_to_fp4_blocks(mask, 64, vbs)) + o_fp4 = sageattn_blackwell_sparse(qs, ks, vs, q2k_idx, q2k_num, kvv, quad)[:, :, :Lpad] + o_tri, _ = bsa64(qb, kb, vb, mask, vbs) + rows = untile # valid (non-pad) rows of the tile buffer + diff = (o_fp4[:, :, rows].float() - o_tri[:, :, rows].float()) + att["quad_vs_triton_rel_l2"] = round((diff.norm() / o_tri[:, :, rows].float().norm()).item(), 4) + od = sageattn_blackwell(qd, kd, vd) + from torch.nn.attention import SDPBackend, sdpa_kernel + if L < 30000: + ref_dense = F.scaled_dot_product_attention(qd, kd, vd) + att["dense_fp4_vs_bf16_rel_l2"] = round(((od.float() - ref_dense.float()).norm() / + ref_dense.float().norm()).item(), 4) + att["quad_finite"] = bool(torch.isfinite(o_fp4[:, :, rows]).all()) + r["attention"] = att + del q, k, v, g, qt, kt, vt, gt, qd, kd, vd, qs, ks, vs, qb, kb, vb, out_a, out_b, zero_ds, o_fp4, o_tri, od + + # ---------------- elementwise ---------------- + ew = {} + x = torch.randn(1, L, HID, device=dev, dtype=torch.bfloat16) + br = torch.randn_like(x) + nw = torch.randn(HID, device=dev, dtype=torch.bfloat16) + tab = [torch.randn(3, HID, device=dev, dtype=torch.bfloat16) * 0.1 for _ in range(6)] + idx = torch.randint(0, 3, (L,), device=dev) + norm = torch.nn.RMSNorm(HID, eps=1e-6, device=dev, dtype=torch.bfloat16) + + def eager_mod(): + n1 = norm(x) * (1.0 + tab[1].index_select(0, idx)) + tab[0].index_select(0, idx) + h = x + tab[2].index_select(0, idx) * br + n2 = norm(h) * (1.0 + tab[4].index_select(0, idx)) + tab[3].index_select(0, idx) + out = h + tab[5].index_select(0, idx) * br + return n1, n2, out + + def fused_mod(): + n1 = fused_rmsnorm_modulate(x, nw, tab[1], tab[0], idx, 1e-6) + h, n2 = fused_residual_gate_rmsnorm_modulate(x, br, tab[2], nw, tab[4], tab[3], idx, 1e-6) + out = h + tab[5].index_select(0, idx) * br + return n1, n2, out + + with torch.no_grad(): + ew["modulate_eager_ms"] = timed(eager_mod) + ew["modulate_fused_ms"] = timed(fused_mod) + ew["modulate_compiled_ms"] = timed(torch.compile(eager_mod)) + packed = torch.randn(1, L, 2 * FFN, device=dev, dtype=torch.bfloat16) + + def eager_swiglu(): + a, b = packed.chunk(2, dim=-1) + return a * F.silu(b) + + ew["swiglu_eager_ms"] = timed(eager_swiglu) + ew["swiglu_fused_ms"] = timed(lambda: minimax_h3_swiglu(packed)) + del packed + qq = torch.randn(1, L, HEADS, HD, device=dev, dtype=torch.bfloat16) + cos = torch.randn(L, 96, device=dev, dtype=torch.bfloat16) + sin = torch.randn(L, 96, device=dev, dtype=torch.bfloat16) + qn = torch.nn.RMSNorm(HD, eps=1e-6, device=dev, dtype=torch.bfloat16) + + def eager_rope(): + outs = [] + for t in (qq, qq): + t = qn(t) + rot, pas = t[..., :96], t[..., 96:] + c, s = cos[None, :, None, :], sin[None, :, None, :] + a, b = rot.chunk(2, dim=-1) + outs.append(torch.cat((rot * c + torch.cat((-b, a), -1) * s, pas), -1).contiguous()) + return outs + + ew["qknorm_rope_eager_ms"] = timed(eager_rope) + ew["qknorm_rope_fused_ms"] = timed( + lambda: (fused_qknorm_rope(qq, qn.weight, cos, sin, 1e-6), fused_qknorm_rope(qq, qn.weight, cos, sin, 1e-6))) + r["elementwise"] = ew + report[name] = r + print("BENCH", name, json.dumps(r), flush=True) + return report + + +def check_tile64(): + """Exact semantics check of the quadrant path on small H3-style layouts.""" + from fastvideo.attention.backends.video_sparse_attn_h3 import _build_block_mask, _h3_tile_geometry, _pool_tiles + import attn_qat_infer.api as fa + from attn_qat_infer.api import sageattn_blackwell, sageattn_blackwell_sparse + + dev = torch.device("cuda") + gen = torch.Generator(device=dev).manual_seed(1) + out = {} + cases = [("odd_tiles", (77, 46), (9, 6, 10), 64), ("even_tiles", (64, 40), (8, 8, 8), 64), + ("tile256", (100, 70), (8, 8, 16), 256)] + for name, prefix, vshape, tt in cases: + H = 4 + shape = {64: (4, 4, 4), 256: (4, 8, 8)}[tt] + _, vbs, untile, n_prefix, n_video = _h3_tile_geometry(prefix, vshape, dev, shape) + n_tiles = vbs.numel() + Lpad = n_tiles * tt + L = sum(prefix) + math.prod(vshape) + q = torch.randn(1, L, H, HD, device=dev, generator=gen, dtype=torch.bfloat16) + k = torch.randn(1, L, H, HD, device=dev, generator=gen, dtype=torch.bfloat16) + v = torch.randn(1, L, H, HD, device=dev, generator=gen, dtype=torch.bfloat16) + + def tile(x): + buf = torch.zeros((1, Lpad, H, HD), device=dev, dtype=x.dtype) + buf[:, untile] = x + return buf.transpose(1, 2).contiguous() # BHSD + + qt, kt, vt = tile(q), tile(k), tile(v) + scores = torch.matmul(_pool_tiles(qt.transpose(1, 2), vbs, tt), _pool_tiles(kt.transpose(1, 2), vbs, tt).transpose(-2, -1)) + mask = _build_block_mask(scores, n_prefix, n_video, 0.8, True) + q2k_idx, q2k_num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, tt, vbs, validate=True) + Lk = q2k_idx.shape[2] * 128 + pad = lambda x: F.pad(x, (0, 0, 0, Lk - Lpad)).contiguous() + o = sageattn_blackwell_sparse(pad(qt), pad(kt), pad(vt), q2k_idx, q2k_num, kvv, quad)[:, :, :Lpad] + tok_tile = torch.arange(n_tiles, device=dev).repeat_interleave(tt) + tok_valid = torch.zeros(Lpad, dtype=torch.bool, device=dev) + tok_valid[untile] = True + tm = mask[:, :, tok_tile][:, :, :, tok_tile] & tok_valid[None, None, None, :] + ref = F.scaled_dot_product_attention(qt.float(), kt.float(), vt.float(), attn_mask=tm) + rows = untile + err = (o.float() - ref)[:, :, rows] + # FP4 noise floor: dense FP4 vs dense fp32 on the same valid tokens + qv, kv_, vv = (x[:, :, rows].contiguous() for x in (qt, kt, vt)) + od = sageattn_blackwell(qv, kv_, vv) + rd = F.scaled_dot_product_attention(qv.float(), kv_.float(), vv.float()) + out[name] = { + "tiles": n_tiles, "density": round(mask.float().mean().item(), 3), + "rel_l2": round((err.norm() / ref[:, :, rows].norm()).item(), 4), + "dense_fp4_rel_l2_floor": round(((od.float() - rd).norm() / rd.norm()).item(), 4), + "cosine": round(F.cosine_similarity(o[:, :, rows].float().flatten(), ref[:, :, rows].flatten(), dim=0).item(), 5), + "finite": bool(torch.isfinite(o[:, :, rows]).all()), + } + if tt == 64: + from fastvideo_kernel.block_sparse_attn import block_sparse_attn as bsa64 + ot, _ = bsa64(qt, kt, vt, mask, vbs) + out[name]["triton_bf16_rel_l2"] = round(((ot.float() - ref)[:, :, rows].norm() / ref[:, :, rows].norm()).item(), 4) + print("CHECK", name, json.dumps(out[name]), flush=True) + return out + + +def density_study(prefix_segments=(256, 810), video_shape=(72, 24, 42)): + """Block densities each kernel granularity would compute, and FP4 sparse timing at 0.8 / 0.9.""" + from fastvideo.attention.backends.video_sparse_attn_h3 import _build_block_mask, _h3_tile_geometry, _pool_tiles + import attn_qat_infer.api as fa + + dev = torch.device("cuda") + gen = torch.Generator(device=dev).manual_seed(0) + _, vbs, untile, n_prefix, n_video = _h3_tile_geometry(tuple(prefix_segments), tuple(video_shape), dev, (4, 4, 4)) + n_tiles = vbs.numel() + Lpad = n_tiles * 64 + q = smooth_tokens(sum(prefix_segments), tuple(video_shape), dev, gen) + k = (q + 0.5 * torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16)).to(torch.bfloat16) + v = torch.randn(q.shape, device=dev, generator=gen, dtype=torch.bfloat16) + + def tile(x): + buf = torch.zeros((1, Lpad, HEADS, HD), device=dev, dtype=x.dtype) + buf[:, untile] = x + return buf + + qt, kt, vt = tile(q), tile(k), tile(v) + scores = torch.matmul(_pool_tiles(qt, vbs, 64), _pool_tiles(kt, vbs, 64).transpose(-2, -1)) + nt2 = n_tiles + n_tiles % 2 + rows = nt2 * 64 + qs, ks, vs = (F.pad(x, (0, 0, 0, 0, 0, rows - Lpad)).contiguous() for x in (qt, kt, vt)) + out = {} + for sparsity in (0.8, 0.9): + mask = _build_block_mask(scores, n_prefix, n_video, sparsity, True) + m = F.pad(mask, (0, nt2 - n_tiles, 0, nt2 - n_tiles), value=False) + B, H = m.shape[:2] + d = {"tile64x64": m.float().mean().item()} + q2 = m.view(B, H, nt2 // 2, 2, nt2) # query pairs + d["q128_k64"] = q2.any(3).float().mean().item() + k2 = m.view(B, H, nt2, nt2 // 2, 2) + d["q64_k128"] = k2.any(4).float().mean().item() + d["q128_k128"] = m.view(B, H, nt2 // 2, 2, nt2 // 2, 2).any(5).any(3).float().mean().item() + idx, num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, 64, vbs) + def bshd(): + qh, kh, vh = (x.transpose(1, 2) for x in (qs, ks, vs)) + ds = fa._zero_delta_s(1, HEADS, rows, dev) + return fa.blockscaled_fp4_attn_sparse(fa.scale_and_quant_fp4(qh), fa.scale_and_quant_fp4_permute(kh), + fa.scale_and_quant_fp4_transpose(vh), ds, rows, idx, num, kvv, quad, + False, True, True, None)[0] + + d["sparse_fp4_quad_ms"] = timed(bshd, iters=3) + out[str(sparsity)] = {key: round(val, 4) for key, val in d.items()} + print("DENSITY", sparsity, json.dumps(out[str(sparsity)]), flush=True) + out["dense_fp4_bshd_ms"] = timed(lambda: fa.sageattn_blackwell(*(x.transpose(1, 2) for x in (qs, ks, vs))), iters=2) + return out diff --git a/scripts/benchmarks/minimax_h3_pro6000/download_personal.py b/scripts/benchmarks/minimax_h3_pro6000/download_personal.py new file mode 100644 index 0000000000..eba4edbc55 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/download_personal.py @@ -0,0 +1,40 @@ +"""Stage every FastH3 input into the personal workspace's volume (CPU only, public repos, no token).""" +import modal + +volume = modal.Volume.from_name("h3-pro6000-weights", create_if_missing=True) +image = modal.Image.debian_slim(python_version="3.12").pip_install("huggingface_hub[hf_transfer]>=0.34") +app = modal.App("h3-pro6000-download-personal", image=image) + + +@app.function(volumes={"/vol": volume}, timeout=10800, cpu=16, memory=32768) +def download() -> dict: + import os + from huggingface_hub import hf_hub_download, snapshot_download + out = {} + out["v2"] = snapshot_download("FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4", local_dir="/vol/v2-nvfp4", max_workers=16) + volume.commit() + out["v4"] = snapshot_download("FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree-NVFP4", local_dir="/vol/v4-nvfp4", + allow_patterns=["transformer/*", "scheduler/*", "audio_scheduler/*", "*.json", + "tokenizer/*", "processor/*"], max_workers=16) + volume.commit() + out["official_vae"] = snapshot_download("MiniMaxAI/MiniMax-H3", local_dir="/vol/official", + allow_patterns=["vae/*", "model_index.json", "modular_model_index.json"], + max_workers=8) + out["comfy_int8_vae"] = hf_hub_download("Comfy-Org/MiniMax-H3", "vae/minimax_h3_video_vae_int8_convrot.safetensors", + local_dir="/vol/comfy") + out["light_vae"] = hf_hub_download("corechan/MiniMax-H3-LightVAE", "lynnreal_light_vae_decoder_fp16.safetensors", + local_dir="/vol/light-vae") + out["kijai"] = hf_hub_download("Kijai/MiniMax-H3-experimental", "minimax_h3_lynnreal_light_vae_int8_convrot.safetensors", + local_dir="/vol/kijai") + volume.commit() + sizes = {} + for root in ("/vol/v2-nvfp4", "/vol/v4-nvfp4", "/vol/official", "/vol/comfy", "/vol/light-vae", "/vol/kijai"): + sizes[root] = round(sum(os.path.getsize(os.path.join(d, f)) for d, _, fs in os.walk(root) for f in fs) / 1e9, 2) + out["sizes_gb"] = sizes + return out + + +@app.local_entrypoint() +def main(): + import json + print("DOWNLOADED", json.dumps(download.remote(), indent=1)) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py new file mode 100644 index 0000000000..a8de3a945c --- /dev/null +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -0,0 +1,217 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Convert a ModelOpt-unified NVFP4 MiniMax-H3 transformer into FastVideo's packed export. + +ModelOpt's unified Hugging Face layout (e.g. ``FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4``) +stores each quantized linear as:: + + .weight uint8 [out, in // 2] two E2M1 values per byte + .weight_scale float8_e4m3fn [out, in // 16] linear layout + .weight_scale_2 float32 [] amax(|W|) / (6 * 448) + .input_scale float32 [] static activation scale + +FastVideo's packed NVFP4 DiT export (``nvfp4_weights.safetensors``, read by +``load_minimax_h3_nvfp4_dit_export``) stores ``::_nvfp4_weight`` (same +bytes), ``::_nvfp4_weight_scale`` (the same E4M3 bytes in FlashInfer's 128x4 +swizzled layout), ``::_nvfp4_alpha`` (= ``weight_scale_2``) and +``::_weight_global_sf`` (= 1 / ``weight_scale_2``). The calibrated weights are +therefore carried over bit for bit; only the activation scale changes, because +FastVideo quantizes activations per call with a unit global scale and the +static ``input_scale`` is dropped. + +``--quantize-attention`` additionally quantizes the dense BF16 attention +projections (``attn.to_{q,k,v,out}``) of every main block exactly as +``convert_model_to_nvfp4`` would at runtime, producing the full +``layer_profile="h3_dit"`` set. ``--quantize-gate`` (with it) also quantizes +each block's VSA compression gate, for ``layer_profile="h3_dit_vsa"``. +Without either the export holds the FFN linears only and must be loaded with +``layer_profile="h3_dit_ffn"``. + +Every exported linear is probed on random BF16 rows through the same +``mm_fp4`` path the loader runs; the relative error against a BF16 matmul with +the dequantized weight must stay under ``--max-probe-error`` (genuine W4A4 +noise on random inputs is about 0.1; a wrong nibble order or scale layout +reads near 1.0). Needs a Blackwell GPU with FlashInfer. + +Usage:: + + python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \\ + --src /path/to/FastH3-V2-NVFP4/transformer --dst /path/to/out/transformer [--quantize-attention] +""" +from __future__ import annotations + +import argparse +import json +import re +import shutil +from pathlib import Path + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +EXPORT_FILENAME = "nvfp4_weights.safetensors" +_BLOCK_ATTN = re.compile(r"^transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out\.0)$") +_BLOCK_GATE = re.compile(r"^transformer_blocks\.\d+\.attn\.to_gate_compress$") +_RENAMES = ((re.compile(r"\.ff\.net\.0\.proj$"), ".ff.fc_in"), (re.compile(r"\.ff\.net\.2$"), ".ff.fc_out"), + (re.compile(r"\.attn\.to_out\.0$"), ".attn.to_out")) +_E2M1 = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0]) + + +def fastvideo_module_name(diffusers_prefix: str) -> str: + name = diffusers_prefix + for pattern, replacement in _RENAMES: + name = pattern.sub(replacement, name) + return name + + +def dequantize_modelopt(packed: torch.Tensor, scale: torch.Tensor, scale_2: torch.Tensor) -> torch.Tensor: + """E2M1 low nibble = even column, high nibble = odd column (ModelOpt / CUTLASS order).""" + lut = _E2M1.to(packed.device) + low = lut[(packed & 0x0F).long()] + high = lut[(packed >> 4).long()] + values = torch.stack((low, high), dim=-1).reshape(packed.shape[0], packed.shape[1] * 2) + block = scale.to(torch.float32).repeat_interleave(16, dim=1) + return values * block * scale_2.to(torch.float32) + + +def _flashinfer(): + from flashinfer import SfLayout, mm_fp4, nvfp4_quantize + try: + from flashinfer import block_scale_interleave + except ImportError: + from flashinfer.fp4_quantization import block_scale_interleave + return SfLayout, mm_fp4, nvfp4_quantize, block_scale_interleave + + +def probe(buffers: dict[str, torch.Tensor], reference: torch.Tensor, rows: int = 256) -> float: + """Relative error of the loader's mm_fp4 path against BF16 x @ W_ref^T.""" + SfLayout, mm_fp4, nvfp4_quantize, _ = _flashinfer() + device = buffers["_nvfp4_weight"].device + x = torch.randn(rows, reference.shape[1], device=device, dtype=torch.bfloat16) + unit = torch.tensor(1.0, device=device, dtype=torch.float32) + x_fp4, x_scale = nvfp4_quantize(x, unit, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + out = mm_fp4(x_fp4, buffers["_nvfp4_weight"].T, x_scale, buffers["_nvfp4_weight_scale"].T, + buffers["_nvfp4_alpha"] / unit, torch.bfloat16, None, backend="auto") + ref = x.float() @ reference.float().T + return ((out.float() - ref).norm() / ref.norm()).item() + + +def convert_modelopt_linear(weight, scale, scale_2, device) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: + """Carry the calibrated bytes over; return (buffers, dequantized weight, scale-byte agreement). + + The agreement compares the swizzled ModelOpt scales with the scales + ``nvfp4_quantize`` derives from the dequantized weight. It is a layout check + (identical swizzle and shape) and is near 1.0 when ModelOpt used max + calibration; blocks whose largest code is below 6 legitimately differ. + """ + SfLayout, _, nvfp4_quantize, block_scale_interleave = _flashinfer() + weight = weight.to(device) + scale = scale.to(device) + scale_2 = scale_2.to(device=device, dtype=torch.float32).reshape(()) + reference = dequantize_modelopt(weight, scale, scale_2) + _, layout_ref = nvfp4_quantize(reference.to(torch.bfloat16), 1.0 / scale_2, sfLayout=SfLayout.layout_128x4, + do_shuffle=False) + swizzled = block_scale_interleave(scale.view(torch.uint8).contiguous()).reshape(layout_ref.shape) + agreement = (swizzled.view(torch.uint8) == layout_ref.view(torch.uint8)).float().mean().item() + buffers = { + "_nvfp4_weight": weight.contiguous(), + "_nvfp4_weight_scale": swizzled.view(layout_ref.dtype).contiguous(), + "_nvfp4_alpha": scale_2.clone(), + "_weight_global_sf": (1.0 / scale_2).to(torch.bfloat16), + } + return buffers, reference, agreement + + +def quantize_dense_linear(weight: torch.Tensor, device) -> tuple[dict[str, torch.Tensor], torch.Tensor]: + """Same math as ``nvfp4_config.convert_model_to_nvfp4``.""" + SfLayout, _, nvfp4_quantize, _ = _flashinfer() + weight = weight.to(device=device, dtype=torch.bfloat16) + global_sf = (448 * 6) / weight.float().abs().nan_to_num().max() + fp4_w, fp4_s = nvfp4_quantize(weight, global_sf, sfLayout=SfLayout.layout_128x4, do_shuffle=False) + global_sf = torch.as_tensor(global_sf, device=device, dtype=torch.float32) + buffers = {"_nvfp4_weight": fp4_w, "_nvfp4_weight_scale": fp4_s, "_nvfp4_alpha": (1.0 / global_sf).float(), + "_weight_global_sf": global_sf.to(torch.bfloat16)} + return buffers, weight + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--src", required=True, type=Path) + parser.add_argument("--dst", required=True, type=Path) + parser.add_argument("--quantize-attention", action="store_true") + parser.add_argument("--quantize-gate", action="store_true", help="also quantize attn.to_gate_compress") + parser.add_argument("--max-probe-error", type=float, default=0.3) + parser.add_argument("--dense-shard-gb", type=float, default=5.0) + args = parser.parse_args() + device = torch.device("cuda") + args.dst.mkdir(parents=True, exist_ok=True) + + index = json.loads((args.src / "diffusion_pytorch_model.safetensors.index.json").read_text()) + weight_map: dict[str, str] = index["weight_map"] + modelopt = sorted(k[:-len(".weight_scale_2")] for k in weight_map if k.endswith(".weight_scale_2")) + modelopt_keys = {f"{p}.{s}" for p in modelopt for s in ("weight", "weight_scale", "weight_scale_2", "input_scale")} + if args.quantize_gate and not args.quantize_attention: + parser.error("--quantize-gate requires --quantize-attention") + dense_pattern = re.compile(_BLOCK_ATTN.pattern + ("|" + _BLOCK_GATE.pattern if args.quantize_gate else "")) + attention = sorted(k[:-len(".weight")] for k in weight_map + if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")])) if args.quantize_attention else [] + attention_keys = {f"{p}.weight" for p in attention} + + readers = {shard: safe_open(str(args.src / shard), framework="pt", device="cpu") for shard in set(weight_map.values())} + tensor = lambda key: readers[weight_map[key]].get_tensor(key) + + export: dict[str, torch.Tensor] = {} + worst = 0.0 + agreements: list[float] = [] + attention_set = set(attention) + for prefix in modelopt + attention: + if prefix in attention_set: + buffers, reference = quantize_dense_linear(tensor(f"{prefix}.weight"), device) + else: + buffers, reference, agreement = convert_modelopt_linear(tensor(f"{prefix}.weight"), + tensor(f"{prefix}.weight_scale"), + tensor(f"{prefix}.weight_scale_2"), device) + agreements.append(agreement) + error = probe(buffers, reference) + worst = max(worst, error) + if error > args.max_probe_error: + raise SystemExit(f"probe error {error:.3f} on {prefix} exceeds {args.max_probe_error}; nothing written") + module = fastvideo_module_name(prefix) + for name, value in buffers.items(): + export[f"{module}::{name}"] = value.cpu() + save_file(export, str(args.dst / EXPORT_FILENAME)) + + dense_keys = [k for k in weight_map if k not in modelopt_keys and k not in attention_keys] + shard, shard_bytes, shards = {}, 0, [] + for key in dense_keys: + value = tensor(key) + shard[key] = value + shard_bytes += value.numel() * value.element_size() + if shard_bytes >= args.dense_shard_gb * 1e9: + shards.append(shard) + shard, shard_bytes = {}, 0 + if shard: + shards.append(shard) + new_map = {} + for i, part in enumerate(shards, 1): + name = f"diffusion_pytorch_model-{i:05d}-of-{len(shards):05d}.safetensors" + save_file(part, str(args.dst / name)) + new_map.update({k: name for k in part}) + (args.dst / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps({"metadata": {}, "weight_map": new_map}, indent=2) + "\n") + config = json.loads((args.src / "config.json").read_text()) + config.pop("quantization_config", None) + (args.dst / "config.json").write_text(json.dumps(config, indent=2) + "\n") + for extra in args.src.iterdir(): + if extra.suffix not in (".safetensors", ".json"): + shutil.copy2(extra, args.dst / extra.name) + print(json.dumps({"exported_linears": len(modelopt) + len(attention), "modelopt_linears": len(modelopt), + "quantized_attention_linears": len(attention), "worst_probe_error": round(worst, 4), + "gate_linears": sum(1 for p in attention if _BLOCK_GATE.match(p)), + "min_scale_byte_agreement": round(min(agreements), 4) if agreements else None, + "mean_scale_byte_agreement": round(sum(agreements) / len(agreements), 4) if agreements else None, + "dense_tensors": len(dense_keys), "dense_shards": len(shards)})) + + +if __name__ == "__main__": + main() From 0a4737cd5733826258919a41daee24fabe6452fe Mon Sep 17 00:00:00 2001 From: aryan5v Date: Fri, 2 Oct 2026 22:37:30 -0700 Subject: [PATCH 008/107] [wip]: H3 single-GPU 5090/4090 shipping path - Pre-quantized FP8 W8A8 checkpoint loader (float8 weight + per-channel weight_scale) - NVFP4 export: optional calibrated activation scale (_nvfp4_input_global_sf); converter gains --quantize-ffn (from bf16) and --act-amax - NVFP4 static/dynamic activation scales via env; FP8 attention projections next to NVFP4 FFN - AdaLN modulation host cache and precomputed tables (skips 24 GiB of AdaLN weights) - Layerwise offload streams large buffers (fix: onload by name, placeholders fail the size test) - CPU-first DiT load for layerwise/AdaLN-cache paths; pinned encoder/VAE swaps; parked modules - NVFP4 text encoder bf16 de-quant fallback for pre-Blackwell GPUs - VSA guard ignores offloaded placeholders; per-stage memory logging; memory cap / report knobs - SP stage profiling and FP8 all-to-all simulation; Modal PRO 6000 bench steps (cherry picked from commit cf04434ac7f960cde143d9ce5c4362647158543f) --- fastvideo/hooks/layerwise_offload.py | 35 +++- fastvideo/layers/quantization/fp8_config.py | 3 + fastvideo/layers/quantization/nvfp4_config.py | 71 ++++++- fastvideo/models/dits/minimax_h3.py | 107 +++++++++- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 188 +++++++++++++----- .../encoders/minimax_h3_checkpoint_nvfp4.py | 38 ++++ fastvideo/models/loader/component_loader.py | 23 ++- fastvideo/models/loader/fsdp_load.py | 73 ++++++- .../basic/minimax_h3/minimax_h3_pipeline.py | 46 ++++- .../pipelines/basic/minimax_h3/vsa_guard.py | 3 +- fastvideo/pipelines/stages/base.py | 6 + fastvideo/worker/gpu_worker.py | 37 ++++ scripts/benchmarks/minimax_h3_pro6000/app.py | 148 +++++++++++++- .../convert_minimax_h3_modelopt_nvfp4_dit.py | 33 ++- 14 files changed, 726 insertions(+), 85 deletions(-) diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 05a9aacf54..f0ce1d52eb 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -1,4 +1,5 @@ from contextlib import contextmanager +from itertools import chain from typing import Any import torch from torch import nn @@ -14,6 +15,30 @@ def _tensor_placeholder(tensor: torch.Tensor, device: torch.device) -> torch.Ten return torch.empty(shape, device=device, dtype=tensor.dtype) +# Buffers at least this large also stream (e.g. packed NVFP4 weights registered as buffers); +# small ones (scales, caches) stay resident. Opt in with FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS=1. +_BUFFER_OFFLOAD_MIN_BYTES = 1 << 20 + + +def _offload_tensors(module: nn.Module, names: dict[str, torch.Tensor] | None = None): + """``(name, tensor)`` for every parameter and, when enabled, every large buffer. + + ``names`` restricts the walk to the tensors chosen at init: an offloaded buffer is a + zero-element placeholder afterwards and would fail the size test. + """ + import os + if names is not None: + for name, tensor in chain(module.named_parameters(), module.named_buffers()): + if name in names: + yield name, tensor + return + yield from module.named_parameters() + if os.environ.get("FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS") == "1": + for name, buf in module.named_buffers(): + if buf is not None and buf.numel() * buf.element_size() >= _BUFFER_OFFLOAD_MIN_BYTES: + yield name, buf + + class LayerwiseOffloadState: def __init__( @@ -35,7 +60,7 @@ def _will_offload(self, name: str) -> bool: @torch.compiler.disable def on_init(self, module: nn.Module): self.module_ref = module - for name, param in self.module_ref.named_parameters(): + 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()) param.data = _tensor_placeholder(param.data, self.device) @@ -44,7 +69,7 @@ def on_init(self, module: nn.Module): def wait_and_replace_params(self): torch.cuda.current_stream().wait_stream(self.async_copy_stream) # now gpu_named_parameters are ready - for name, param in self.module_ref.named_parameters(): + for name, param in _offload_tensors(self.module_ref, self.cpu_named_parameters): if not self._will_offload(name): continue if name not in self.gpu_named_parameters: @@ -56,7 +81,7 @@ def wait_and_replace_params(self): def prefetch_params(self): compute_stream = torch.cuda.current_stream() with torch.cuda.stream(self.async_copy_stream): - for name, param in self.module_ref.named_parameters(): + for name, param in _offload_tensors(self.module_ref, self.cpu_named_parameters): if not self._will_offload(name): continue assert name not in self.gpu_named_parameters @@ -66,7 +91,7 @@ def prefetch_params(self): @torch.compiler.disable def release_gpu_params(self): - for name, param in self.module_ref.named_parameters(): + for name, param in _offload_tensors(self.module_ref, self.cpu_named_parameters): if self._will_offload(name): param.data = _tensor_placeholder(param.data, self.device) del self.gpu_named_parameters[name] @@ -83,7 +108,7 @@ def on_attach(self, module: nn.Module): self.state.on_init(module) # pyright: ignore def on_detach(self, module: nn.Module): - named_parameters = dict(module.named_parameters()) + 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: diff --git a/fastvideo/layers/quantization/fp8_config.py b/fastvideo/layers/quantization/fp8_config.py index e73c372d65..cabc7cfda2 100644 --- a/fastvideo/layers/quantization/fp8_config.py +++ b/fastvideo/layers/quantization/fp8_config.py @@ -32,6 +32,9 @@ _FP8_SUFFIXES = ( "ffn.fc_in", "ffn.fc_out", + # MiniMax-H3 blocks name their MLP ``ff``. + "ff.fc_in", + "ff.fc_out", "to_q", "to_k", "to_v", diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 3c30f3b921..403af3b5ac 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -80,6 +80,7 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: _MINIMAX_H3_NVFP4_FF_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.ff\.(?:fc_in|fc_out)$") _MINIMAX_H3_NVFP4_DIT_PREFIX = re.compile( r"(?:^|\.)transformer_blocks\.\d+\.(?:attn\.to_(?:q|k|v|out)|ff\.(?:fc_in|fc_out))$") +_H3_BLOCK_ATTN_PROJ = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out)$") _MINIMAX_H3_NVFP4_VSA_GATE_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.attn\.to_gate_compress$") H3_NVFP4_DIT_EXPORT_FILENAME = "nvfp4_weights.safetensors" H3_NVFP4_DIT_KEY_SEP = "::" @@ -89,6 +90,9 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: "_nvfp4_alpha", "_weight_global_sf", ) +# Optional per-layer static activation global scale (448 * 6 / calibrated input amax). +# Exports without it quantize activations with the unit global scale. +H3_NVFP4_DIT_INPUT_SF_NAME = "_nvfp4_input_global_sf" def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: @@ -362,6 +366,18 @@ def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: return x +_AMAX_TABLES: dict[str, dict[str, float]] = {} + + +def _load_amax_table(path: str) -> dict[str, float]: + if path not in _AMAX_TABLES: + import json + with open(path) as f: + raw = json.load(f) + _AMAX_TABLES[path] = {k: float(v["all"] if isinstance(v, dict) else v) for k, v in raw.items()} + return _AMAX_TABLES[path] + + class NVFP4QuantizeMethod(QuantizeMethodBase): def __init__(self, layer_prefix: str = ""): @@ -389,6 +405,34 @@ def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, layer.register_parameter("weight", weight) set_weight_attrs(weight, extra_weight_attrs) + def _static_activation_global_sf(self) -> torch.Tensor | None: + """FASTVIDEO_NVFP4_ACT_AMAX: JSON of calibrated input amax per layer ("b." or full prefix).""" + if getattr(self, "_static_sf_checked", False): + return self._static_sf + self._static_sf_checked, self._static_sf = True, None + path = os.environ.get("FASTVIDEO_NVFP4_ACT_AMAX") + if path: + table = _load_amax_table(path) + prefix = self.layer_prefix or "" + match = re.search(r"transformer_blocks\.(\d+)\.(.+)$", prefix) + keys = [prefix] + ([f"b{match.group(1)}.{match.group(2)}"] if match else []) + amax = next((table[k] for k in keys if k in table), None) + if amax is not None: + self._static_sf = torch.tensor((448.0 * 6.0) / max(amax, 1e-12), dtype=torch.float32, + device="cuda") + return self._static_sf + + def _dynamic_activation_scale(self) -> bool: + """FASTVIDEO_NVFP4_DYNAMIC_ACT: "all", or comma-separated layer-name suffixes (e.g. "ff.fc_out").""" + cached = getattr(self, "_dynamic_act_cached", None) + if cached is None: + selected = os.environ.get("FASTVIDEO_NVFP4_DYNAMIC_ACT", "") + suffixes = [part.strip() for part in selected.split(",") if part.strip()] + prefix = self.layer_prefix or "" + cached = "all" in suffixes or any(prefix.endswith(suffix) for suffix in suffixes) + self._dynamic_act_cached = cached + return cached + def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: SfLayout, _, _ = _require_flashinfer() x = _coerce_fp4_input_dtype(x) @@ -448,7 +492,17 @@ def apply( else: x = _coerce_fp4_input_dtype(x) x = x.view(-1, x.shape[-1]) - x_global_sf = self.x_global_sf + static_sf = self._static_activation_global_sf() + if static_sf is None: + static_sf = getattr(layer, H3_NVFP4_DIT_INPUT_SF_NAME, None) + if static_sf is not None: + x_global_sf = static_sf + elif self._dynamic_activation_scale(): + # A unit global scale caps FP8 block scales at |x| = 6 * 448; inputs such as H3's ff.fc_out + # (post-SwiGLU) exceed that, so derive the global scale from this call's amax. + x_global_sf = (448.0 * 6.0) / x.abs().amax().float().clamp(min=1e-12) + else: + x_global_sf = self.x_global_sf x_fp4, x_scale = _nvfp4_quantize( x, x_global_sf, @@ -552,6 +606,11 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str): method = NVFP4QuantizeMethod(layer_prefix=prefix) method._retain_original_weights = self.retain_original_weights return method + if (self.layer_profile == "h3_dit_ffn" and os.environ.get("FASTVIDEO_H3_FP8_ATTENTION") == "1" + and _H3_BLOCK_ATTN_PROJ.search(prefix) is not None): + # Mixed precision: NVFP4 MLPs, FP8 (per-tensor weight, dynamic per-tensor activation) attention. + from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod + return FP8QuantizeMethod(granularity=os.environ.get("FASTVIDEO_H3_FP8_GRANULARITY", "tensor")) return None @@ -650,7 +709,8 @@ def load_minimax_h3_nvfp4_dit_export( """Load a packed NVFP4H3 DiT export onto already-tagged NVFP4 linears. Keys are ``::`` with the four buffers - ``convert_model_to_nvfp4`` registers. Every export prefix must match an + ``convert_model_to_nvfp4`` registers, plus an optional calibrated + ``_nvfp4_input_global_sf``. Every export prefix must match an NVFP4 linear, and every NVFP4 linear must appear in the export. """ from safetensors import safe_open @@ -673,7 +733,7 @@ def load_minimax_h3_nvfp4_dit_export( missing_buffers = [name for name in H3_NVFP4_DIT_BUFFER_NAMES if name not in buffers] if missing_buffers: raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} is missing {missing_buffers}") - extra_buffers = sorted(set(buffers) - set(H3_NVFP4_DIT_BUFFER_NAMES)) + extra_buffers = sorted(set(buffers) - set(H3_NVFP4_DIT_BUFFER_NAMES) - {H3_NVFP4_DIT_INPUT_SF_NAME}) if extra_buffers: raise ValueError(f"MiniMax-H3 NVFP4 DiT export layer {prefix!r} has unknown buffers {extra_buffers}") module = _module_by_nvfp4_export_prefix(modules, prefix) @@ -682,7 +742,9 @@ def load_minimax_h3_nvfp4_dit_export( if not isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod): raise RuntimeError("MiniMax-H3 NVFP4 DiT export layer " f"{prefix!r} is not an NVFP4 linear; set NVFP4Config(layer_profile='h3_dit')") - for buffer_name in H3_NVFP4_DIT_BUFFER_NAMES: + for buffer_name in H3_NVFP4_DIT_BUFFER_NAMES + (H3_NVFP4_DIT_INPUT_SF_NAME, ): + if buffer_name not in buffers: + continue tensor = reader.get_tensor(buffers[buffer_name]).to(device=device) module.register_buffer(buffer_name, tensor, persistent=False) module.register_parameter("weight", None) @@ -700,6 +762,7 @@ def load_minimax_h3_nvfp4_dit_export( __all__ = [ "H3_NVFP4_DIT_BUFFER_NAMES", "H3_NVFP4_DIT_EXPORT_FILENAME", + "H3_NVFP4_DIT_INPUT_SF_NAME", "NVFP4Config", "NVFP4QuantizeMethod", "convert_model_to_nvfp4", diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index a15d233fa0..c092e4300f 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -4,6 +4,8 @@ from __future__ import annotations import contextlib +import json +import os import math from collections.abc import Iterator, Mapping from typing import Any @@ -34,7 +36,7 @@ from fastvideo.layers.visual_embedding import Timesteps from fastvideo.logger import init_logger from fastvideo.models.dits.base import BaseDiT -from fastvideo.models.dits.minimax_h3_vsa_fp4 import (vsa_fp4_attention, vsa_fp4_attention_sp, +from fastvideo.models.dits.minimax_h3_vsa_fp4 import (STAGES, vsa_fp4_attention, vsa_fp4_attention_sp, vsa_fp4_requested) from fastvideo.models.dits.minimax_h3_fusions import ( HAVE_TRITON, @@ -317,7 +319,8 @@ def forward( else: hidden_states = vsa_fp4_attention_sp(self, hidden_states, rotary_emb, meta, use_fused_rope, get_sp_group()) - hidden_states, _ = self.to_out(hidden_states) + 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) @@ -458,7 +461,51 @@ def __init__( prefix=f"{prefix}.linear", ) + def enable_host_cache(self, table: dict | None = None) -> None: + """Keep the projection in pinned host memory and cache its output per timestep set. + + The modulation is a pure function of the timestep embedding, and few-step checkpoints sample a fixed + timestep ladder, so each block's output is a small constant table. The weights (the largest bf16 tensors + in the DiT) then never occupy device memory: a cache miss copies them in for one matmul. + """ + weight, bias = self.linear.weight, self.linear.bias + if table is not None and (weight is None or bias is None): + # Table-only load skipped these tensors entirely. + self._host_weight = self._host_bias = None + self._modulation_cache = dict(table) + self._cache_key = None + return + weight, bias = weight.data, bias.data + if table is None: + self._host_weight = weight.to("cpu").pin_memory() + self._host_bias = bias.to("cpu").pin_memory() + else: + # Precomputed modulation for a fixed timestep ladder: the projection weights are not needed at all. + self._host_weight = self._host_bias = None + self.linear.weight.data = torch.empty(0, dtype=weight.dtype) + self.linear.bias.data = torch.empty(0, dtype=bias.dtype) + self._modulation_cache: dict[Any, torch.Tensor] = dict(table or {}) + self._cache_key: Any = None + def forward(self, temb: torch.Tensor) -> tuple[torch.Tensor, ...]: + cache = getattr(self, "_modulation_cache", None) + if cache is not None: + out = cache.get(self._cache_key) + if out is None: + if self._host_weight is None: + raise RuntimeError(f"No precomputed AdaLN modulation for timestep key {self._cache_key}; the " + "table only covers the checkpoint's fixed ladder. Load without " + "FASTVIDEO_H3_ADALN_TABLE to use other timesteps.") + x = F.silu(temb) if self.apply_silu else temb + weight = self._host_weight.to(temb.device, non_blocking=True) + bias = self._host_bias.to(temb.device, non_blocking=True) + out = F.linear(x.to(weight.dtype), weight, bias) + if self._cache_key is not None: + cache[self._cache_key] = out + if os.environ.get("FASTVIDEO_H3_ADALN_DUMP"): + # Projection inputs, kept only when dumping, for offline low-rank fits. + self.__dict__.setdefault("_modulation_inputs", {})[self._cache_key] = x.detach() + return out.view(-1, 6 * self.hidden_size).chunk(6, dim=-1) if self.apply_silu: temb = F.silu(temb) temb, _ = self.linear(temb.to(self.linear.weight.dtype)) @@ -601,7 +648,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) - with nvtx_range("minimax_h3.transformer_block.feed_forward"): + 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(): return _gated_residual(hidden_states, gate_mlp, adaln_indices, feed_forward_output) @@ -930,6 +977,45 @@ def _rotary_for(self, position_ids: torch.Tensor, dtype: torch.dtype) -> tuple[t self._rope_cache = (position_ids, dtype, value) return value + def enable_adaln_host_cache(self, table_path: str | None = None) -> None: + """Move every block's AdaLN projection to pinned host memory behind a per-timestep cache. + + With ``table_path`` (written by FASTVIDEO_H3_ADALN_DUMP), the cache is prefilled from precomputed + modulation tables and the projection weights are dropped entirely. + """ + tables = None + if table_path: + import ast + raw = torch.load(table_path, map_location="cpu") + tables = {int(i): {ast.literal_eval(k): v for k, v in blk.items()} for i, blk in raw.items()} + for index, block in enumerate(self.transformer_blocks): + block.adaln_proj.enable_host_cache(None if tables is None else tables[index]) + self._adaln_host_cache = True + self._adaln_dumped_entries = -1 + + def _move_adaln_tables(self, device: torch.device) -> None: + for block in self.transformer_blocks: + cache = block.adaln_proj._modulation_cache + for key, value in cache.items(): + if value.device != device: + cache[key] = value.to(device) + + def _maybe_dump_adaln_tables(self) -> None: + path = os.environ.get("FASTVIDEO_H3_ADALN_DUMP") + if not path: + return + entries = sum(len(b.adaln_proj._modulation_cache) for b in self.transformer_blocks) + if entries == self._adaln_dumped_entries: + return + self._adaln_dumped_entries = entries + if model_parallel_is_initialized() and get_sp_group().rank_in_group != 0: + return + torch.save({i: {repr(k): v.detach().cpu() for k, v in b.adaln_proj._modulation_cache.items()} + for i, b in enumerate(self.transformer_blocks)}, path) + first = self.transformer_blocks[0].adaln_proj.__dict__.get("_modulation_inputs", {}) + torch.save({repr(k): v.detach().cpu() for k, v in first.items()}, path + ".inputs") + logger.info("Dumped AdaLN modulation tables (%d entries) to %s", entries, path) + def _refined_text(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor: """The prompt embedding is constant across the denoising loop and the refiner blocks are timestep-free, so refine once per generation @@ -1011,10 +1097,17 @@ def forward( local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0) rotary_emb = (rotary_cos, rotary_sin) + if getattr(self, "_adaln_host_cache", False): + # One host read per forward keys every block's modulation cache by the timestep values. + key = (tuple(timestep.reshape(-1).tolist()), tuple(temb.shape), str(temb.dtype)) + for block in self.transformer_blocks: + block.adaln_proj._cache_key = key + self._move_adaln_tables(temb.device) + # 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): - with nvtx_range(f"minimax_h3.transformer_block.{block_index}"): + with nvtx_range(f"minimax_h3.transformer_block.{block_index}"), STAGES.span("block_total"): packed_hidden_states = block( packed_hidden_states, temb, @@ -1036,6 +1129,12 @@ def forward( video_output = video_output.index_select(1, video_indices) audio_output = audio_output.index_select(1, audio_indices) + if getattr(self, "_adaln_host_cache", False): + self._maybe_dump_adaln_tables() + if STAGES.enabled: + stages = STAGES.flush() + if not model_parallel_is_initialized() or get_sp_group().rank_in_group == 0: + logger.info("H3_STAGE_MS %s", json.dumps(stages)) return video_output, audio_output diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index c83e6f7357..1291c47e32 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -29,6 +29,48 @@ VSA_FP4_ENV = "FASTVIDEO_H3_VSA_FP4" _BLOCK = 128 + +class _StageTimer: + """Opt-in (``FASTVIDEO_H3_SP_PROFILE=1``) CUDA-event spans summed per stage over one DiT forward. + + Events are only recorded during the forward; ``flush`` synchronizes once + and returns milliseconds per stage, so enabling it costs one sync per step. + """ + + def __init__(self) -> None: + self.enabled = os.environ.get("FASTVIDEO_H3_SP_PROFILE", "0") == "1" + self._spans: list[tuple[str, torch.cuda.Event, torch.cuda.Event]] = [] + + def span(self, name: str): + import contextlib + if not self.enabled: + return contextlib.nullcontext() + return self._record(name) + + def _record(self, name: str): + import contextlib + + @contextlib.contextmanager + def ctx(): + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + yield + end.record() + self._spans.append((name, start, end)) + + return ctx() + + def flush(self) -> dict[str, float]: + torch.cuda.synchronize() + totals: dict[str, float] = {} + for name, start, end in self._spans: + totals[name] = totals.get(name, 0.0) + start.elapsed_time(end) + self._spans.clear() + return {k: round(v, 1) for k, v in totals.items()} + + +STAGES = _StageTimer() + _fp4_api: Any = None @@ -101,35 +143,46 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ api = _api() layout = _layout_for(meta, rotary_emb) heads, dim = attn.num_attention_heads, attn.attention_head_dim - x_tiles = layout.gather_in(hidden_states) - 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)) - if use_fused_rope: - from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope - cos, sin = layout.cos.to(query.dtype), layout.sin.to(query.dtype) - query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) - key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) - else: - rope = (layout.cos, layout.sin) - query = attn._apply_rotary_emb(attn.norm_q(query), rope) - key = attn._apply_rotary_emb(attn.norm_k(key), rope) + with STAGES.span("qkv_proj_rope"): + x_tiles = layout.gather_in(hidden_states) + 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)) + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + cos, sin = layout.cos.to(query.dtype), layout.sin.to(query.dtype) + query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) + else: + rope = (layout.cos, layout.sin) + query = attn._apply_rotary_emb(attn.norm_q(query), rope) + key = attn._apply_rotary_emb(attn.norm_k(key), rope) + + sim_fp8 = os.environ.get(SIM_SP_FP8_ENV, "0") == "1" + if sim_fp8: + query, key, value = (_fp8_roundtrip(t) for t in (query, key, value)) vbs = meta.variable_block_sizes logical = layout.n_tiles * layout.tile - q_pooled = _pool_tiles(query[:, :logical], vbs, layout.tile) - k_pooled = _pool_tiles(key[:, :logical], vbs, layout.tile) - 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) - q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) - out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) - out = out.transpose(1, 2).index_select(1, layout.untile) # [B, L, H, D], packed order + with STAGES.span("select_mask"): + q_pooled = _pool_tiles(query[:, :logical], vbs, layout.tile) + k_pooled = _pool_tiles(key[:, :logical], vbs, layout.tile) + 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) + q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) + with STAGES.span("fp4_attention"): + out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) + with STAGES.span("out_untile"): + out = out.transpose(1, 2).index_select(1, layout.untile) # [B, L, H, D], packed order + if sim_fp8: + out = _fp8_roundtrip(out) if attn.to_gate_compress is not None and attn._gate_active(): - gate, _ = attn.to_gate_compress(hidden_states) - v_pooled = _pool_tiles(value[:, :logical], vbs, layout.tile) - out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled).permute(0, 2, 1, 3).to(out.dtype) - out = out.addcmul_(out_c.index_select(1, layout.row_tile), gate.unflatten(-1, (heads, dim))) + with STAGES.span("gate_compress"): + gate, _ = attn.to_gate_compress(hidden_states) + v_pooled = _pool_tiles(value[:, :logical], vbs, layout.tile) + out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled).permute(0, 2, 1, 3).to(out.dtype) + out = out.addcmul_(out_c.index_select(1, layout.row_tile), gate.unflatten(-1, (heads, dim))) return out.flatten(2, 3) @@ -149,6 +202,19 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ _FP8 = torch.float8_e4m3fn _FP8_MAX = 448.0 +# Debug: apply the SP path's FP8 rounding on one GPU (q/k/v after RoPE and the +# attention output, one scale per token and head), to separate exchange +# rounding from sharding errors when comparing SP>1 against SP=1. +SIM_SP_FP8_ENV = "FASTVIDEO_H3_SIM_SP_FP8" + + +@torch.compile(dynamic=True, fullgraph=True) +def _fp8_roundtrip(x: torch.Tensor) -> torch.Tensor: + """``[..., H, D]`` BF16 -> FP8 with a per-(token, head) scale -> BF16, as the SP exchange does.""" + xf = x.float() + scale = (xf.abs().amax(dim=-1) / _FP8_MAX).clamp_min(1e-12) + return ((xf / scale[..., None]).to(_FP8).float() * scale[..., None]).to(torch.bfloat16) + @torch.compile(dynamic=False, fullgraph=True) def _pack_heads_fp8(query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, @@ -242,41 +308,53 @@ def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tup layout = _SPTileLayout(meta, rank, local_rows) meta._h3_fp4_sp_layout = layout # type: ignore[attr-defined] - query, key, value = (t.unflatten(-1, (heads, dim)) - for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), hidden_states)) - if use_fused_rope: - from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope - cos, sin = rotary_emb[0].to(query.dtype), rotary_emb[1].to(query.dtype) - query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) - key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) - else: - query = attn._apply_rotary_emb(attn.norm_q(query), rotary_emb) - key = attn._apply_rotary_emb(attn.norm_k(key), rotary_emb) - - payload, scale = _pack_heads_fp8(query[0], key[0], value[0], world) - payload, scale = _all_to_all(payload, scale, sp_group.device_group) - qkv = layout.tiles_from(_unpack_seq_fp8(payload, scale, layout.seq_len)) # [3, R, Hs, D] + with STAGES.span("qkv_proj_rope"): + query, key, value = (t.unflatten(-1, (heads, dim)) + for t in _shared_input_projections((attn.to_q, attn.to_k, attn.to_v), hidden_states)) + if use_fused_rope: + from fastvideo.models.dits.minimax_h3_fusions import fused_qknorm_rope + cos, sin = rotary_emb[0].to(query.dtype), rotary_emb[1].to(query.dtype) + query = fused_qknorm_rope(query, attn.norm_q.weight, cos, sin, attn.norm_q.eps) + key = fused_qknorm_rope(key, attn.norm_k.weight, cos, sin, attn.norm_k.eps) + else: + query = attn._apply_rotary_emb(attn.norm_q(query), rotary_emb) + key = attn._apply_rotary_emb(attn.norm_k(key), rotary_emb) + + with STAGES.span("qkv_pack"): + payload, scale = _pack_heads_fp8(query[0], key[0], value[0], world) + with STAGES.span("qkv_all_to_all"): + payload, scale = _all_to_all(payload, scale, sp_group.device_group) + with STAGES.span("qkv_unpack_tile"): + qkv = layout.tiles_from(_unpack_seq_fp8(payload, scale, layout.seq_len)) # [3, R, Hs, D] q_t, k_t, v_t = qkv[0:1], qkv[1:2], qkv[2:3] vbs = meta.variable_block_sizes logical = layout.n_tiles * layout.tile - scores = torch.matmul(_pool_tiles(q_t[:, :logical], vbs, layout.tile), - _pool_tiles(k_t[:, :logical], vbs, layout.tile).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) - q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) - out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad) - - payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows) - payload, scale = _all_to_all(payload, scale, sp_group.device_group) - out = _unpack_heads_fp8(payload, scale) # [rows, H, D] + with STAGES.span("select_mask"): + scores = torch.matmul(_pool_tiles(q_t[:, :logical], vbs, layout.tile), + _pool_tiles(k_t[:, :logical], vbs, layout.tile).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) + q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) + with STAGES.span("fp4_attention"): + out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad) + + with STAGES.span("out_pack"): + payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows) + with STAGES.span("out_all_to_all"): + payload, scale = _all_to_all(payload, scale, sp_group.device_group) + with STAGES.span("out_unpack"): + out = _unpack_heads_fp8(payload, scale) # [rows, H, D] if attn.to_gate_compress is not None and attn._gate_active(): - v_pooled = _pool_tiles(v_t[:, :logical], vbs, layout.tile) - out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled)[0].to(out.dtype) # [Hs, n_tiles, D] - gathered = torch.empty((world, *out_c.shape), dtype=out_c.dtype, device=out_c.device) - dist.all_gather_into_tensor(gathered, out_c.contiguous(), group=sp_group.device_group) - out_c_all = gathered.flatten(0, 1).transpose(0, 1) # [n_tiles, H, D] - gate, _ = attn.to_gate_compress(hidden_states) - out = _apply_gate(out, out_c_all, layout.local_row_tile, gate[0].unflatten(-1, (heads, dim))) + with STAGES.span("gate_compress"): + v_pooled = _pool_tiles(v_t[:, :logical], vbs, layout.tile) + out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled)[0].to(out.dtype) # [Hs, n_tiles, D] + gathered = torch.empty((world, *out_c.shape), dtype=out_c.dtype, device=out_c.device) + with STAGES.span("gate_all_gather"): + dist.all_gather_into_tensor(gathered, out_c.contiguous(), group=sp_group.device_group) + with STAGES.span("gate_apply"): + out_c_all = gathered.flatten(0, 1).transpose(0, 1) # [n_tiles, H, D] + gate, _ = attn.to_gate_compress(hidden_states) + out = _apply_gate(out, out_c_all, layout.local_row_tile, gate[0].unflatten(-1, (heads, dim))) return out.flatten(1, 2).unsqueeze(0) diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 0c96e636cf..7082e26256 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -182,6 +182,38 @@ def _nvfp4_linear( ) +_E2M1_VALUES = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0) + + +def _fp4_gemm_supported(device: torch.device) -> bool: + """FP4 tensor-core GEMMs exist on Blackwell (sm_100 / sm_120) and newer.""" + return torch.cuda.is_available() and torch.cuda.get_device_capability(device)[0] >= 10 + + +def unswizzle_128x4_scales(scale: torch.Tensor, rows: int, cols: int) -> torch.Tensor: + """FlashInfer ``layout_128x4`` block scales -> row-major ``[rows, cols]`` (E4M3 bytes). + + The swizzled buffer holds the padded matrix ``[ceil(rows/128)*128, ceil(cols/4)*4]`` as + ``(row_tile, col_tile, row % 32, (row // 32) % 4, col % 4)``. + """ + pad_rows, pad_cols = -(-rows // 128) * 128, -(-cols // 4) * 4 + tiles = scale.reshape(-1)[:pad_rows * pad_cols].view(pad_rows // 128, pad_cols // 4, 32, 4, 4) + return tiles.permute(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols] + + +def dequantize_serialized_nvfp4(weight_packed: torch.Tensor, weight_scale: torch.Tensor, global_scale: float, + dtype: torch.dtype = torch.bfloat16) -> torch.Tensor: + """Packed E2M1 ``[out, in // 2]`` + swizzled E4M3 scales -> dense ``[out, in]`` weight.""" + out_features, in_features = weight_packed.shape[0], weight_packed.shape[1] * 2 + lut = torch.tensor(_E2M1_VALUES, device=weight_packed.device, dtype=torch.float32) + packed = weight_packed.view(torch.uint8) + values = torch.stack((lut[(packed & 0x0F).long()], lut[(packed >> 4).long()]), dim=-1).reshape( + out_features, in_features) + scales = unswizzle_128x4_scales(weight_scale.view(torch.uint8), out_features, in_features // NVFP4_GROUP_SIZE) + scales = scales.view(torch.float8_e4m3fn).float().repeat_interleave(NVFP4_GROUP_SIZE, dim=1) + return (values * scales / global_scale).to(dtype) + + class MiniMaxH3SerializedNVFP4Config(QuantizationConfig): """Serialized 16-group NVFP4 contract for the H3 text encoder. @@ -384,6 +416,12 @@ def process_weights_after_loading(self, layer: nn.Module) -> None: @staticmethod def _apply_finalized(layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None) -> 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, + float(layer.weight_global_scale.item()), x.dtype) + return torch.nn.functional.linear(x, weight, None if bias is None else bias.to(x.dtype)) x = _coerce_fp4_input_dtype(x) original_shape = x.shape if x.numel() == 0: diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index ae50f6172f..05312e92a2 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -1139,6 +1139,16 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): os.environ.get("RANK", "0"), resolved.name if resolved else "automatic selection", local_main_process_only=False) + # Layerwise offload keeps every block's weights in pinned host memory, so load them on the CPU and + # attach the hooks before anything moves to the GPU; loading on the GPU first would need the whole + # DiT resident once, which is exactly what offload exists to avoid on small cards. + layerwise_load = (fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload + and not fastvideo_args.use_fsdp_inference) + # The AdaLN host cache also needs the projection weights to stay off the device from the start. + adaln_table = os.environ.get("FASTVIDEO_H3_ADALN_TABLE") or None + adaln_host_cache = (fastvideo_args.inference_mode and not fastvideo_args.use_fsdp_inference + and (os.environ.get("FASTVIDEO_H3_ADALN_CACHE") == "1" or adaln_table is not None)) + layerwise_load = layerwise_load or adaln_host_cache model = maybe_load_fsdp_model( model_cls=model_cls, init_params={ @@ -1146,7 +1156,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): "hf_config": hf_config }, weight_dir_list=safetensors_list, - device=get_local_torch_device(), + device=torch.device("cpu") if layerwise_load else get_local_torch_device(), hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim, hsdp_shard_dim=fastvideo_args.hsdp_shard_dim, strict=strict_load, @@ -1180,15 +1190,26 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): model = model.eval() + if adaln_host_cache and hasattr(model, "enable_adaln_host_cache"): + model.enable_adaln_host_cache(adaln_table) + logger.info("AdaLN modulation: %s", "precomputed tables from " + adaln_table if adaln_table + else "projections in pinned host memory behind a per-timestep cache") + if layerwise_load and not fastvideo_args.dit_layerwise_offload: + model = model.to(get_local_torch_device()) + if fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload: # Check if model has nn.ModuleList for layerwise offload compatibility has_module_list = any(isinstance(m, nn.ModuleList) for m in model.children()) if has_module_list: enable_layerwise_offload(model) + # Blocks now hold placeholders; the remaining (non-block) weights and buffers belong on the GPU. + model = model.to(get_local_torch_device()) else: logger.warning( "Layerwise offload requested but model %s does not have " "nn.ModuleList structure. Skipping layerwise offload.", cls_name) + if layerwise_load: + model = model.to(get_local_torch_device()) return model diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 911f878963..b63b50a067 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -36,6 +36,54 @@ def _summarize_param_names(names: set[str]) -> str: return ", ".join(f"{family} x{count}" if count > 1 else family for family, count in sorted(families.items())) +def _prequantized_fp8_prefixes(weight_files: list[str]) -> list[str]: + """Checkpoint prefixes stored as FP8 W8A8: float8_e4m3fn ``weight`` plus per-channel float32 ``weight_scale``.""" + from safetensors import safe_open + + fp8_weights: set[str] = set() + scales: set[str] = set() + for path in weight_files: + with safe_open(path, framework="pt", device="cpu") as reader: + for key in reader.keys(): # noqa: SIM118 + if key.endswith(".weight") and reader.get_slice(key).get_dtype() == "F8_E4M3": + fp8_weights.add(key[:-len(".weight")]) + elif key.endswith(".weight_scale") and reader.get_slice(key).get_dtype() == "F32": + scales.add(key[:-len(".weight_scale")]) + if fp8_weights - scales: + raise ValueError(f"FP8 checkpoint weights without a weight_scale: {sorted(fp8_weights - scales)[:4]}") + return sorted(fp8_weights) + + +def _load_prequantized_fp8(model: nn.Module, weight_files: list[str], targets: dict[str, str], + device: torch.device) -> None: + """Attach FP8 W8A8 checkpoint tensors to their FP8 linears as ``convert_model_to_fp8`` would (per channel).""" + from safetensors import safe_open + + from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod + + modules = dict(model.named_modules()) + wanted = {f"{prefix}.{suffix}" for prefix in targets for suffix in ("weight", "weight_scale")} + tensors: dict[str, torch.Tensor] = {} + for path in weight_files: + with safe_open(path, framework="pt", device="cpu") as reader: + for key in wanted.intersection(reader.keys()): + tensors[key] = reader.get_tensor(key) + for prefix, target in targets.items(): + module = modules.get(target[:-len(".weight")]) + quant_method = getattr(module, "quant_method", None) + if not isinstance(quant_method, FP8QuantizeMethod): + raise RuntimeError(f"Pre-quantized FP8 weight {prefix!r} maps to {target!r}, which is not an FP8 linear; " + "load the checkpoint with an FP8 quant_config covering it.") + # The checkpoint scales are per output channel; activations then use per-token scales. + quant_method.granularity = "channel" + module.register_buffer("_fp8_weight", tensors[f"{prefix}.weight"].to(device).contiguous(), persistent=False) + module.register_buffer("_fp8_weight_scale", + tensors[f"{prefix}.weight_scale"].reshape(-1).to(device=device, dtype=torch.float32), + persistent=False) + module._parameters.pop("weight", None) + logger.info("Loaded %d pre-quantized FP8 linears", len(targets)) + + def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lora_merge: bool = False) -> None: """Quantize inference linear weights after checkpoint loading. @@ -80,16 +128,23 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor for mod in model.modules(): qm = getattr(mod, "quant_method", None) if isinstance(qm, NVFP4QuantizeMethod): + mixed_fp8 = any(isinstance(getattr(m, "quant_method", None), FP8QuantizeMethod) for m in model.modules()) if any( getattr(module, "_nvfp4_weight", None) is not None for module in model.modules() if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod)): logger.info("NVFP4 packed export already populated; skipping runtime weight conversion") + if mixed_fp8: + logger.info("Converting the FP8 linears of a mixed NVFP4/FP8 model") + convert_model_to_fp8(model) return if defer_weight_conversion_until_lora_merge: logger.info("Deferring NVFP4 weight conversion until the inference LoRA merge completes") return logger.info("Converting loaded model weights for NVFP4 linear layers") convert_model_to_nvfp4(model) + if any(isinstance(getattr(m, "quant_method", None), FP8QuantizeMethod) for m in model.modules()): + logger.info("Converting the FP8 linears of a mixed NVFP4/FP8 model") + convert_model_to_fp8(model) return if isinstance(qm, NVFP4QATQuantizeMethod): logger.info("Converting loaded model weights for NVFP4-QAT linear layers") @@ -324,6 +379,20 @@ def maybe_load_fsdp_model( "This LoRA adapter provides MiniMax H3 VSA compression gates, but the selected attention backend " "did not construct them. Use attention_backend='VIDEO_SPARSE_ATTN_H3'. Missing parameters: " + ", ".join(missing_vsa_gates[:3]) + (" ..." if len(missing_vsa_gates) > 3 else "")) + skip_param_names = set(nvfp4_skip_param_names) + fp8_targets: dict[str, str] = {} + fp8_prefixes = _prequantized_fp8_prefixes(weight_dir_list) + if fp8_prefixes: + fp8_targets = {prefix: param_names_mapping_fn(f"{prefix}.weight")[0] for prefix in fp8_prefixes} + skip_param_names |= set(fp8_targets.values()) + fp8_keys = {f"{prefix}.{suffix}" for prefix in fp8_prefixes for suffix in ("weight", "weight_scale")} + weight_iterator = ((name, tensor) for name, tensor in weight_iterator if name not in fp8_keys) + if os.environ.get("FASTVIDEO_H3_ADALN_TABLE"): + # Precomputed AdaLN modulation replaces the per-block projections; never read their weights. + skip_param_names |= {name for name, _ in model.named_parameters() + if re.fullmatch(r"transformer_blocks\.\d+\.adaln_proj\.linear\.(weight|bias)", name)} + logger.info("Skipping %d AdaLN projection tensors (precomputed modulation tables)", + sum("adaln_proj" in n for n in skip_param_names)) load_model_from_full_model_state_dict( model, weight_iterator, @@ -333,10 +402,12 @@ def maybe_load_fsdp_model( cpu_offload=cpu_offload, param_names_mapping=param_names_mapping_fn, dense_lora_patch=dense_lora_patch, - skip_param_names=nvfp4_skip_param_names or None, + skip_param_names=skip_param_names or None, ) if packed_nvfp4_export is not None: load_minimax_h3_nvfp4_dit_export(model, packed_nvfp4_export, device=device) + if fp8_targets: + _load_prequantized_fp8(model, weight_dir_list, fp8_targets, device=device) if hasattr(model, "materialize_non_persistent_buffers"): model.materialize_non_persistent_buffers(device=device, dtype=default_dtype) for n, p in chain(model.named_parameters(), model.named_buffers()): diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 1310f53fc9..afdd902d08 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -3,6 +3,7 @@ from __future__ import annotations import gc +import os import json import math from dataclasses import dataclass @@ -75,6 +76,33 @@ def _checkpoint_has_vsa_gates(transformer_dir: Path) -> bool: return False +def _pinned_swap(module: Any, device: torch.device) -> None: + """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 are copied every time. + """ + store = module.__dict__.setdefault("_pinned_host_tensors", {}) + params = dict(module.named_parameters()) + for name, tensor in list(params.items()) + list(module.named_buffers()): + if tensor is None: + continue + if device.type == "cpu": + if tensor.device.type == "cpu": + continue + host = store.get(name) if name in params else None + if host is None or host.shape != tensor.shape or host.dtype != tensor.dtype: + host = torch.empty(tensor.shape, dtype=tensor.dtype, pin_memory=True) + host.copy_(tensor) + if name in params: + store[name] = host + tensor.data = host + 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() + + def _module_has_dtensor_params(module: Any) -> bool: parameters = getattr(module, "parameters", None) if not callable(parameters): @@ -429,12 +457,26 @@ def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: def _move_module(self, module: Any, device: str | torch.device) -> bool: if _module_has_dtensor_params(module): return False - module.to(device) + if os.environ.get("FASTVIDEO_H3_PINNED_SWAP", "1") == "1": + _pinned_swap(module, torch.device(device)) + else: + module.to(device) return True + @staticmethod + def _parked_module_names() -> tuple[str, ...]: + """Denoise modules parked on the host while the text encoder runs (FASTVIDEO_H3_PARK_MODULES). + + Cards with room for the DiT next to the encoder park only the VAEs and keep the DiT resident. + """ + requested = os.environ.get("FASTVIDEO_H3_PARK_MODULES") + if not requested: + return _DENOISE_MODULE_NAMES + return tuple(name for name in requested.split(",") if name in _DENOISE_MODULE_NAMES) + def _park_denoise_modules(self) -> None: parked = False - for name in _DENOISE_MODULE_NAMES: + for name in self._parked_module_names(): module = self.get_module(name) if module is None: continue diff --git a/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py index 4bed84a162..bdb3f8e546 100644 --- a/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py +++ b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py @@ -21,7 +21,8 @@ def refuse_zero_initialized_h3_vsa(transformer: Any) -> None: if gate is None: continue weight = getattr(gate, "weight", None) - if weight is None: + # Layerwise offload leaves a zero-element placeholder on the module; it says nothing about the gate. + if weight is None or weight.numel() == 0: continue saw_weight = True if bool((weight != 0).any()): diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index 838ce677ad..c87d946f65 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -188,6 +188,12 @@ def _execute( torch.cuda.synchronize() execution_time = time.perf_counter() - start_time logger.info("[%s] Execution completed in %s ms", stage_name, execution_time * 1000) + if torch.cuda.is_available(): + gib = 1024**3 + logger.info("[%s] Memory peak_allocated=%.2f GiB reserved=%.2f GiB resident_after=%.2f GiB", + stage_name, torch.cuda.max_memory_allocated() / gib, + torch.cuda.memory_reserved() / gib, torch.cuda.memory_allocated() / gib) + torch.cuda.reset_peak_memory_stats() batch.logging_info.add_stage_execution_time(stage_key, execution_time) batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name) component_metric = self.performance_component_metric diff --git a/fastvideo/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index d0712091c6..4fdcae2c5e 100644 --- a/fastvideo/worker/gpu_worker.py +++ b/fastvideo/worker/gpu_worker.py @@ -21,6 +21,35 @@ def _log_cuda_device_uuid(rank: int, device: torch.device) -> None: logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False) + +def _log_pipeline_memory(pipeline) -> None: + """Debug (FASTVIDEO_MEMORY_REPORT=1): bytes held per pipeline component, by device and dtype, plus the + largest tensors, so the resident footprint can be attributed before choosing offload placements.""" + gib = 1024**3 + for name, module in getattr(pipeline, "modules", {}).items(): + if not isinstance(module, torch.nn.Module): + continue + by_kind: dict[str, int] = {} + largest: list[tuple[int, str, str]] = [] + seen: set[int] = set() + for tname, t in list(module.named_parameters()) + list(module.named_buffers()): + if t is None or id(t) in seen: + continue + seen.add(id(t)) + nbytes = t.numel() * t.element_size() + key = f"{t.device.type}/{str(t.dtype).replace('torch.', '')}" + by_kind[key] = by_kind.get(key, 0) + nbytes + largest.append((nbytes, tname, key)) + largest.sort(reverse=True) + total = sum(by_kind.values()) + logger.info("MEMREPORT %s total=%.2f GiB %s", name, total / gib, + {k: round(v / gib, 2) for k, v in sorted(by_kind.items(), key=lambda kv: -kv[1])}) + for nbytes, tname, key in largest[:8]: + logger.info("MEMREPORT %s %.3f GiB %s %s", name, nbytes / gib, key, tname) + if torch.cuda.is_available(): + logger.info("MEMREPORT cuda allocated=%.2f GiB reserved=%.2f GiB", torch.cuda.memory_allocated() / gib, + torch.cuda.memory_reserved() / gib) + class Worker: def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str): @@ -71,6 +100,12 @@ def init_device(self) -> None: # Set the CUDA device BEFORE any CUDA calls if current_platform.is_cuda_alike(): torch.cuda.set_device(self.device) + # Debug: FASTVIDEO_CUDA_MEMORY_CAP_GIB emulates a smaller card by capping this process's allocator. + cap_gib = os.environ.get("FASTVIDEO_CUDA_MEMORY_CAP_GIB") + if cap_gib: + total = torch.cuda.get_device_properties(self.device).total_memory + torch.cuda.set_per_process_memory_fraction(min(1.0, float(cap_gib) * 1024**3 / total), self.device) + logger.info("Capped CUDA allocator at %s GiB of %.1f GiB", cap_gib, total / 1024**3) self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0] if current_platform.is_cuda(): _log_cuda_device_uuid(self.rank, self.device) @@ -90,6 +125,8 @@ def init_device(self) -> None: self.distributed_init_method) self.pipeline = build_pipeline(self.fastvideo_args) + if os.environ.get("FASTVIDEO_MEMORY_REPORT") == "1" and self.rank == 0: + _log_pipeline_memory(self.pipeline) def execute_forward(self, forward_batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: output_batch = self.pipeline.forward(forward_batch, self.fastvideo_args) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index 543af88054..9b2058fade 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -7,7 +7,9 @@ import modal -WORKTREE = pathlib.Path(__file__).resolve().parents[3] +_HERE = pathlib.Path(__file__).resolve() +# Inside the container this file is /root/app.py; the repo is already baked into the image there. +WORKTREE = _HERE.parents[3] if len(_HERE.parents) > 3 else pathlib.Path("/src/fastvideo") CUTLASS_COMMIT = "e67e63c331d6e4b729047c95cf6b92c8454cba89" volume = modal.Volume.from_name("h3-pro6000-weights", create_if_missing=True) @@ -18,7 +20,7 @@ .add_local_dir(WORKTREE, "/src/fastvideo", copy=True, ignore=[".git", "**/__pycache__", "fastvideo-kernel/include/cutlass/**", "fastvideo-kernel/include/tk/**", "fastvideo/third_party/eval/**", "docs/**", - "assets/**", "comfyui/**", "apps/**", "**/*.mp4"]) + "assets/**", "comfyui/**", "apps/**", "**/*.mp4", "**/*.log"]) .run_commands("cd /src/fastvideo && UV_TORCH_BACKEND=cu130 uv pip install --system -e . --no-sources") .run_commands("uv pip install --system 'cmake==3.31.6' ninja 'scikit-build-core>=0.10' pybind11 hf_transfer") .run_commands(f"git clone --filter=blob:none https://github.com/NVIDIA/cutlass.git /cutlass && " @@ -186,8 +188,11 @@ def convert(minimal: bool = True) -> dict: def run_variant(name: str, model: str, profile: str, attention: str, decode: str, vae_compile: bool, height: int = 480, width: int = 832, warmups: int = 1, num_frames: int = 124, env: dict | None = None, prompts: tuple = ("kitesurf", "chef"), sparsity: float = 0.8, - steps: int = 9, num_gpus: int = 1, parallel_decode: bool = False, pre_runs: tuple = ()) -> dict: + steps: int = 9, num_gpus: int = 1, parallel_decode: bool = False, pre_runs: tuple = (), + prompt_texts: dict | None = None, offload: dict | None = None, + experimental_extra: dict | None = None) -> dict: os.environ.update(env or {}) + texts = {**PROMPTS, **(prompt_texts or {})} _install_kernel() import torch from fastvideo import VideoGenerator @@ -198,13 +203,14 @@ def run_variant(name: str, model: str, profile: str, attention: str, decode: str "vae_parallel_decode": parallel_decode, "video_decode_backend": decode} if attention == "VIDEO_SPARSE_ATTN_H3": experimental.update({"VSA_sparsity": sparsity, "VSA_tile_size": 64}) + experimental.update(experimental_extra or {}) config = { "model_path": f"/vol/fv/{model}", "engine": {"num_gpus": num_gpus, "use_fsdp_inference": False, "quantization": {"transformer_quant": "NVFP4", "layer_profile": profile}, "parallelism": {"tp_size": 1, "sp_size": num_gpus}, "offload": {"dit": False, "dit_layerwise": False, "text_encoder": False, "vae": False, - "pin_cpu_memory": num_gpus == 1, "lazy_module_load": False}, + "pin_cpu_memory": num_gpus == 1, "lazy_module_load": False, **(offload or {})}, "compile": {"enabled": False, "vae_enabled": vae_compile}}, "pipeline": {"experimental": experimental}, } @@ -219,7 +225,7 @@ def run_variant(name: str, model: str, profile: str, attention: str, decode: str try: # Optional runs at other shapes first (e.g. a 480p correctness clip), same loaded model. for j, (ph, pw, pf, pid) in enumerate(pre_runs): - request = {"prompt": PROMPTS[pid], "negative_prompt": "", + request = {"prompt": texts[pid], "negative_prompt": "", "sampling": {"seed": 20260929, "height": ph, "width": pw, "num_frames": pf, "fps": 24, "num_inference_steps": steps, "guidance_scale": 1.0, "batch_cfg": False}, "output": {"output_path": str(out_dir / f"pre{j:02d}_{pid}_{ph}p.mp4"), "save_video": True, @@ -229,9 +235,16 @@ def run_variant(name: str, model: str, profile: str, attention: str, decode: str results.setdefault("pre_runs", []).append({"shape": [ph, pw, pf], "prompt": pid, "wall_s": round(time.perf_counter() - t, 2), "video": getattr(result, "video_path", None)}) - order = [prompts[0]] * warmups + list(prompts) + print("PRE_RUN", json.dumps(results["pre_runs"][-1]), flush=True) + if pre_runs: + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + volume.commit() + # Warm every distinct prompt: prompt length changes the packed sequence, and shape-specialized + # compiled kernels would otherwise recompile inside the first timed run of each prompt. + distinct = list(dict.fromkeys(prompts)) + order = [distinct[i % len(distinct)] for i in range(warmups)] + list(prompts) for i, pid in enumerate(order): - prompt = PROMPTS[pid] + prompt = texts[pid] request = {"prompt": prompt, "negative_prompt": "", "sampling": {"seed": 20260929, "height": height, "width": width, "num_frames": num_frames, "fps": 24, "num_inference_steps": steps, "guidance_scale": 1.0, "batch_cfg": False}, @@ -248,6 +261,7 @@ def run_variant(name: str, model: str, profile: str, attention: str, decode: str results["peak_mem_gb_device"] = _sh("nvidia-smi --query-gpu=memory.used --format=csv,noheader").strip() finally: generator.shutdown() + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) volume.commit() return results @@ -297,11 +311,65 @@ def run_variant2(*args, **kwargs) -> dict: return run_variant.local(*args, **kwargs) -@app.function(gpu="RTX-PRO-6000:8", memory=65536, cpu=16, timeout=5400, volumes={"/vol": volume}) +@app.function(gpu="RTX-PRO-6000:8", memory=196608, cpu=16, timeout=5400, volumes={"/vol": volume}) def run_variant8(*args, **kwargs) -> dict: return run_variant.local(*args, **kwargs) +@app.function(gpu="RTX-PRO-6000:4", memory=131072, cpu=8, timeout=5400, volumes={"/vol": volume}) +def run_variant4(*args, **kwargs) -> dict: + return run_variant.local(*args, **kwargs) + + +@app.function(gpu="RTX-PRO-6000:8", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) +def bench8_pair(prompt_texts: dict, timed: tuple, warmups: int, models: tuple = ("v2", "v2_profile")) -> list: + """8-GPU 10 s 768p V2 8-step benchmark, then a stage-profiled pass, in one container (shared compile caches).""" + specs = {"v2": ("sp8_v2_8step_768p10s_1k", "v2_vsa_light", 0.8, 9, {}, timed, warmups), + # One warm + one profiled generation; CUDA-event spans log per DiT forward as H3_STAGE_MS. + "v2_profile": ("sp8_v2_8step_profile", "v2_vsa_light", 0.8, 9, {"FASTVIDEO_H3_SP_PROFILE": "1"}, + timed[:1], 1)} + out = [] + for m in models: + name, model, sparsity, steps, extra_env, prompts, warm = specs[m] + r = run_variant.local(name, model, "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", + vae_compile=True, height=768, width=1344, num_frames=243, warmups=warm, + env={**FAST_ENV, **extra_env}, prompts=prompts, sparsity=sparsity, steps=steps, + num_gpus=8, parallel_decode=True, prompt_texts=prompt_texts) + print("RESULT", json.dumps(r), flush=True) + out.append(r) + return out + + +@app.function(gpu="RTX-PRO-6000", memory=196608, cpu=16, timeout=3 * 3600, volumes={"/vol": volume}) +def memladder(prompt_texts: dict, pid: str, configs: dict | None = None) -> list: + """V2 8-step 10 s 768p on one GPU under memory placements; stage logs carry per-stage peaks. + + ``configs`` maps a name to (offload overrides, experimental overrides, extra env). The env can set + FASTVIDEO_CUDA_MEMORY_CAP_GIB to emulate a smaller card. + """ + seq = {"h3_sequential_load": True} + lw = {"text_encoder": True, "vae": True, "dit_layerwise": True} + configs = configs or { + "A_resident": ({}, {}, {}), + "B_seq_encoder_vae_offload": ({"text_encoder": True, "vae": True}, seq, {}), + "C_plus_dit_layerwise": (lw, seq, {}), + } + out = [] + for name, (offload, extra, env_extra) in configs.items(): + try: + r = run_variant.local(f"mem_{name}", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", + decode="h3-vae", vae_compile=False, height=768, width=1344, num_frames=243, + warmups=1, env={**FAST_ENV, **env_extra}, prompts=(pid,), sparsity=0.8, steps=9, + num_gpus=1, prompt_texts=prompt_texts, offload=offload, experimental_extra=extra) + except Exception as e: # keep the ladder going; record the failure + r = {"name": f"mem_{name}", "error": repr(e)[:3000]} + print("RESULT", json.dumps(r), flush=True) + out.append(r) + for k in env_extra: + os.environ.pop(k, None) + return out + + @app.function(cpu=8, memory=32768, timeout=1800, volumes={"/vol": volume}) def compare_videos(a: str, b: str) -> dict: """Frame PSNR between two MP4s on the volume (same seed and prompt).""" @@ -320,7 +388,7 @@ def compare_videos(a: str, b: str) -> dict: @app.local_entrypoint() -def main(step: str = "all"): +def main(step: str = "all", ladder: str = "base"): if step == "prep_personal": print("KERNEL", build_kernel.remote()) print("CONVERT", json.dumps(convert.remote(minimal=True), indent=1)[:6000]) @@ -346,6 +414,68 @@ def main(step: str = "all"): print("COMPARE", json.dumps(compare_videos.remote("/vol/outputs/sp1_480p/00_kitesurf.mp4", r["pre_runs"][0]["video"]))) return + if step == "simfp8": + # SP=1 with and without the SP exchange's FP8 rounding, same settings as sp1_480p. + common = dict(attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", vae_compile=False, height=480, width=832, + num_frames=124, warmups=0, prompts=("kitesurf",), sparsity=0.8, steps=9) + plain = run_variant.spawn("sp1_480p_rerun", "v2_vsa_light", "h3_dit_vsa", env=FAST_ENV, **common) + sim = run_variant.spawn("sp1_480p_simfp8", "v2_vsa_light", "h3_dit_vsa", + env={**FAST_ENV, "FASTVIDEO_H3_SIM_SP_FP8": "1"}, **common) + r_plain, r_sim = plain.get(), sim.get() + print("RESULT", json.dumps(r_plain)) + print("RESULT", json.dumps(r_sim)) + return + if step == "memreport": + r = run_variant.remote("memreport", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", + decode="h3-vae", vae_compile=False, height=480, width=832, num_frames=124, warmups=0, + env={**FAST_ENV, "FASTVIDEO_MEMORY_REPORT": "1"}, prompts=("kitesurf",), sparsity=0.8, + steps=9, offload={"text_encoder": True, "vae": True}, + experimental_extra={"h3_sequential_load": False}) + print("RESULT", json.dumps(r)) + return + if step == "memladder": + rows = [json.loads(line) for line in open(WORKTREE.parent / "UniServe-sm120fp4" / "uniserve_eval" / "workloads" + / "fast_h3" / "latency.jsonl")] + pid = "latency-ceramics-005" + text = {r["id"]: r["prompt"] for r in rows if r["id"] == pid} + configs = None + if ladder == "caps": + seq = {"h3_sequential_load": True} + lw = {"text_encoder": True, "vae": True, "dit_layerwise": True} + configs = { + "D32_C_cap32": (lw, seq, {"FASTVIDEO_CUDA_MEMORY_CAP_GIB": "32"}), + "D24_buffers_cap24": (lw, seq, {"FASTVIDEO_CUDA_MEMORY_CAP_GIB": "24", + "FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS": "1"}), + "D16_buffers_tile8_cap16": (lw, seq, {"FASTVIDEO_CUDA_MEMORY_CAP_GIB": "16", + "FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS": "1", + "FASTVIDEO_H3_VAE_TILE_BATCH": "8"}), + } + for r in memladder.remote(text, pid, configs): + print("RESULT", json.dumps(r)) + return + if step == "bench8_pair": + # UniServe's 10 s / ~1K-token latency prompts, so the numbers line up with its published protocol. + rows = [json.loads(line) for line in open(WORKTREE.parent / "UniServe-sm120fp4" / "uniserve_eval" / "workloads" + / "fast_h3" / "latency.jsonl")] + ten = {r["id"]: r["prompt"] for r in rows if r["seconds"] == 10 and r["prompt_len"] == 1000} + ids = ("latency-ceramics-005", "latency-harbor-005") + texts = {i: ten[i] for i in ids} + for r in bench8_pair.remote(texts, (ids[0], ids[1], ids[0], ids[1]), 2): + timed = sorted(x["wall_s"] for x in r["runs"] if not x["warmup"]) + print("SUMMARY", r["name"], "timed", timed, "median", timed[len(timed) // 2] if len(timed) % 2 + else (timed[len(timed) // 2 - 1] + timed[len(timed) // 2]) / 2) + return + if step == "sp4": + r = run_variant4.remote( + "sp4_v2_8step", "v2_vsa_light", "h3_dit_vsa", attention="VIDEO_SPARSE_ATTN_H3", decode="h3-vae", + vae_compile=True, height=768, width=1344, num_frames=243, warmups=1, env=FAST_ENV, + prompts=("kitesurf", "chef", "kitesurf", "chef"), sparsity=0.8, steps=9, num_gpus=4, parallel_decode=True, + pre_runs=((480, 832, 124, "kitesurf"), )) + print("RESULT", json.dumps(r)) + if r.get("pre_runs"): + print("COMPARE", json.dumps(compare_videos.remote("/vol/outputs/sp1_480p/00_kitesurf.mp4", + r["pre_runs"][0]["video"]))) + return if step in ("bench8_v2", "bench8_v4"): v2 = step == "bench8_v2" r = run_variant8.remote( diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py index a8de3a945c..8b2ff1b448 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -26,6 +26,14 @@ Without either the export holds the FFN linears only and must be loaded with ``layer_profile="h3_dit_ffn"``. +``--quantize-ffn`` takes the FFN linears from a dense BF16 source instead of a +ModelOpt export, with the same round-to-nearest weight math as the runtime +(ModelOpt's max calibration produces the same weight codes). ``--act-amax`` +adds a calibrated static activation scale per linear +(``_nvfp4_input_global_sf`` = 448 * 6 / amax) from a JSON of input amax keyed +``b.`` (e.g. ``b3.ff.fc_in``); without it activations use the +unit global scale, which saturates inputs above 2688 (H3's ``ff.fc_out``). + Every exported linear is probed on random BF16 rows through the same ``mm_fp4`` path the loader runs; the relative error against a BF16 matmul with the dequantized weight must stay under ``--max-probe-error`` (genuine W4A4 @@ -51,6 +59,7 @@ EXPORT_FILENAME = "nvfp4_weights.safetensors" _BLOCK_ATTN = re.compile(r"^transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out\.0)$") +_BLOCK_FFN = re.compile(r"^transformer_blocks\.\d+\.ff\.net\.(?:0\.proj|2)$") _BLOCK_GATE = re.compile(r"^transformer_blocks\.\d+\.attn\.to_gate_compress$") _RENAMES = ((re.compile(r"\.ff\.net\.0\.proj$"), ".ff.fc_in"), (re.compile(r"\.ff\.net\.2$"), ".ff.fc_out"), (re.compile(r"\.attn\.to_out\.0$"), ".attn.to_out")) @@ -140,6 +149,8 @@ def main() -> None: parser.add_argument("--dst", required=True, type=Path) parser.add_argument("--quantize-attention", action="store_true") parser.add_argument("--quantize-gate", action="store_true", help="also quantize attn.to_gate_compress") + parser.add_argument("--quantize-ffn", action="store_true", help="quantize the FFN linears from a bf16 source") + parser.add_argument("--act-amax", type=Path, help="JSON of calibrated input amax per linear") parser.add_argument("--max-probe-error", type=float, default=0.3) parser.add_argument("--dense-shard-gb", type=float, default=5.0) args = parser.parse_args() @@ -152,9 +163,17 @@ def main() -> None: modelopt_keys = {f"{p}.{s}" for p in modelopt for s in ("weight", "weight_scale", "weight_scale_2", "input_scale")} if args.quantize_gate and not args.quantize_attention: parser.error("--quantize-gate requires --quantize-attention") - dense_pattern = re.compile(_BLOCK_ATTN.pattern + ("|" + _BLOCK_GATE.pattern if args.quantize_gate else "")) + selected = ([_BLOCK_ATTN.pattern] if args.quantize_attention else []) + ( + [_BLOCK_GATE.pattern] if args.quantize_gate else []) + ([_BLOCK_FFN.pattern] if args.quantize_ffn else []) + dense_pattern = re.compile("|".join(selected)) if selected else None attention = sorted(k[:-len(".weight")] for k in weight_map - if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")])) if args.quantize_attention else [] + if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")])) if dense_pattern else [] + if args.quantize_ffn and any(_BLOCK_FFN.match(p) for p in modelopt): + parser.error("--quantize-ffn needs a bf16 source; this one already holds ModelOpt FFN linears") + amax_table = None + if args.act_amax: + raw = json.loads(args.act_amax.read_text()) + amax_table = {k: float(v["all"] if isinstance(v, dict) else v) for k, v in raw.items()} attention_keys = {f"{p}.weight" for p in attention} readers = {shard: safe_open(str(args.src / shard), framework="pt", device="cpu") for shard in set(weight_map.values())} @@ -177,6 +196,13 @@ def main() -> None: if error > args.max_probe_error: raise SystemExit(f"probe error {error:.3f} on {prefix} exceeds {args.max_probe_error}; nothing written") module = fastvideo_module_name(prefix) + if amax_table is not None: + block = re.match(r"transformer_blocks\.(\d+)\.(.+)$", module) + key = f"b{block.group(1)}.{block.group(2)}" + if key not in amax_table: + raise SystemExit(f"--act-amax has no entry {key!r} for {module}; nothing written") + buffers["_nvfp4_input_global_sf"] = torch.tensor((448.0 * 6.0) / max(amax_table[key], 1e-12), + dtype=torch.float32) for name, value in buffers.items(): export[f"{module}::{name}"] = value.cpu() save_file(export, str(args.dst / EXPORT_FILENAME)) @@ -206,7 +232,8 @@ def main() -> None: if extra.suffix not in (".safetensors", ".json"): shutil.copy2(extra, args.dst / extra.name) print(json.dumps({"exported_linears": len(modelopt) + len(attention), "modelopt_linears": len(modelopt), - "quantized_attention_linears": len(attention), "worst_probe_error": round(worst, 4), + "quantized_dense_linears": len(attention), "worst_probe_error": round(worst, 4), + "static_activation_scales": len(attention) + len(modelopt) if amax_table else 0, "gate_linears": sum(1 for p in attention if _BLOCK_GATE.match(p)), "min_scale_byte_agreement": round(min(agreements), 4) if agreements else None, "mean_scale_byte_agreement": round(sum(agreements) / len(agreements), 4) if agreements else None, From 4012c372106919746c77994437ef63f51a7312ac Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 09:12:48 -0700 Subject: [PATCH 009/107] [wip]: H3 4090 groundwork, step splice, encoder fallback on sm89 - FP8 on sm89: per-tensor GEMM + Triton per-token x per-channel scale epilogue (torch rowwise _scaled_mm runs ~70 TFLOPS there, below bf16) and a fused one-launch per-token quantize (5x faster than the torch chain) - FASTVIDEO_H3_FFN_CHUNK_TOKENS: inference-only FFN token chunking - FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS: keep the first N blocks resident - Serialized NVFP4 text encoder: allow sm80-sm90 through the bf16 de-quant path - FASTVIDEO_H3_SPLICE_TRANSFORMER / _FROM_STEP: second checkpoint runs late DMD steps (cherry picked from commit e2ed39c019d65a9404fa30ffe1e3a38b757d896d) --- fastvideo/hooks/layerwise_offload.py | 6 ++ fastvideo/layers/quantization/fp8_config.py | 18 ++++ fastvideo/layers/quantization/fp8_kernels.py | 83 +++++++++++++++++++ fastvideo/models/dits/minimax_h3.py | 28 +++++++ .../encoders/minimax_h3_checkpoint_nvfp4.py | 18 ++-- fastvideo/models/loader/component_loader.py | 11 +++ 6 files changed, 159 insertions(+), 5 deletions(-) create mode 100644 fastvideo/layers/quantization/fp8_kernels.py diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index f0ce1d52eb..fefb908b08 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -167,9 +167,15 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): return state_list = [] async_stream = torch.cuda.Stream() + # 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 + resident = int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0")) for name, submodule in model.named_children(): if isinstance(submodule, nn.ModuleList): for idx, module_entry in enumerate(submodule): + if idx < resident: + continue state = LayerwiseOffloadState(async_copy_stream=async_stream, device=device) state_list.append(state) hook_mgr = ModuleHookManager.get_from_or_default(module_entry) diff --git a/fastvideo/layers/quantization/fp8_config.py b/fastvideo/layers/quantization/fp8_config.py index cabc7cfda2..5b7e3f3cbf 100644 --- a/fastvideo/layers/quantization/fp8_config.py +++ b/fastvideo/layers/quantization/fp8_config.py @@ -3,6 +3,8 @@ Matches linear layers by suffix (``to_q/k/v/to_out``, ``ffn.fc_in/fc_out``). Supports per-tensor (default, fast) and per-channel (higher accuracy) granularity. +On sm89 the per-channel path runs the per-tensor GEMM plus a Triton scale epilogue +(``fp8_kernels``), since torch's rowwise-scaled kernel there is slower than bf16. Falls back to bf16 dequant on GPUs older than sm89. """ from __future__ import annotations @@ -77,8 +79,18 @@ def _quantize_tensorwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor return x_fp8, x_scale.view(1) +def _rowwise_scaled_mm_is_slow() -> bool: + if getattr(torch.version, "hip", None) or not torch.cuda.is_available(): + return False + from fastvideo.layers.quantization.fp8_kernels import rowwise_scaled_mm_is_slow + return rowwise_scaled_mm_is_slow() + + def _quantize_rowwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``.""" + if x_2d.is_cuda and _rowwise_scaled_mm_is_slow(): + from fastvideo.layers.quantization.fp8_kernels import quantize_rowwise_fp8 + return quantize_rowwise_fp8(x_2d) x_absmax = x_2d.abs().amax(dim=-1, keepdim=True).float() x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE) x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE) @@ -159,6 +171,12 @@ def apply( w_fp8 = layer._fp8_weight w_scale = layer._fp8_weight_scale + if self.granularity == "channel" and _rowwise_scaled_mm_is_slow(): + from fastvideo.layers.quantization.fp8_kernels import scaled_mm_token_channel + out = scaled_mm_token_channel(x_fp8, x_scale, w_fp8.t(), w_scale) + if bias is not None: + out = out + bias + return out.view(*original_shape[:-1], out_dim) scale_b = w_scale.view(1, -1) if self.granularity == "channel" else w_scale out = torch._scaled_mm( diff --git a/fastvideo/layers/quantization/fp8_kernels.py b/fastvideo/layers/quantization/fp8_kernels.py new file mode 100644 index 0000000000..5f8de6be83 --- /dev/null +++ b/fastvideo/layers/quantization/fp8_kernels.py @@ -0,0 +1,83 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Triton helpers for per-token x per-channel FP8 linears on GPUs whose rowwise ``_scaled_mm`` is slow. + +On sm89 (RTX 4090 / L40S / RTX 6000 Ada) ``torch._scaled_mm`` with rowwise scales runs at +~70 TFLOPS, below bf16 (~160), while the per-tensor kernel reaches 220-305 TFLOPS. The +per-token x per-channel result is recovered exactly by running the per-tensor kernel with unit +scales and applying ``out[i, j] *= sx[i] * sw[j]`` in one pass over the output, which costs +5-10% of the GEMM instead of 2-4x. +""" +from __future__ import annotations + +import functools + +import torch +import triton +import triton.language as tl + +FP8_MAX = 448.0 +FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0) + + +@functools.cache +def rowwise_scaled_mm_is_slow() -> bool: + """Ada (sm89) has no fast rowwise-scaled FP8 GEMM in torch; Hopper and Blackwell do.""" + return torch.cuda.is_available() and torch.cuda.get_device_capability() == (8, 9) + + +@triton.jit +def _quantize_rowwise_kernel(x_ptr, q_ptr, s_ptr, K, stride_x, stride_q, BLOCK_K: tl.constexpr): + row = tl.program_id(0) + x_row = x_ptr + row * stride_x + amax = tl.zeros((BLOCK_K, ), dtype=tl.float32) + for k in range(0, K, BLOCK_K): + cols = k + tl.arange(0, BLOCK_K) + amax = tl.maximum(amax, tl.abs(tl.load(x_row + cols, mask=cols < K, other=0.0).to(tl.float32))) + scale = tl.maximum(tl.max(amax, axis=0) / 448.0, 1.0 / (448.0 * 512.0)) + tl.store(s_ptr + row, scale) + inv = 1.0 / scale + for k in range(0, K, BLOCK_K): + cols = k + tl.arange(0, BLOCK_K) + v = tl.load(x_row + cols, mask=cols < K, other=0.0).to(tl.float32) * inv + v = tl.minimum(tl.maximum(v, -448.0), 448.0) + tl.store(q_ptr + row * stride_q + cols, v.to(tl.float8e4nv), mask=cols < K) + + +def quantize_rowwise_fp8(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Per-token FP8 quantization in one launch. Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``.""" + x_2d = x_2d.contiguous() + M, K = x_2d.shape + q = torch.empty((M, K), device=x_2d.device, dtype=torch.float8_e4m3fn) + s = torch.empty((M, 1), device=x_2d.device, dtype=torch.float32) + if M: + _quantize_rowwise_kernel[(M, )](x_2d, q, s, K, x_2d.stride(0), q.stride(0), BLOCK_K=1024, num_warps=4) + return q, s + + +@triton.jit +def _scale_rows_cols_kernel(o_ptr, sx_ptr, sw_ptr, M, N, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr): + rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) + cols = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N) + mask = (rows[:, None] < M) & (cols[None, :] < N) + ptrs = o_ptr + rows[:, None] * N + cols[None, :] + v = tl.load(ptrs, mask=mask, other=0.0).to(tl.float32) + sx = tl.load(sx_ptr + rows, mask=rows < M, other=0.0) + sw = tl.load(sw_ptr + cols, mask=cols < N, other=0.0) + tl.store(ptrs, (v * sx[:, None] * sw[None, :]).to(o_ptr.dtype.element_ty), mask=mask) + + +def scaled_mm_token_channel(x_fp8: torch.Tensor, x_scale: torch.Tensor, w_fp8_t: torch.Tensor, + w_scale: torch.Tensor) -> torch.Tensor: + """``(x_fp8 * x_scale) @ (w_fp8_t * w_scale)`` in bf16 via the fast per-tensor GEMM plus a scale epilogue. + + The unit-scale GEMM output is at most 448^2 * K, far inside bf16 range, and its relative + precision is that of any bf16 output, so the epilogue loses nothing against rowwise scaling. + """ + one = torch.ones((), device=x_fp8.device, dtype=torch.float32) + out = torch._scaled_mm(x_fp8, w_fp8_t, scale_a=one, scale_b=one, out_dtype=torch.bfloat16) + if isinstance(out, tuple): + out = out[0] + M, N = out.shape + grid = (triton.cdiv(M, 64), triton.cdiv(N, 128)) + _scale_rows_cols_kernel[grid](out, x_scale.reshape(-1), w_scale.reshape(-1), M, N, BLOCK_M=64, BLOCK_N=128) + return out diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index c092e4300f..ccda038983 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -140,8 +140,22 @@ def __init__( self.fuse_swiglu = fuse_swiglu self.use_mxfp8 = isinstance(self.fc_in.quant_method, MXFP8QuantizeMethod) and isinstance( self.fc_out.quant_method, MXFP8QuantizeMethod) + # Inference-only token chunking: the 2 * ffn_dim intermediate is ~5.3x the block input + # (4.5 GiB at 78k tokens), so chunks bound the activation peak on 24-32 GB GPUs. + self.chunk_tokens = int(os.environ.get("FASTVIDEO_H3_FFN_CHUNK_TOKENS", "0")) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + tokens = hidden_states.shape[-2] if hidden_states.dim() > 1 else 0 + if (self.chunk_tokens and tokens > self.chunk_tokens and not torch.is_grad_enabled() + and not torch.compiler.is_compiling() and hidden_states.numel() == tokens * hidden_states.shape[-1]): + out = torch.empty_like(hidden_states) + for start in range(0, tokens, self.chunk_tokens): + rows = slice(start, start + self.chunk_tokens) + out[..., rows, :] = self._forward(hidden_states[..., rows, :].contiguous()) + return out + return self._forward(hidden_states) + + def _forward(self, hidden_states: torch.Tensor) -> torch.Tensor: if self.use_mxfp8: from fastvideo.layers.mxfp8linear import mxfp8_swiglu_feed_forward @@ -977,6 +991,16 @@ def _rotary_for(self, position_ids: torch.Tensor, dtype: torch.dtype) -> tuple[t self._rope_cache = (position_ids, dtype, value) return value + def attach_step_splice(self, late: nn.Module, from_step: int) -> None: + """Hand denoising steps ``from_step`` onward to ``late`` (same architecture, other weights). + + Early DMD steps fix layout and object count, late ones texture and detail, so two checkpoints + can split the trajectory. ``late`` is kept out of this module's children: it is placed, offloaded + and checkpointed on its own. + """ + object.__setattr__(self, "_splice_late", late) + self._splice_from_step = int(from_step) + def enable_adaln_host_cache(self, table_path: str | None = None) -> None: """Move every block's AdaLN projection to pinned host memory behind a per-timestep cache. @@ -1055,6 +1079,10 @@ def forward( text_indices: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: """Predict video and audio velocities from one caller-defined packed layout.""" + late = self.__dict__.get("_splice_late") + if late is not None and get_forward_context().current_timestep >= self._splice_from_step: + return late(hidden_states, audio_hidden_states, encoder_hidden_states, timestep, timestep_indices, + token_tags, position_ids, video_indices, audio_indices, text_indices) if position_ids.ndim != 2 or position_ids.shape[-1] != 3: raise ValueError(f"position_ids must have shape (seq_len, 3), got {tuple(position_ids.shape)}.") sequence_length = position_ids.shape[0] diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 7082e26256..5479f797f1 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -54,6 +54,9 @@ _require_flashinfer, ) from fastvideo.models.utils import set_weight_attrs +from fastvideo.logger import init_logger + +logger = init_logger(__name__) NVFP4_GROUP_SIZE = 16 NVFP4_SCALE_LAYOUT = "128x4" @@ -299,15 +302,20 @@ def validate_runtime(self, device: torch.device) -> None: raise RuntimeError(f"MiniMax-H3 serialized NVFP4 requires a CUDA device; got {device.type!r}") capability = torch.cuda.get_device_capability(device) capability_number = capability[0] * 10 + capability[1] + if get_tp_world_size() > 1: + raise NotImplementedError("MiniMax-H3 serialized NVFP4 supports a single GPU: packed FP4 columns and " + "128x4 swizzled scale rows cannot be narrowed per tensor-parallel rank") if capability_number < self.get_min_capability(): - raise RuntimeError("MiniMax-H3 serialized NVFP4 requires GPU capability " - f"sm{self.get_min_capability()} or newer, got sm{capability_number}") + if capability_number < 80: + raise RuntimeError("MiniMax-H3 serialized NVFP4 needs bf16 compute (sm80+) for its de-quantized " + f"fallback, got sm{capability_number}") + logger.warning( + "MiniMax-H3 serialized NVFP4 on sm%d: no FP4 GEMM, each linear de-quantizes its weight to bf16 " + "per call", capability_number) + return if capability[0] not in (10, 12): raise RuntimeError("MiniMax-H3 serialized NVFP4 runs FlashInfer's Blackwell FP4 GEMM; " f"got unsupported sm{capability_number}") - if get_tp_world_size() > 1: - raise NotImplementedError("MiniMax-H3 serialized NVFP4 supports a single GPU: packed FP4 columns and " - "128x4 swizzled scale rows cannot be narrowed per tensor-parallel rank") sf_layout, _, _ = _require_flashinfer() if not hasattr(sf_layout, "layout_128x4"): raise RuntimeError("The installed flashinfer has no SfLayout.layout_128x4; MiniMax-H3 serialized NVFP4 " diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 05312e92a2..3180062809 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -1210,6 +1210,17 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): "nn.ModuleList structure. Skipping layerwise offload.", cls_name) if layerwise_load: model = model.to(get_local_torch_device()) + # FASTVIDEO_H3_SPLICE_TRANSFORMER=: a second checkpoint of the same architecture + # runs denoising steps FASTVIDEO_H3_SPLICE_FROM_STEP (default 4) onward. + splice_path = os.environ.pop("FASTVIDEO_H3_SPLICE_TRANSFORMER", None) + if splice_path and hasattr(model, "attach_step_splice"): + try: + late = self.load(splice_path, fastvideo_args) + finally: + os.environ["FASTVIDEO_H3_SPLICE_TRANSFORMER"] = splice_path + from_step = int(os.environ.get("FASTVIDEO_H3_SPLICE_FROM_STEP", "4")) + model.attach_step_splice(late, from_step) + logger.info("Step splice: steps >= %d run the transformer from %s", from_step, splice_path) return model From 4a72920130db7b2e07718491d8bd424d39f77067 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:34:58 -0700 Subject: [PATCH 010/107] [feat]: H3 headline benchmark (480p 5 s, 768p 10 s) for GPU clusters and Modal bench_headline.py times the release protocol (two fixed prompts, one warmup, two timed runs each, generate_video wall time) from a checkpoint's fastvideo_inference.json contract, with optional W&B logging. headline_app.py runs it on 1/4/8 RTX PRO 6000 Blackwell GPUs on Modal from a FastVideo HF repo. --- .../minimax_h3_pro6000/bench_headline.py | 99 +++++++++++++++++++ .../minimax_h3_pro6000/headline_app.py | 80 +++++++++++++++ .../minimax_h3_pro6000/headline_prompts.json | 4 + 3 files changed, 183 insertions(+) create mode 100644 scripts/benchmarks/minimax_h3_pro6000/bench_headline.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_app.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_prompts.json diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py new file mode 100644 index 0000000000..e8f92ecd34 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -0,0 +1,99 @@ +"""Headline e2e benchmark for the FastH3 local release: 480p 5 s and 768p 10 s. + +usage: python bench_headline.py [--settings 480p5s,768p10s] + +Protocol (RELEASE_PLAN §5): prompts latency-ceramics-005 + latency-harbor-005, one untimed warmup per +setting, then each prompt timed twice; e2e = generate_video wall time (encode + DiT + decode + mp4 write). +Per-stage times come from FASTVIDEO_STAGE_LOGGING. Clips and results.json go to //, +and one W&B run per invocation (group headline-) when WANDB_PROJECT is set. +""" +import argparse +import json +import os +import statistics +import time + +SETTINGS = {"480p5s": (832, 480, 124), "768p10s": (1344, 768, 243)} # 17n+5 frames at 24 fps +PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("run_name") + ap.add_argument("model_dir") + ap.add_argument("num_gpus", type=int) + ap.add_argument("nvfp4_profile") + ap.add_argument("--settings", default="480p5s,768p10s") + ap.add_argument("--prompts", default=os.path.join(os.path.dirname(os.path.abspath(__file__)), "headline_prompts.json")) + ap.add_argument("--out", default=os.environ.get("HEADLINE_OUT", "headline")) + ap.add_argument("--timed", type=int, default=2) + ap.add_argument("--sparsity", type=float, default=None, help="default: the checkpoint contract's vsa_sparsity") + a = ap.parse_args() + + contract = json.load(open(os.path.join(a.model_dir, "fastvideo_inference.json"))) + steps = contract["dmd_denoising_steps"] + sparsity = a.sparsity if a.sparsity is not None else float(contract.get("vsa_sparsity", 0.8)) + texts = json.load(open(a.prompts)) + out_dir = os.path.join(a.out, a.run_name) + os.makedirs(out_dir, exist_ok=True) + + engine = {"num_gpus": a.num_gpus, "use_fsdp_inference": False, + "parallelism": {"tp_size": 1, "sp_size": a.num_gpus}, + "offload": {"dit": False, "dit_layerwise": False, "text_encoder": False, "vae": False, + "pin_cpu_memory": False, "lazy_module_load": False}, + "compile": {"enabled": False, "vae_enabled": os.environ.get("HEADLINE_VAE_COMPILE", "1") == "1"}} + if a.nvfp4_profile != "none": + engine["quantization"] = {"transformer_quant": "NVFP4", "layer_profile": a.nvfp4_profile} + config = {"model_path": a.model_dir, "engine": engine, + "pipeline": {"experimental": {"attention_backend": "VIDEO_SPARSE_ATTN_H3", "VSA_sparsity": sparsity, + "VSA_tile_size": 64, "h3_sequential_load": False, + "inference_torch_compile": False}}} + env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA"))} + run = None + if os.environ.get("WANDB_PROJECT"): + import wandb + run = wandb.init(project=os.environ["WANDB_PROJECT"], entity=os.environ.get("WANDB_ENTITY"), + name=a.run_name, group=os.environ.get("HEADLINE_GROUP", "headline"), job_type="benchmark", + dir=out_dir, config={"model_dir": a.model_dir, "num_gpus": a.num_gpus, + "nvfp4_profile": a.nvfp4_profile, "dmd_steps": steps, + "vsa_sparsity": sparsity, "settings": a.settings, "engine": engine, + "env": env, "device": os.environ.get("HEADLINE_DEVICE", "")}) + + from fastvideo import VideoGenerator + t0 = time.perf_counter() + generator = VideoGenerator.from_config(config) + results = {"run_name": a.run_name, "model_dir": a.model_dir, "num_gpus": a.num_gpus, + "nvfp4_profile": a.nvfp4_profile, "load_s": round(time.perf_counter() - t0, 1), "env": env, + "settings": {}} + for name in a.settings.split(","): + width, height, frames = SETTINGS[name] + runs = [] + plan = [(PROMPT_IDS[0], True)] + [(pid, False) for _ in range(a.timed) for pid in PROMPT_IDS] + for i, (pid, warmup) in enumerate(plan): + path = os.path.join(out_dir, f"{name}_{pid}_{'warmup' if warmup else i}.mp4") + t = time.perf_counter() + generator.generate_video(prompt=texts[pid], height=height, width=width, num_frames=frames, fps=24, + guidance_scale=1.0, num_inference_steps=len(steps) + 1, seed=1234, + output_path=path, save_video=True) + wall = round(time.perf_counter() - t, 2) + runs.append({"prompt": pid, "warmup": warmup, "e2e_s": wall, "path": path}) + print("RUN", name, pid, "warmup" if warmup else "timed", wall, flush=True) + if run is not None and not warmup: + import wandb + run.log({f"{name}/e2e_s": wall, f"{name}/{pid}": wandb.Video(path, fps=24, format="mp4")}) + timed = [r["e2e_s"] for r in runs if not r["warmup"]] + results["settings"][name] = {"width": width, "height": height, "frames": frames, + "e2e_median_s": statistics.median(timed), "e2e_min_s": min(timed), "runs": runs} + print("SETTING", name, json.dumps(results["settings"][name]), flush=True) + if run is not None: + run.summary[f"{name}_e2e_median_s"] = statistics.median(timed) + json.dump(results, open(os.path.join(out_dir, "results.json"), "w"), indent=1) + generator.shutdown() + if run is not None: + run.summary["load_s"] = results["load_s"] + run.finish() + print("HEADLINE_DONE", a.run_name, flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py new file mode 100644 index 0000000000..1c20e677db --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py @@ -0,0 +1,80 @@ +"""Headline e2e numbers (480p 5 s, 768p 10 s) for a FastH3 HF repo on 1 / 4 / 8 RTX PRO 6000 Blackwell GPUs. + +Reuses the image and kernel wheel of ``app.py`` (run ``modal run app.py --step build`` once per volume). + + MODAL_PROFILE=aryan5v modal run --detach headline_app.py --repo FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 \\ + --profile h3_dit_ffn --gpus 1,4,8 + +Results land in the ``h3-pro6000-weights`` volume under ``outputs/headline//`` (``results.json`` + clips). +""" +import json +import os +import pathlib +import subprocess + +import modal + +from app import FAST_ENV, _install_kernel, _sh, image, volume + +HERE = pathlib.Path(__file__).resolve().parent +headline_image = (image.add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") + .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) +app = modal.App("h3-pro6000-headline", image=headline_image) +SECRETS = [modal.Secret.from_name("hf-fastvideo")] + + +@app.function(cpu=8, memory=32768, timeout=3600, volumes={"/vol": volume}, secrets=SECRETS) +def fetch(repo: str) -> str: + from huggingface_hub import snapshot_download + local = f"/vol/models/{repo.split('/')[-1]}" + snapshot_download(repo, local_dir=local, token=os.environ["HF_TOKEN"], max_workers=16) + volume.commit() + return _sh(f"du -sh {local}/*") + + +def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dict: + _install_kernel() + model = f"/vol/models/{repo.split('/')[-1]}" + run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}" + env = dict(os.environ, **FAST_ENV, **(extra_env or {}), HEADLINE_OUT="/vol/outputs/headline", + HEADLINE_DEVICE=f"{gpus}x RTX PRO 6000", PYTHONPATH="/src/fastvideo") + proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, + "--prompts", "/root/headline_prompts.json"], env=env, capture_output=True, text=True, + cwd="/root") + volume.commit() + tail = (proc.stdout + proc.stderr)[-6000:] + result_path = pathlib.Path("/vol/outputs/headline") / run_name / "results.json" + results = json.loads(result_path.read_text()) if result_path.exists() else {} + return {"run": run_name, "returncode": proc.returncode, "results": results, + "log_tail": tail if proc.returncode else tail[-1500:]} + + +@app.function(gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=2 * 3600, volumes={"/vol": volume}) +def headline1(repo: str, profile: str, extra_env: dict | None = None) -> dict: + return _headline(repo, 1, profile, extra_env) + + +@app.function(gpu="RTX-PRO-6000:4", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) +def headline4(repo: str, profile: str, extra_env: dict | None = None) -> dict: + return _headline(repo, 4, profile, extra_env) + + +@app.function(gpu="RTX-PRO-6000:8", memory=262144, cpu=32, timeout=2 * 3600, volumes={"/vol": volume}) +def headline8(repo: str, profile: str, extra_env: dict | None = None) -> dict: + return _headline(repo, 8, profile, extra_env) + + +@app.local_entrypoint() +def main(repo: str, profile: str = "h3_dit_ffn", gpus: str = "1,4,8", skip_fetch: bool = False): + if not skip_fetch: + print(fetch.remote(repo)) + fns = {"1": headline1, "4": headline4, "8": headline8} + calls = [fns[g].spawn(repo, profile) for g in gpus.split(",")] + out = HERE / "headline_results" + out.mkdir(exist_ok=True) + for call in calls: + res = call.get() + print(json.dumps({k: v for k, v in res.items() if k != "log_tail"}, indent=1)[:3000]) + if res["returncode"]: + print(res["log_tail"]) + (out / f"{res['run']}.json").write_text(json.dumps(res, indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_prompts.json b/scripts/benchmarks/minimax_h3_pro6000/headline_prompts.json new file mode 100644 index 0000000000..553a5a7d67 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_prompts.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 f0688b0056eb01edd889be275dba35a0cc2268e7 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:37:42 -0700 Subject: [PATCH 011/107] [bugfix]: ship app.py into the headline Modal image --- scripts/benchmarks/minimax_h3_pro6000/headline_app.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py index 1c20e677db..cc25bdba6d 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py @@ -17,7 +17,8 @@ from app import FAST_ENV, _install_kernel, _sh, image, volume HERE = pathlib.Path(__file__).resolve().parent -headline_image = (image.add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") +headline_image = (image.add_local_file(HERE / "app.py", "/root/app.py") + .add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) app = modal.App("h3-pro6000-headline", image=headline_image) SECRETS = [modal.Secret.from_name("hf-fastvideo")] From 0a23563bff05133a8397ed91000f859b966b19c5 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:39:54 -0700 Subject: [PATCH 012/107] [bugfix]: ship app.py into the headline Modal image as local python source --- scripts/benchmarks/minimax_h3_pro6000/headline_app.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py index cc25bdba6d..4458da76e2 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py @@ -17,7 +17,7 @@ from app import FAST_ENV, _install_kernel, _sh, image, volume HERE = pathlib.Path(__file__).resolve().parent -headline_image = (image.add_local_file(HERE / "app.py", "/root/app.py") +headline_image = (image.add_local_python_source("app") .add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) app = modal.App("h3-pro6000-headline", image=headline_image) From a87344c55f61b14f4efe1a52631b2aca3526656e Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:53:00 -0700 Subject: [PATCH 013/107] [bugfix]: run the headline benchmark from app.py (one Modal app, no cross-module import) --- scripts/benchmarks/minimax_h3_pro6000/app.py | 65 +++++++++++++++ .../minimax_h3_pro6000/headline_app.py | 81 ------------------- ...0x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json | 6 ++ ...0x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json | 6 ++ 4 files changed, 77 insertions(+), 81 deletions(-) delete mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_app.py create mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json create mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index 9b2058fade..4352d9e812 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -536,3 +536,68 @@ def main(step: str = "all", ladder: str = "base"): print("FAILED", v[0], repr(result)[:3000]) else: print("RESULT", json.dumps(result)) + + +# --- Headline benchmark (480p 5 s, 768p 10 s) on 1 / 4 / 8 GPUs from a FastVideo HF repo. --- +# MODAL_PROFILE=aryan5v modal run --detach app.py::headline --repo FastVideo/ --profile h3_dit_ffn --gpus 1,4,8 +HERE = pathlib.Path(__file__).resolve().parent +headline_image = (image.add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") + .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) +SECRETS = [modal.Secret.from_name("hf-fastvideo")] + + +@app.function(cpu=8, memory=32768, timeout=3600, volumes={"/vol": volume}, secrets=SECRETS, image=headline_image) +def headline_fetch(repo: str) -> str: + from huggingface_hub import snapshot_download + local = f"/vol/models/{repo.split('/')[-1]}" + snapshot_download(repo, local_dir=local, token=os.environ["HF_TOKEN"], max_workers=16) + volume.commit() + return _sh(f"du -sh {local}/*") + + +def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dict: + _install_kernel() + model = f"/vol/models/{repo.split('/')[-1]}" + run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}" + env = dict(os.environ, **FAST_ENV, **(extra_env or {}), HEADLINE_OUT="/vol/outputs/headline", + HEADLINE_DEVICE=f"{gpus}x RTX PRO 6000", PYTHONPATH="/src/fastvideo") + proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, + "--prompts", "/root/headline_prompts.json"], env=env, capture_output=True, text=True, + cwd="/root") + volume.commit() + tail = (proc.stdout + proc.stderr)[-6000:] + result_path = pathlib.Path("/vol/outputs/headline") / run_name / "results.json" + results = json.loads(result_path.read_text()) if result_path.exists() else {} + return {"run": run_name, "returncode": proc.returncode, "results": results, + "log_tail": tail if proc.returncode else tail[-1500:]} + + +@app.function(image=headline_image, gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=2 * 3600, volumes={"/vol": volume}) +def headline1(repo: str, profile: str, extra_env: dict | None = None) -> dict: + return _headline(repo, 1, profile, extra_env) + + +@app.function(image=headline_image, gpu="RTX-PRO-6000:4", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) +def headline4(repo: str, profile: str, extra_env: dict | None = None) -> dict: + return _headline(repo, 4, profile, extra_env) + + +@app.function(image=headline_image, gpu="RTX-PRO-6000:8", memory=262144, cpu=32, timeout=2 * 3600, volumes={"/vol": volume}) +def headline8(repo: str, profile: str, extra_env: dict | None = None) -> dict: + return _headline(repo, 8, profile, extra_env) + + +@app.local_entrypoint() +def headline(repo: str, profile: str = "h3_dit_ffn", gpus: str = "1,4,8", skip_fetch: bool = False): + if not skip_fetch: + print(headline_fetch.remote(repo)) + fns = {"1": headline1, "4": headline4, "8": headline8} + calls = [fns[g].spawn(repo, profile) for g in gpus.split(",")] + out = HERE / "headline_results" + out.mkdir(exist_ok=True) + for call in calls: + res = call.get() + print(json.dumps({k: v for k, v in res.items() if k != "log_tail"}, indent=1)[:3000]) + if res["returncode"]: + print(res["log_tail"]) + (out / f"{res['run']}.json").write_text(json.dumps(res, indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py b/scripts/benchmarks/minimax_h3_pro6000/headline_app.py deleted file mode 100644 index 4458da76e2..0000000000 --- a/scripts/benchmarks/minimax_h3_pro6000/headline_app.py +++ /dev/null @@ -1,81 +0,0 @@ -"""Headline e2e numbers (480p 5 s, 768p 10 s) for a FastH3 HF repo on 1 / 4 / 8 RTX PRO 6000 Blackwell GPUs. - -Reuses the image and kernel wheel of ``app.py`` (run ``modal run app.py --step build`` once per volume). - - MODAL_PROFILE=aryan5v modal run --detach headline_app.py --repo FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 \\ - --profile h3_dit_ffn --gpus 1,4,8 - -Results land in the ``h3-pro6000-weights`` volume under ``outputs/headline//`` (``results.json`` + clips). -""" -import json -import os -import pathlib -import subprocess - -import modal - -from app import FAST_ENV, _install_kernel, _sh, image, volume - -HERE = pathlib.Path(__file__).resolve().parent -headline_image = (image.add_local_python_source("app") - .add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") - .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) -app = modal.App("h3-pro6000-headline", image=headline_image) -SECRETS = [modal.Secret.from_name("hf-fastvideo")] - - -@app.function(cpu=8, memory=32768, timeout=3600, volumes={"/vol": volume}, secrets=SECRETS) -def fetch(repo: str) -> str: - from huggingface_hub import snapshot_download - local = f"/vol/models/{repo.split('/')[-1]}" - snapshot_download(repo, local_dir=local, token=os.environ["HF_TOKEN"], max_workers=16) - volume.commit() - return _sh(f"du -sh {local}/*") - - -def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dict: - _install_kernel() - model = f"/vol/models/{repo.split('/')[-1]}" - run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}" - env = dict(os.environ, **FAST_ENV, **(extra_env or {}), HEADLINE_OUT="/vol/outputs/headline", - HEADLINE_DEVICE=f"{gpus}x RTX PRO 6000", PYTHONPATH="/src/fastvideo") - proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, - "--prompts", "/root/headline_prompts.json"], env=env, capture_output=True, text=True, - cwd="/root") - volume.commit() - tail = (proc.stdout + proc.stderr)[-6000:] - result_path = pathlib.Path("/vol/outputs/headline") / run_name / "results.json" - results = json.loads(result_path.read_text()) if result_path.exists() else {} - return {"run": run_name, "returncode": proc.returncode, "results": results, - "log_tail": tail if proc.returncode else tail[-1500:]} - - -@app.function(gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=2 * 3600, volumes={"/vol": volume}) -def headline1(repo: str, profile: str, extra_env: dict | None = None) -> dict: - return _headline(repo, 1, profile, extra_env) - - -@app.function(gpu="RTX-PRO-6000:4", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) -def headline4(repo: str, profile: str, extra_env: dict | None = None) -> dict: - return _headline(repo, 4, profile, extra_env) - - -@app.function(gpu="RTX-PRO-6000:8", memory=262144, cpu=32, timeout=2 * 3600, volumes={"/vol": volume}) -def headline8(repo: str, profile: str, extra_env: dict | None = None) -> dict: - return _headline(repo, 8, profile, extra_env) - - -@app.local_entrypoint() -def main(repo: str, profile: str = "h3_dit_ffn", gpus: str = "1,4,8", skip_fetch: bool = False): - if not skip_fetch: - print(fetch.remote(repo)) - fns = {"1": headline1, "4": headline4, "8": headline8} - calls = [fns[g].spawn(repo, profile) for g in gpus.split(",")] - out = HERE / "headline_results" - out.mkdir(exist_ok=True) - for call in calls: - res = call.get() - print(json.dumps({k: v for k, v in res.items() if k != "log_tail"}, indent=1)[:3000]) - if res["returncode"]: - print(res["log_tail"]) - (out / f"{res['run']}.json").write_text(json.dumps(res, indent=1)) diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json b/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json new file mode 100644 index 0000000000..f827ac0bd6 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json @@ -0,0 +1,6 @@ +{ + "run": "pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300", + "returncode": -3, + "results": {}, + "log_tail": "INFO 10-03 17:41:41.508 [__init__.py:47] CUDA is available\nINFO 10-03 17:41:43.971 [utils.py:691] Diffusers version: 0.36.0.dev0\nINFO 10-03 17:41:43.973 [multiproc_executor.py:107] Use master port: 52996\nINFO 10-03 17:41:55.492 [__init__.py:47] CUDA is available\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] WorkerMultiprocProc failed to start.\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] Traceback (most recent call last):\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] self.worker.init_device()\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] ^^\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:41:58.093 [multiproc_executor.py:682] Worker 0 initialization failed: name 'os' is not defined\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] Worker 0 hit an exception: Traceback (most recent call last):\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] self.worker.init_device()\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] ^^\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] \n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/multiprocessing/resource_tracker.py:254: UserWarning: resource_tracker: There appear to be 6 leaked semaphore objects to clean up at shutdown\n warnings.warn('resource_tracker: There appear to be %d '\n" +} \ No newline at end of file diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json b/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json new file mode 100644 index 0000000000..19f67ac323 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json @@ -0,0 +1,6 @@ +{ + "run": "pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300", + "returncode": -3, + "results": {}, + "log_tail": "INFO 10-03 17:45:31.047 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:33.414 [utils.py:691] Diffusers version: 0.36.0.dev0\nINFO 10-03 17:45:33.416 [multiproc_executor.py:107] Use master port: 51677\nINFO 10-03 17:45:59.703 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:59.746 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:59.749 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:59.801 [__init__.py:47] CUDA is available\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] WorkerMultiprocProc failed to start.\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] Traceback (most recent call last):\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] self.worker.init_device()\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] ^^\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:46:04.272 [multiproc_executor.py:682] Worker 3 initialization failed: name 'os' is not defined\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] Worker 3 hit an exception: Traceback (most recent call last):\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] self.worker.init_device()\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] ^^\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] \n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/multiprocessing/resource_tracker.py:254: UserWarning: resource_tracker: There appear to be 6 leaked semaphore objects to clean up at shutdown\n warnings.warn('resource_tracker: There appear to be %d '\n" +} \ No newline at end of file From edbe18bbfdffc7c6e8c2fdffe3c71184eac76b15 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 10:57:09 -0700 Subject: [PATCH 014/107] [bugfix]: import os in gpu_worker; apply pre-commit formatting; ignore local headline results --- .gitignore | 1 + fastvideo/layers/quantization/nvfp4_config.py | 3 +-- fastvideo/pipelines/stages/base.py | 6 ++++-- fastvideo/worker/gpu_worker.py | 12 ++++++++---- .../benchmarks/minimax_h3_pro6000/bench_headline.py | 2 ++ .../pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json | 6 ------ .../pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json | 6 ------ 7 files changed, 16 insertions(+), 20 deletions(-) delete mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json delete mode 100644 scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json diff --git a/.gitignore b/.gitignore index 13c5fbc871..88035e8c2b 100644 --- a/.gitignore +++ b/.gitignore @@ -142,3 +142,4 @@ fastvideo/tests/ssim/reference_videos/** *.nvimlog .nvimlog .python-version +scripts/benchmarks/minimax_h3_pro6000/headline_results/ diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 403af3b5ac..7678ae1db1 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -418,8 +418,7 @@ def _static_activation_global_sf(self) -> torch.Tensor | None: keys = [prefix] + ([f"b{match.group(1)}.{match.group(2)}"] if match else []) amax = next((table[k] for k in keys if k in table), None) if amax is not None: - self._static_sf = torch.tensor((448.0 * 6.0) / max(amax, 1e-12), dtype=torch.float32, - device="cuda") + self._static_sf = torch.tensor((448.0 * 6.0) / max(amax, 1e-12), dtype=torch.float32, device="cuda") return self._static_sf def _dynamic_activation_scale(self) -> bool: diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index c87d946f65..aa086a7d8f 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -191,8 +191,10 @@ def _execute( if torch.cuda.is_available(): gib = 1024**3 logger.info("[%s] Memory peak_allocated=%.2f GiB reserved=%.2f GiB resident_after=%.2f GiB", - stage_name, torch.cuda.max_memory_allocated() / gib, - torch.cuda.memory_reserved() / gib, torch.cuda.memory_allocated() / gib) + stage_name, + torch.cuda.max_memory_allocated() / gib, + torch.cuda.memory_reserved() / gib, + torch.cuda.memory_allocated() / gib) torch.cuda.reset_peak_memory_stats() batch.logging_info.add_stage_execution_time(stage_key, execution_time) batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name) diff --git a/fastvideo/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index 4fdcae2c5e..beb2af9385 100644 --- a/fastvideo/worker/gpu_worker.py +++ b/fastvideo/worker/gpu_worker.py @@ -1,4 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 +import os from typing import Any, cast import torch @@ -21,7 +22,6 @@ def _log_cuda_device_uuid(rank: int, device: torch.device) -> None: logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False) - def _log_pipeline_memory(pipeline) -> None: """Debug (FASTVIDEO_MEMORY_REPORT=1): bytes held per pipeline component, by device and dtype, plus the largest tensors, so the resident footprint can be attributed before choosing offload placements.""" @@ -42,14 +42,18 @@ def _log_pipeline_memory(pipeline) -> None: largest.append((nbytes, tname, key)) largest.sort(reverse=True) total = sum(by_kind.values()) - logger.info("MEMREPORT %s total=%.2f GiB %s", name, total / gib, - {k: round(v / gib, 2) for k, v in sorted(by_kind.items(), key=lambda kv: -kv[1])}) + logger.info("MEMREPORT %s total=%.2f GiB %s", name, total / gib, { + k: round(v / gib, 2) + for k, v in sorted(by_kind.items(), key=lambda kv: -kv[1]) + }) for nbytes, tname, key in largest[:8]: logger.info("MEMREPORT %s %.3f GiB %s %s", name, nbytes / gib, key, tname) if torch.cuda.is_available(): - logger.info("MEMREPORT cuda allocated=%.2f GiB reserved=%.2f GiB", torch.cuda.memory_allocated() / gib, + logger.info("MEMREPORT cuda allocated=%.2f GiB reserved=%.2f GiB", + torch.cuda.memory_allocated() / gib, torch.cuda.memory_reserved() / gib) + class Worker: def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str): diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py index e8f92ecd34..85d3c6de87 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -42,6 +42,8 @@ def main(): "offload": {"dit": False, "dit_layerwise": False, "text_encoder": False, "vae": False, "pin_cpu_memory": False, "lazy_module_load": False}, "compile": {"enabled": False, "vae_enabled": os.environ.get("HEADLINE_VAE_COMPILE", "1") == "1"}} + if os.environ.get("HEADLINE_BACKEND"): + engine["execution_backend"] = os.environ["HEADLINE_BACKEND"] # "ray" for multi-node if a.nvfp4_profile != "none": engine["quantization"] = {"transformer_quant": "NVFP4", "layer_profile": a.nvfp4_profile} config = {"model_path": a.model_dir, "engine": engine, diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json b/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json deleted file mode 100644 index f827ac0bd6..0000000000 --- a/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300.json +++ /dev/null @@ -1,6 +0,0 @@ -{ - "run": "pro6000x1-FastH3-Pruned-8Step-NVFP4-ckpt300", - "returncode": -3, - "results": {}, - "log_tail": "INFO 10-03 17:41:41.508 [__init__.py:47] CUDA is available\nINFO 10-03 17:41:43.971 [utils.py:691] Diffusers version: 0.36.0.dev0\nINFO 10-03 17:41:43.973 [multiproc_executor.py:107] Use master port: 52996\nINFO 10-03 17:41:55.492 [__init__.py:47] CUDA is available\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] WorkerMultiprocProc failed to start.\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] Traceback (most recent call last):\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] self.worker.init_device()\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] ^^\nERROR 10-03 17:41:58.091 [multiproc_executor.py:622] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:41:58.093 [multiproc_executor.py:682] Worker 0 initialization failed: name 'os' is not defined\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] Worker 0 hit an exception: Traceback (most recent call last):\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] self.worker.init_device()\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] ^^\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:41:58.093 [multiproc_executor.py:643] \n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/multiprocessing/resource_tracker.py:254: UserWarning: resource_tracker: There appear to be 6 leaked semaphore objects to clean up at shutdown\n warnings.warn('resource_tracker: There appear to be %d '\n" -} \ No newline at end of file diff --git a/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json b/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json deleted file mode 100644 index 19f67ac323..0000000000 --- a/scripts/benchmarks/minimax_h3_pro6000/headline_results/pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300.json +++ /dev/null @@ -1,6 +0,0 @@ -{ - "run": "pro6000x4-FastH3-Pruned-8Step-NVFP4-ckpt300", - "returncode": -3, - "results": {}, - "log_tail": "INFO 10-03 17:45:31.047 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:33.414 [utils.py:691] Diffusers version: 0.36.0.dev0\nINFO 10-03 17:45:33.416 [multiproc_executor.py:107] Use master port: 51677\nINFO 10-03 17:45:59.703 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:59.746 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:59.749 [__init__.py:47] CUDA is available\nINFO 10-03 17:45:59.801 [__init__.py:47] CUDA is available\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] WorkerMultiprocProc failed to start.\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] Traceback (most recent call last):\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] self.worker.init_device()\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] ^^\nERROR 10-03 17:46:04.269 [multiproc_executor.py:622] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:46:04.272 [multiproc_executor.py:682] Worker 3 initialization failed: name 'os' is not defined\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] Worker 3 hit an exception: Traceback (most recent call last):\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 583, in worker_main\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] worker = WorkerMultiprocProc(*args, **kwargs)\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/multiproc_executor.py\", line 504, in __init__\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] self.worker.init_device()\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] File \"/src/fastvideo/fastvideo/worker/gpu_worker.py\", line 104, in init_device\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] cap_gib = os.environ.get(\"FASTVIDEO_CUDA_MEMORY_CAP_GIB\")\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] ^^\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] NameError: name 'os' is not defined. Did you forget to import 'os'\nERROR 10-03 17:46:04.272 [multiproc_executor.py:643] \n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/site-packages/huggingface_hub/constants.py:302: FutureWarning: The `HF_HUB_ENABLE_HF_TRANSFER` environment variable is deprecated as 'hf_transfer' is not used anymore. Please use `HF_XET_HIGH_PERFORMANCE` instead to enable high performance transfer with Xet. Visit https://huggingface.co/docs/huggingface_hub/package_reference/environment_variables#hfxethighperformance for more details.\n warnings.warn(\n/usr/local/lib/python3.12/multiprocessing/resource_tracker.py:254: UserWarning: resource_tracker: There appear to be 6 leaked semaphore objects to clean up at shutdown\n warnings.warn('resource_tracker: There appear to be %d '\n" -} \ No newline at end of file From dd10fc7bc4ddb32ac32811d23f468c3ca806978d Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:08:29 -0700 Subject: [PATCH 015/107] [bugfix]: H3 FP4 attention: new _build_block_mask signature; share Q/K/V quantization only at unit scale Upstream's _build_block_mask now takes per-region video tile spans and sparsities; the FP4 VSA paths still passed the old five arguments and raised TypeError on the first block. The shared Q/K/V quantization assumed every NVFP4 layer used the unit activation scale; with calibrated (export or env) or dynamic scales each projection now quantizes its own input. Found by Greptile review on #45. --- fastvideo/layers/quantization/nvfp4_config.py | 9 +++++++++ fastvideo/models/dits/minimax_h3_vsa_fp4.py | 16 ++++++++++------ 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 7678ae1db1..ded48937fa 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -432,6 +432,15 @@ def _dynamic_activation_scale(self) -> bool: self._dynamic_act_cached = cached return cached + def uses_unit_activation_scale(self, layer: torch.nn.Module) -> bool: + """Whether ``apply`` quantizes this layer's input with the unit global scale. + + False when a calibrated scale (env table or the export's ``_nvfp4_input_global_sf``) or a dynamic + per-call scale applies; such inputs cannot share one pre-quantized copy across layers. + """ + return (self._static_activation_global_sf() is None and getattr(layer, H3_NVFP4_DIT_INPUT_SF_NAME, None) is None + and not self._dynamic_activation_scale()) + def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: SfLayout, _, _ = _require_flashinfer() x = _coerce_fp4_input_dtype(x) diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index 1291c47e32..5429f76a3f 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -123,15 +123,17 @@ 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. + """Run projections of one input, quantizing it once when all are NVFP4 with the unit activation scale. - NVFP4 activations use a unit global scale for every layer, so one - quantized copy is exactly what each layer would have produced. + 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. """ from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod methods = [linear.quant_method for linear in linears] - if not all(type(m) is NVFP4QuantizeMethod and m.wants_prequantized_input() for m in methods): + 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)): 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)] @@ -168,7 +170,8 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ k_pooled = _pool_tiles(key[:, :logical], vbs, layout.tile) 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) q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) with STAGES.span("fp4_attention"): out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) @@ -334,7 +337,8 @@ def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tup scores = torch.matmul(_pool_tiles(q_t[:, :logical], vbs, layout.tile), _pool_tiles(k_t[:, :logical], vbs, layout.tile).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) q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) with STAGES.span("fp4_attention"): out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad) From a6faddd3d156bade2fbe8d2a47cf338f95d2f7f6 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:10:07 -0700 Subject: [PATCH 016/107] [bugfix]: address #45 review: int64 FP8 epilogue offsets, order-independent NVFP4/FP8 conversion, splice scope - fp8_kernels: int64 row offsets (a 78k-token fc_in output has 2.2e9 elements) - _maybe_quantize_model: handle NVFP4 (+ mixed FP8) before the per-module walk - step splice: primary transformer component only; no env mutation - layerwise offload: tolerate malformed / all-resident FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS - NVFP4 encoder fallback: same input dtype contract as the FP4 path - benchmarks: new _build_block_mask contract in bench_code.py, extra_env overrides, --timed validation, generator shutdown on failure - drop the stale playground 'Older clip' assertion (that UI change did not survive the rebase) --- fastvideo/hooks/layerwise_offload.py | 11 +++- fastvideo/layers/quantization/fp8_kernels.py | 5 +- .../encoders/minimax_h3_checkpoint_nvfp4.py | 2 +- fastvideo/models/loader/component_loader.py | 9 ++-- fastvideo/models/loader/fsdp_load.py | 37 +++++++------ .../entrypoints/test_openai_video_client.py | 1 - scripts/benchmarks/minimax_h3_pro6000/app.py | 4 +- .../minimax_h3_pro6000/bench_code.py | 6 +-- .../minimax_h3_pro6000/bench_headline.py | 52 ++++++++++--------- 9 files changed, 71 insertions(+), 56 deletions(-) diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index fefb908b08..1356219b2f 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -170,7 +170,12 @@ 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 - resident = int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0")) + 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): @@ -190,6 +195,10 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False): hook_mgr.append_forward_hook(hook) break if len(state_list) == 0: + if resident > 0: + logger.info("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%d keeps every block resident; nothing to offload", + resident) + return raise ValueError("No nn.ModuleList found in the model for layerwise offloading.") # circular linking of states diff --git a/fastvideo/layers/quantization/fp8_kernels.py b/fastvideo/layers/quantization/fp8_kernels.py index 5f8de6be83..a589262c3a 100644 --- a/fastvideo/layers/quantization/fp8_kernels.py +++ b/fastvideo/layers/quantization/fp8_kernels.py @@ -27,7 +27,7 @@ def rowwise_scaled_mm_is_slow() -> bool: @triton.jit def _quantize_rowwise_kernel(x_ptr, q_ptr, s_ptr, K, stride_x, stride_q, BLOCK_K: tl.constexpr): - row = tl.program_id(0) + row = tl.program_id(0).to(tl.int64) x_row = x_ptr + row * stride_x amax = tl.zeros((BLOCK_K, ), dtype=tl.float32) for k in range(0, K, BLOCK_K): @@ -56,9 +56,10 @@ def quantize_rowwise_fp8(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor @triton.jit def _scale_rows_cols_kernel(o_ptr, sx_ptr, sw_ptr, M, N, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr): - rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) + rows = tl.program_id(0).to(tl.int64) * BLOCK_M + tl.arange(0, BLOCK_M) cols = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N) mask = (rows[:, None] < M) & (cols[None, :] < N) + # int64 offsets: a 78k-token fc_in output has 2.2e9 elements, past int32. ptrs = o_ptr + rows[:, None] * N + cols[None, :] v = tl.load(ptrs, mask=mask, other=0.0).to(tl.float32) sx = tl.load(sx_ptr + rows, mask=rows < M, other=0.0) diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 5479f797f1..686346e8ba 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -424,13 +424,13 @@ def process_weights_after_loading(self, layer: nn.Module) -> None: @staticmethod def _apply_finalized(layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor: + x = _coerce_fp4_input_dtype(x) 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, float(layer.weight_global_scale.item()), x.dtype) return torch.nn.functional.linear(x, weight, None if bias is None else bias.to(x.dtype)) - x = _coerce_fp4_input_dtype(x) original_shape = x.shape if x.numel() == 0: # An empty prompt has nothing to quantize; the FP4 kernels are not defined for zero rows. diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 3180062809..7e668a1709 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -1212,12 +1212,15 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): model = model.to(get_local_torch_device()) # FASTVIDEO_H3_SPLICE_TRANSFORMER=: a second checkpoint of the same architecture # runs denoising steps FASTVIDEO_H3_SPLICE_FROM_STEP (default 4) onward. - splice_path = os.environ.pop("FASTVIDEO_H3_SPLICE_TRANSFORMER", None) - if splice_path and hasattr(model, "attach_step_splice"): + # Only the primary ``transformer`` component splices; the spliced load itself never does. + splice_path = os.environ.get("FASTVIDEO_H3_SPLICE_TRANSFORMER") + if (splice_path and not getattr(self, "_loading_splice", False) and hasattr(model, "attach_step_splice") + and os.path.basename(os.path.normpath(model_path)) == "transformer"): + self._loading_splice = True try: late = self.load(splice_path, fastvideo_args) finally: - os.environ["FASTVIDEO_H3_SPLICE_TRANSFORMER"] = splice_path + self._loading_splice = False from_step = int(os.environ.get("FASTVIDEO_H3_SPLICE_FROM_STEP", "4")) model.attach_step_splice(late, from_step) logger.info("Step splice: steps >= %d run the transformer from %s", from_step, splice_path) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index b63b50a067..7eb7986a98 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -123,29 +123,28 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor convert_model_to_mxfp8, ) + # NVFP4 may share a model with FP8 linears (e.g. NVFP4 FFN + FP8 attention); handle that pair before the + # per-module walk so the result does not depend on which quantized module comes first. + nvfp4_modules = [m for m in model.modules() if isinstance(getattr(m, "quant_method", None), NVFP4QuantizeMethod)] + if nvfp4_modules: + mixed_fp8 = any(isinstance(getattr(m, "quant_method", None), FP8QuantizeMethod) for m in model.modules()) + if any(getattr(module, "_nvfp4_weight", None) is not None for module in nvfp4_modules): + logger.info("NVFP4 packed export already populated; skipping runtime weight conversion") + elif defer_weight_conversion_until_lora_merge: + logger.info("Deferring NVFP4 weight conversion until the inference LoRA merge completes") + return + else: + logger.info("Converting loaded model weights for NVFP4 linear layers") + convert_model_to_nvfp4(model) + if mixed_fp8: + logger.info("Converting the FP8 linears of a mixed NVFP4/FP8 model") + convert_model_to_fp8(model) + return + qat_train_attached = 0 qat_train_skipped = 0 for mod in model.modules(): qm = getattr(mod, "quant_method", None) - if isinstance(qm, NVFP4QuantizeMethod): - mixed_fp8 = any(isinstance(getattr(m, "quant_method", None), FP8QuantizeMethod) for m in model.modules()) - if any( - getattr(module, "_nvfp4_weight", None) is not None for module in model.modules() - if isinstance(getattr(module, "quant_method", None), NVFP4QuantizeMethod)): - logger.info("NVFP4 packed export already populated; skipping runtime weight conversion") - if mixed_fp8: - logger.info("Converting the FP8 linears of a mixed NVFP4/FP8 model") - convert_model_to_fp8(model) - return - if defer_weight_conversion_until_lora_merge: - logger.info("Deferring NVFP4 weight conversion until the inference LoRA merge completes") - return - logger.info("Converting loaded model weights for NVFP4 linear layers") - convert_model_to_nvfp4(model) - if any(isinstance(getattr(m, "quant_method", None), FP8QuantizeMethod) for m in model.modules()): - logger.info("Converting the FP8 linears of a mixed NVFP4/FP8 model") - convert_model_to_fp8(model) - return if isinstance(qm, NVFP4QATQuantizeMethod): logger.info("Converting loaded model weights for NVFP4-QAT linear layers") convert_model_to_fp4(model) diff --git a/fastvideo/tests/entrypoints/test_openai_video_client.py b/fastvideo/tests/entrypoints/test_openai_video_client.py index 7355ce8c44..3e48b8e647 100644 --- a/fastvideo/tests/entrypoints/test_openai_video_client.py +++ b/fastvideo/tests/entrypoints/test_openai_video_client.py @@ -191,7 +191,6 @@ def test_playground_assets_and_config_do_not_generate(local_server): assert response.headers["X-Content-Type-Options"] == "nosniff" html = response.read().decode() assert "Generate video" in html - assert "Older clip" in html assert 'src="./playground.js"' in html for asset in ["playground.js", "playground.css"]: with urlopen(origin + "/playground/" + asset) as response: diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index 4352d9e812..40dc13c35a 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -559,8 +559,8 @@ def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dic _install_kernel() model = f"/vol/models/{repo.split('/')[-1]}" run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}" - env = dict(os.environ, **FAST_ENV, **(extra_env or {}), HEADLINE_OUT="/vol/outputs/headline", - HEADLINE_DEVICE=f"{gpus}x RTX PRO 6000", PYTHONPATH="/src/fastvideo") + env = {**os.environ, **FAST_ENV, **(extra_env or {}), "HEADLINE_OUT": "/vol/outputs/headline", + "HEADLINE_DEVICE": f"{gpus}x RTX PRO 6000", "PYTHONPATH": "/src/fastvideo"} proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, "--prompts", "/root/headline_prompts.json"], env=env, capture_output=True, text=True, cwd="/root") diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_code.py b/scripts/benchmarks/minimax_h3_pro6000/bench_code.py index 3f3cc590b9..c45557b839 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_code.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_code.py @@ -106,7 +106,7 @@ def vsa_full(sparsity=0.8, with_gate=True): qp = _pool_tiles(qt, vbs, 64) kp = _pool_tiles(kt, vbs, 64) scores = torch.matmul(qp, kp.transpose(-2, -1)) / HD**0.5 - mask = _build_block_mask(scores, n_prefix, n_video, sparsity, True) + mask = _build_block_mask(scores, n_prefix, sparsity, True, ((n_prefix, n_prefix + n_video), ), (sparsity, )) out, _ = bsa64(qt.transpose(1, 2).contiguous(), kt.transpose(1, 2).contiguous(), vt.transpose(1, 2).contiguous(), mask, vbs) out = out.transpose(1, 2).contiguous() @@ -257,7 +257,7 @@ def tile(x): qt, kt, vt = tile(q), tile(k), tile(v) scores = torch.matmul(_pool_tiles(qt.transpose(1, 2), vbs, tt), _pool_tiles(kt.transpose(1, 2), vbs, tt).transpose(-2, -1)) - mask = _build_block_mask(scores, n_prefix, n_video, 0.8, True) + mask = _build_block_mask(scores, n_prefix, 0.8, True, ((n_prefix, n_prefix + n_video), ), (0.8, )) q2k_idx, q2k_num, kvv, quad = fa.vsa_tile_mask_to_fp4_blocks(mask, tt, vbs, validate=True) Lk = q2k_idx.shape[2] * 128 pad = lambda x: F.pad(x, (0, 0, 0, Lk - Lpad)).contiguous() @@ -314,7 +314,7 @@ def tile(x): qs, ks, vs = (F.pad(x, (0, 0, 0, 0, 0, rows - Lpad)).contiguous() for x in (qt, kt, vt)) out = {} for sparsity in (0.8, 0.9): - mask = _build_block_mask(scores, n_prefix, n_video, sparsity, True) + mask = _build_block_mask(scores, n_prefix, sparsity, True, ((n_prefix, n_prefix + n_video), ), (sparsity, )) m = F.pad(mask, (0, nt2 - n_tiles, 0, nt2 - n_tiles), value=False) B, H = m.shape[:2] d = {"tile64x64": m.float().mean().item()} diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py index 85d3c6de87..3815a20549 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -29,6 +29,8 @@ def main(): ap.add_argument("--timed", type=int, default=2) ap.add_argument("--sparsity", type=float, default=None, help="default: the checkpoint contract's vsa_sparsity") a = ap.parse_args() + if a.timed < 1: + ap.error("--timed must be at least 1") contract = json.load(open(os.path.join(a.model_dir, "fastvideo_inference.json"))) steps = contract["dmd_denoising_steps"] @@ -67,30 +69,32 @@ def main(): results = {"run_name": a.run_name, "model_dir": a.model_dir, "num_gpus": a.num_gpus, "nvfp4_profile": a.nvfp4_profile, "load_s": round(time.perf_counter() - t0, 1), "env": env, "settings": {}} - for name in a.settings.split(","): - width, height, frames = SETTINGS[name] - runs = [] - plan = [(PROMPT_IDS[0], True)] + [(pid, False) for _ in range(a.timed) for pid in PROMPT_IDS] - for i, (pid, warmup) in enumerate(plan): - path = os.path.join(out_dir, f"{name}_{pid}_{'warmup' if warmup else i}.mp4") - t = time.perf_counter() - generator.generate_video(prompt=texts[pid], height=height, width=width, num_frames=frames, fps=24, - guidance_scale=1.0, num_inference_steps=len(steps) + 1, seed=1234, - output_path=path, save_video=True) - wall = round(time.perf_counter() - t, 2) - runs.append({"prompt": pid, "warmup": warmup, "e2e_s": wall, "path": path}) - print("RUN", name, pid, "warmup" if warmup else "timed", wall, flush=True) - if run is not None and not warmup: - import wandb - run.log({f"{name}/e2e_s": wall, f"{name}/{pid}": wandb.Video(path, fps=24, format="mp4")}) - timed = [r["e2e_s"] for r in runs if not r["warmup"]] - results["settings"][name] = {"width": width, "height": height, "frames": frames, - "e2e_median_s": statistics.median(timed), "e2e_min_s": min(timed), "runs": runs} - print("SETTING", name, json.dumps(results["settings"][name]), flush=True) - if run is not None: - run.summary[f"{name}_e2e_median_s"] = statistics.median(timed) - json.dump(results, open(os.path.join(out_dir, "results.json"), "w"), indent=1) - generator.shutdown() + try: + for name in a.settings.split(","): + width, height, frames = SETTINGS[name] + runs = [] + plan = [(PROMPT_IDS[0], True)] + [(pid, False) for _ in range(a.timed) for pid in PROMPT_IDS] + for i, (pid, warmup) in enumerate(plan): + path = os.path.join(out_dir, f"{name}_{pid}_{'warmup' if warmup else i}.mp4") + t = time.perf_counter() + generator.generate_video(prompt=texts[pid], height=height, width=width, num_frames=frames, fps=24, + guidance_scale=1.0, num_inference_steps=len(steps) + 1, seed=1234, + output_path=path, save_video=True) + wall = round(time.perf_counter() - t, 2) + runs.append({"prompt": pid, "warmup": warmup, "e2e_s": wall, "path": path}) + print("RUN", name, pid, "warmup" if warmup else "timed", wall, flush=True) + if run is not None and not warmup: + import wandb + run.log({f"{name}/e2e_s": wall, f"{name}/{pid}": wandb.Video(path, fps=24, format="mp4")}) + timed = [r["e2e_s"] for r in runs if not r["warmup"]] + results["settings"][name] = {"width": width, "height": height, "frames": frames, + "e2e_median_s": statistics.median(timed), "e2e_min_s": min(timed), "runs": runs} + print("SETTING", name, json.dumps(results["settings"][name]), flush=True) + if run is not None: + run.summary[f"{name}_e2e_median_s"] = statistics.median(timed) + json.dump(results, open(os.path.join(out_dir, "results.json"), "w"), indent=1) + finally: + generator.shutdown() if run is not None: run.summary["load_s"] = results["load_s"] run.finish() From ccebd84a7adeb6bf87154a4ea56f70c0efa1e6b0 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:10:51 -0700 Subject: [PATCH 017/107] [misc]: headline Modal runs keep the full log on the volume and report error lines --- scripts/benchmarks/minimax_h3_pro6000/app.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index 40dc13c35a..a1b72b63d0 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -564,12 +564,17 @@ def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dic proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, "--prompts", "/root/headline_prompts.json"], env=env, capture_output=True, text=True, cwd="/root") + log_text = proc.stdout + proc.stderr + run_dir = pathlib.Path("/vol/outputs/headline") / run_name + run_dir.mkdir(parents=True, exist_ok=True) + (run_dir / "run.log").write_text(log_text) volume.commit() - tail = (proc.stdout + proc.stderr)[-6000:] - result_path = pathlib.Path("/vol/outputs/headline") / run_name / "results.json" + errors = [line for line in log_text.splitlines() + if any(k in line for k in ("Error", "error:", "Traceback", "Killed", "OOM", "out of memory"))][-40:] + result_path = run_dir / "results.json" results = json.loads(result_path.read_text()) if result_path.exists() else {} - return {"run": run_name, "returncode": proc.returncode, "results": results, - "log_tail": tail if proc.returncode else tail[-1500:]} + return {"run": run_name, "returncode": proc.returncode, "results": results, "errors": errors, + "log_tail": log_text[-1500:]} @app.function(image=headline_image, gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=2 * 3600, volumes={"/vol": volume}) From 54f335c5be1faeca9fdc21060d511fa1cefd6865 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:19:52 -0700 Subject: [PATCH 018/107] [bugfix]: restore import os in fsdp_load (dropped in the rebase onto main) --- fastvideo/models/loader/fsdp_load.py | 1 + 1 file changed, 1 insertion(+) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 7eb7986a98..01ccf3a75a 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -6,6 +6,7 @@ from __future__ import annotations import contextlib +import os import re from collections.abc import Callable, Generator from itertools import chain From 4d3371711d9b56982a3edc0219245571355be286 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:43:22 -0700 Subject: [PATCH 019/107] [feat]: headline benchmark: HEADLINE_VAE_PARALLEL=1 decodes VAE tiles on every GPU --- .../benchmarks/minimax_h3_pro6000/bench_headline.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py index 3815a20549..7d623392f2 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -48,10 +48,12 @@ def main(): engine["execution_backend"] = os.environ["HEADLINE_BACKEND"] # "ray" for multi-node if a.nvfp4_profile != "none": engine["quantization"] = {"transformer_quant": "NVFP4", "layer_profile": a.nvfp4_profile} - config = {"model_path": a.model_dir, "engine": engine, - "pipeline": {"experimental": {"attention_backend": "VIDEO_SPARSE_ATTN_H3", "VSA_sparsity": sparsity, - "VSA_tile_size": 64, "h3_sequential_load": False, - "inference_torch_compile": False}}} + experimental = {"attention_backend": "VIDEO_SPARSE_ATTN_H3", "VSA_sparsity": sparsity, "VSA_tile_size": 64, + "h3_sequential_load": False, "inference_torch_compile": False} + if os.environ.get("HEADLINE_VAE_PARALLEL") == "1" and a.num_gpus > 1: + # Decode VAE tiles on every GPU instead of rank 0 only. + experimental.update(vae_parallel_decode=True, vae_parallel_decode_strategy="gather") + config = {"model_path": a.model_dir, "engine": engine, "pipeline": {"experimental": experimental}} env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA"))} run = None if os.environ.get("WANDB_PROJECT"): From 5386585d96a994590c6f40d28d2bffefff52835d Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:54:29 -0700 Subject: [PATCH 020/107] [feat]: headline benchmark: engine/pipeline placement overrides for memory-limited GPUs --- scripts/benchmarks/minimax_h3_pro6000/bench_headline.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py index 7d623392f2..669581d804 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -53,6 +53,15 @@ def main(): if os.environ.get("HEADLINE_VAE_PARALLEL") == "1" and a.num_gpus > 1: # Decode VAE tiles on every GPU instead of rank 0 only. experimental.update(vae_parallel_decode=True, vae_parallel_decode_strategy="gather") + # Placement overrides for memory-limited GPUs, e.g. on a 32 GB RTX 5090: + # HEADLINE_ENGINE_JSON='{"offload": {"text_encoder": true, "pin_cpu_memory": true}}' + # HEADLINE_EXPERIMENTAL_JSON='{"h3_sequential_load": true}' + for key, value in json.loads(os.environ.get("HEADLINE_ENGINE_JSON", "{}")).items(): + if isinstance(value, dict) and isinstance(engine.get(key), dict): + engine[key].update(value) + else: + engine[key] = value + experimental.update(json.loads(os.environ.get("HEADLINE_EXPERIMENTAL_JSON", "{}"))) config = {"model_path": a.model_dir, "engine": engine, "pipeline": {"experimental": experimental}} env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA"))} run = None From cb79431b1141ee13a9315bc3113adede3579c104 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 11:56:38 -0700 Subject: [PATCH 021/107] [feat]: headline Modal entrypoint takes extra env and a run tag --- scripts/benchmarks/minimax_h3_pro6000/app.py | 21 ++++++++++---------- 1 file changed, 11 insertions(+), 10 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index a1b72b63d0..0db7b113a7 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -555,10 +555,10 @@ def headline_fetch(repo: str) -> str: return _sh(f"du -sh {local}/*") -def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dict: +def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None, tag: str = "") -> dict: _install_kernel() model = f"/vol/models/{repo.split('/')[-1]}" - run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}" + run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}{tag}" env = {**os.environ, **FAST_ENV, **(extra_env or {}), "HEADLINE_OUT": "/vol/outputs/headline", "HEADLINE_DEVICE": f"{gpus}x RTX PRO 6000", "PYTHONPATH": "/src/fastvideo"} proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, @@ -578,26 +578,27 @@ def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None) -> dic @app.function(image=headline_image, gpu="RTX-PRO-6000", memory=131072, cpu=8, timeout=2 * 3600, volumes={"/vol": volume}) -def headline1(repo: str, profile: str, extra_env: dict | None = None) -> dict: - return _headline(repo, 1, profile, extra_env) +def headline1(repo: str, profile: str, extra_env: dict | None = None, tag: str = "") -> dict: + return _headline(repo, 1, profile, extra_env, tag) @app.function(image=headline_image, gpu="RTX-PRO-6000:4", memory=196608, cpu=16, timeout=2 * 3600, volumes={"/vol": volume}) -def headline4(repo: str, profile: str, extra_env: dict | None = None) -> dict: - return _headline(repo, 4, profile, extra_env) +def headline4(repo: str, profile: str, extra_env: dict | None = None, tag: str = "") -> dict: + return _headline(repo, 4, profile, extra_env, tag) @app.function(image=headline_image, gpu="RTX-PRO-6000:8", memory=262144, cpu=32, timeout=2 * 3600, volumes={"/vol": volume}) -def headline8(repo: str, profile: str, extra_env: dict | None = None) -> dict: - return _headline(repo, 8, profile, extra_env) +def headline8(repo: str, profile: str, extra_env: dict | None = None, tag: str = "") -> dict: + return _headline(repo, 8, profile, extra_env, tag) @app.local_entrypoint() -def headline(repo: str, profile: str = "h3_dit_ffn", gpus: str = "1,4,8", skip_fetch: bool = False): +def headline(repo: str, profile: str = "h3_dit_ffn", gpus: str = "1,4,8", skip_fetch: bool = False, + extra_env: str = "{}", tag: str = ""): if not skip_fetch: print(headline_fetch.remote(repo)) fns = {"1": headline1, "4": headline4, "8": headline8} - calls = [fns[g].spawn(repo, profile) for g in gpus.split(",")] + calls = [fns[g].spawn(repo, profile, json.loads(extra_env), tag) for g in gpus.split(",")] out = HERE / "headline_results" out.mkdir(exist_ok=True) for call in calls: From a97d23f094c09b33b07da5ced4b714f50b54c0c1 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 12:03:45 -0700 Subject: [PATCH 022/107] [bugfix]: H3 pinned swaps use exact-size cudaHostRegister arenas torch's pinned allocator rounds every block up to a power of two, so parking the pruned NVFP4 DiT (~20 GB) next to the NVFP4 encoder (15 GB) overran a 60 GB container on the RTX 5090. One registered arena per module pins exactly the bytes needed; buffers now keep a persistent host copy as well. --- .../basic/minimax_h3/minimax_h3_pipeline.py | 61 +++++++++++++++---- 1 file changed, 48 insertions(+), 13 deletions(-) diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index afdd902d08..c09e39b0e9 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -76,29 +76,64 @@ 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 a persistent pinned host copy. + """Move a module's tensors between the GPU and persistent, exactly sized pinned host copies. 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 are copied every time. + 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. """ store = module.__dict__.setdefault("_pinned_host_tensors", {}) + arenas = module.__dict__.setdefault("_pinned_host_arenas", []) params = dict(module.named_parameters()) - for name, tensor in list(params.items()) + list(module.named_buffers()): - if tensor is None: - continue - if device.type == "cpu": + named = [(name, tensor) for name, tensor in list(params.items()) + list(module.named_buffers()) + if tensor is not None] + 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: if tensor.device.type == "cpu": continue - host = store.get(name) if name in params else None - if host is None or host.shape != tensor.shape or host.dtype != tensor.dtype: - host = torch.empty(tensor.shape, dtype=tensor.dtype, pin_memory=True) + host = store[name] + if name not in params and name not in fresh: host.copy_(tensor) - if name in params: - store[name] = host tensor.data = host - elif tensor.device != device: - tensor.data = tensor.data.to(device, non_blocking=True) + else: + for _, tensor in named: + if 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() From a21676815b568921a45f9bf19104cd8125c97b6f Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sat, 3 Oct 2026 09:38:20 -0700 Subject: [PATCH 023/107] [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 024/107] [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 025/107] [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 026/107] [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 027/107] [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 028/107] [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 029/107] [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 030/107] [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 031/107] [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 032/107] [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 033/107] [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 034/107] [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 035/107] [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 036/107] [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 037/107] [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 038/107] [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 039/107] [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 040/107] [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 041/107] [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 042/107] [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 043/107] [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 044/107] [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 045/107] [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 046/107] [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. From 5a377cae78081ffa3dee8861142a4f5fe6e8d355 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Thu, 1 Oct 2026 14:03:15 -0700 Subject: [PATCH 047/107] [docs]: FastH3 NVFP4 on RTX PRO 6000 Usage (conversion, profiles, env switches), the end-to-end and per-block measurements, the block densities each kernel granularity computes, what was tried and not shipped (dense FP4 for VSA students, the 8-GPU FP8 Ulysses exchange kept on h3-sm120-experimental), and known limitations. --- docs/inference/fasth3_rtx_pro_6000.md | 260 ++++++++++++++++++++++++++ docs/inference/optimizations.md | 5 + mkdocs.yml | 1 + 3 files changed, 266 insertions(+) create mode 100644 docs/inference/fasth3_rtx_pro_6000.md diff --git a/docs/inference/fasth3_rtx_pro_6000.md b/docs/inference/fasth3_rtx_pro_6000.md new file mode 100644 index 0000000000..9debdf0430 --- /dev/null +++ b/docs/inference/fasth3_rtx_pro_6000.md @@ -0,0 +1,260 @@ +# FastH3 NVFP4 on RTX PRO 6000 (sm_120) + +This page covers serving the FastH3 NVFP4 checkpoints on one RTX PRO 6000 +Blackwell (sm_120, 96 GB) with every component resident: the NVFP4 text +encoder, the NVFP4 denoiser and an INT8 light VAE. It also records what was +measured and tried along the way. Every switch is opt-in and defaults to the +existing behavior. + +## Results + +All results are for one RTX PRO 6000 (Modal), a 10.1 s clip at 1344x768 +(243 frames, 73.6k packed tokens), the light INT8 VAE and warm runs. + +| Checkpoint | Denoise | Video decode | End to end | Peak memory | +| --- | ---: | ---: | ---: | ---: | +| [`FastH3-4-step-Preview-v1-VSA-DataFree-NVFP4`](https://huggingface.co/FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree-NVFP4) (4 forwards, VSA 0.9) | 30.1–31.0 s | 9.0 s | **41.2 / 42.4 s** | 79.6 GB | +| [`FastH3-8-Step-V2-NVFP4`](https://huggingface.co/FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4) (8 forwards, VSA 0.8) | 75.6 s | 9.0 s | **86.5 s** | — | + +These end-to-end runs predate the warp-skip kernel change below, which cuts +sparse attention by a further 1.5x, so they are upper bounds. Conditioning +takes about 0.12 s, frame post-processing plus MP4 writing about 0.9 s, and +the first clip at a new shape about 200 s (VAE compile). + +At 480p (124 frames, 15.1k tokens) V2 8-step denoises in 12.5 s on a cold +run, down from 16.8 s warm before this work. + +## Usage + +### 1. Convert the checkpoint + +The published NVFP4 checkpoints use ModelOpt's unified Hugging Face layout. +`convert_minimax_h3_modelopt_nvfp4_dit.py` repacks it into FastVideo's packed +export, `transformer/nvfp4_weights.safetensors`. The calibrated FFN weights +and scales are carried over bit for bit and only the scale bytes are +swizzled. Optionally it also quantizes the BF16 attention projections and the +VSA compression gates: + +```bash +python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \ + --src /models/FastH3-8-Step-V2-NVFP4/transformer \ + --dst /models/fasth3-v2-fv/transformer \ + --quantize-attention --quantize-gate +``` + +Every exported linear is probed through the loader's `mm_fp4` path against a +BF16 matmul with the dequantized weight. Conversion refuses to write when any +error exceeds `--max-probe-error` (default 0.3; the converted V2 and 4-step +checkpoints probe at 0.14). The rest of the model folder (text encoder, +VAEs, schedulers, `fastvideo_inference.json`) is used as-is; the +NVFP4 text encoder comes from +`convert_minimax_h3_text_encoder_nvfp4.py`. + +The flags select the `layer_profile` to load the result with: + +| Flags | Linears in NVFP4 | `layer_profile` | +| --- | --- | --- | +| *(none)* | FFN `fc_in`/`fc_out` | `h3_dit_ffn` | +| `--quantize-attention` | + attention `to_{q,k,v,out}` | `h3_dit` | +| `--quantize-attention --quantize-gate` | + VSA `to_gate_compress` | `h3_dit_vsa` | + +### 2. Generate + +```python +import os +os.environ.update({ + "FASTVIDEO_H3_VSA_FP4": "1", # sparse FP4 attention + "FASTVIDEO_MINIMAX_H3_FUSIONS": "all", # Triton norm/modulate/RoPE/SwiGLU fusions + "FASTVIDEO_NVFP4_MM_BACKEND": "cutlass", # see "FP4 GEMM backend" below + "FASTVIDEO_H3_VAE_TILE_BATCH": "28", # one decoder call per 1344x768 tile grid +}) +from fastvideo import VideoGenerator + +generator = VideoGenerator.from_config({ + "model_path": "/models/fasth3-v2-fv", + "engine": { + "num_gpus": 1, + "quantization": {"transformer_quant": "NVFP4", "layer_profile": "h3_dit_vsa"}, + "compile": {"enabled": False, "vae_enabled": True}, + }, + "pipeline": {"experimental": {"attention_backend": "VIDEO_SPARSE_ATTN_H3", + "VSA_sparsity": 0.8, "VSA_tile_size": 64}}, +}) +generator.generate({"prompt": "...", "sampling": {"height": 768, "width": 1344, "num_frames": 243, + "num_inference_steps": 9}}) +``` + +Use `VSA_sparsity` 0.9 and `num_inference_steps` 5 for the 4-step checkpoint +(see its `fastvideo_inference.json`). Frame counts must be `17n + 5`: 243 +frames is the closest to 10 s. + +## What changed + +### Block-sparse FP4 attention for VSA tiles (`fastvideo-kernel`) + +SageAttention3's sm_120 FP4 kernel (`attn_qat_infer`) gains a block-sparse +forward, `fwd_sparse`, exposed as `sageattn_blackwell_sparse` (head-major +inputs) and `sageattn_blackwell_sparse_bshd` (sequence-major inputs, quantized +in place without a transpose). `vsa_tile_mask_to_fp4_blocks` turns a VSA tile +mask into the kernel's lists: + +- **Block lists.** Query block `m` visits only the 128-token KV blocks in + `q2k_idx[b, h, m, :q2k_num[b, h, m]]`. +- **Quadrant masks for 64-token tiles.** The kernel computes on 128x128 + blocks, but V2 and the 4-step preview use 64-token VSA tiles. Each listed + block carries a 4-bit `q2k_quad` (one bit per 64x64 quadrant). The kernel + masks unselected quadrants to `-inf`, so the result is exactly VSA's tile-64 + semantics. +- **Valid counts per 64-column half.** `kv_valid` gives the valid tokens in + each 64-column half, so partially filled tiles can pad mid-block. +- **Warp-level skipping.** Each MMA warp owns 16 query rows and so sits + inside one 64-row half. A warp skips a listed block that its half did not + select, and the P·V chunk of a key half it did not select. Masked scores + contribute exactly zero, and the warp sharing its tensor-core partition runs + faster meanwhile. This recovers most of the work that pairing 64-token tiles + into 128-token blocks adds. +- **First-visited block.** Lists run in descending block order because the + kernel visits them last entry first. Block 0 (the first prefix tile, which + VSA-H3's exempt mode gives every query) is therefore visited first, so every + row starts from a finite running max. `validate=True` checks this. The model + integration uses only exempt mode. + +The dense and sparse entry points also stop allocating `delta_s`. With Q +smoothing off (the default), each call used to allocate and zero a +`[B, H, L/128, L]` fp32 tensor: 9.5 GB at 73k tokens. Its int32 batch stride +also overflowed the TMA descriptor ("Failed to initialize the TMA descriptor +1", then an illegal instruction), so FP4 attention could not run 10 s 1344x768 +clips at all. A cached `[B, H, 1, L]` zero row read with `per_block_mean=False` +replaces it, and the outputs are bit-identical. + +Correctness (`fastvideo-kernel/tests/test_attn_qat_infer_sparse.py`, RTX PRO +6000): + +| Layout | Error vs token-masked fp32 | Dense FP4 floor | +| --- | ---: | ---: | +| 64-token tiles, odd count, partial tiles | 0.190 | 0.191 | +| 64-token tiles, even count | 0.191 | 0.193 | +| 256-token tiles, partial tile | 0.195 | 0.196 | + +Errors are relative L2. The sparse kernel sits exactly at the FP4 noise +floor; random Gaussian inputs make that floor large. Full block lists +reproduce the dense kernel bit for bit. + +### Model integration (`FASTVIDEO_H3_VSA_FP4=1`) + +`fastvideo/models/dits/minimax_h3_vsa_fp4.py` replaces only the attention core +of `MiniMaxH3Attention`. VSA-H3's tile pooling, top-k mask, exempt prefix and +gated compression branch are unchanged. Per block, it: + +1. Gathers the attention input into tile order once (one `hidden_size`-wide + pass). Pad rows stay zero, so the q/k/v pad rows are exactly zero through + the bias-free projections, RMSNorm and RoPE. +2. Quantizes that input once and shares it between `to_q`, `to_k` and `to_v`. + NVFP4 activations use a unit global scale, so this is exact. +3. Applies QK-norm and RoPE with tile-ordered `cos`/`sin`, computed once per + step. +4. Runs the sparse FP4 kernel on sequence-major tensors and gathers the output + back to packed order before `to_out`. + +This replaces the generic path's concat, four tile scatters and three +transposes. The route applies only to no-grad, non-compiled, single +sequence-parallel-rank calls in exempt mode; everything else keeps the +existing path. + +Two smaller pieces ship alongside it: + +- **Packed gate check.** With `--quantize-gate`, `to_gate_compress` loses its + BF16 weight, so the gate-activity check reads the packed E2M1 bytes instead. +- **Compiled residual.** With `FASTVIDEO_MINIMAX_H3_FUSIONS` enabled, each + block's final `hidden + gate[indices] * ffn_out` runs as one compiled op + instead of materializing the gathered gate. + +### FP4 GEMM backend (`FASTVIDEO_NVFP4_MM_BACKEND`) + +FlashInfer's `mm_fp4(backend="auto")` picks a kernel about 2x slower than +`cutlass` or `cudnn` on sm_120 once activations reach tens of thousands of +rows. At 15k rows all three match. + +| Linear | 73.6k rows: `auto` | 73.6k rows: `cutlass` | 73.6k rows: `cudnn` | 15.1k rows: `auto` | +| --- | ---: | ---: | ---: | ---: | +| `to_q` (5376→7168) | 7.96 ms | 3.96 ms | 4.24 ms | 0.88 ms | +| `to_out` (7168→5376) | 9.07 ms | 4.09 ms | 4.27 ms | 0.92 ms | +| `fc_in` (5376→28672) | 21.54 ms | 16.59 ms | 16.17 ms | 3.01 ms | +| `fc_out` (14336→5376) | 18.06 ms | 8.09 ms | 8.43 ms | 1.65 ms | + +### Batched VAE tile decode (`FASTVIDEO_H3_VAE_TILE_BATCH`) + +The H3 video VAE decodes 256-pixel spatial tiles one at a time. A 1344x768 +clip is a 4x7 grid per temporal chunk, so a 10 s clip is roughly 400 small +decoder calls. The ViT decoder treats batch entries independently, so +`FASTVIDEO_H3_VAE_TILE_BATCH=N` decodes up to `N` equal-shaped tiles per call; +28 covers a full 1344x768 grid. The decoded tiles are the same as per-tile +decoding. + +## Per-block measurements + +One H3 transformer block (hidden 5376, 56 heads, FFN 14336), RTX PRO 6000: + +| Component | 480p, 124 f (15.1k tokens) | 768p, 243 f (73.6k tokens) | +| --- | ---: | ---: | +| VSA Triton BF16 attention (kernel + pooling/mask) | 14.8 ms | 180.1 ms | +| Tile scatter of q/k/v/gate + gather (generic path) | 2.7 ms | 12.7 ms | +| Dense FP4 attention (SageAttention3) | 10.5 ms | 211.9 ms | +| Sparse FP4, quadrant masks | 7.9 ms | 123.4 ms | +| Sparse FP4, quadrant masks + warp skip | — | **83.0 ms** (VSA 0.8) / **49.3 ms** (VSA 0.9) | +| Dense BF16 SDPA | 18.1 ms | — | +| Modulation: eager / fused / compiled | 4.07 / 1.42 / 0.58 ms | 20.3 / 7.0 / 2.8 ms | +| SwiGLU: eager / fused | 1.53 / 0.88 ms | 7.37 / 4.24 ms | +| QK-norm + RoPE: eager / fused | 4.64 / 1.12 ms | 22.4 / 5.2 ms | + +Before this work, a 480p block cost about 40 ms: 17.7 ms of attention, 12 ms +of linears and 10 ms of eager elementwise ops. Over 50 blocks that is 2.0 s +per step, which matches the measured 2.1 s. + +Block density each kernel granularity computes at 768p (fraction of the +dense attention). "Selected" is what VSA needs; the other columns are what +each block shape computes: + +| VSA sparsity | Selected (64x64) | 128x128 blocks | 64-row x 128 | 128 x 64-col | +| --- | ---: | ---: | ---: | ---: | +| 0.8 | 0.222 | 0.434 | 0.317 | 0.313 | +| 0.9 | 0.125 | 0.254 | 0.180 | 0.177 | + +## What was tried and not shipped + +- **Dense FP4 attention for VSA-trained students.** `ATTN_QAT_INFER` does not + build `to_gate_compress`. A VSA-distilled checkpoint such as V2 carries + trained gates, so the strict loader refuses it ("Parameter + ...to_gate_compress.weight not found"). Dense FP4 attention is also slower + than sparse FP4 at 768p (212 vs 83 ms per block). +- **Multi-GPU (Ulysses) FP8 exchange.** On 8x RTX PRO 6000 (PCIe only), NCCL + all-to-all moves about 21 GB/s per GPU, with NCCL P2P on or off. A BF16 + q/k/v/gate exchange at 73.6k tokens therefore costs 24 ms per block, and the + attention output another 6.6 ms; with FP4 payloads q/k/v/gate drop to 7.2 ms. + The branch `h3-sm120-experimental` keeps a sequence-parallel path that: + - sends q/k/v as FP8 with one scale per token and head; + - never sends the VSA gate, applying it on each rank after a small + all-gather of the per-tile compression output; + - returns the attention output as FP8. + + It is estimated at about 18–20 s per 10 s clip for V2 8-step on 8 GPUs. It + has not executed yet (8-GPU capacity was unavailable), so it is not part of + this change. The same branch holds the Modal drivers behind every number on + this page. +- **64-row query blocks.** The kernel traits allow `kBlockM = 64`, which would + remove the query-side pairing waste. Warp-level skipping recovers most of + that waste without a second kernel instantiation, so it was not built. + +## Known limitations + +- **End-to-end quality.** The kernel matches a masked reference at the FP4 + noise floor. Generated videos have not yet been A/B-compared against the + BF16 Triton VSA path on the H3 audio/video metrics. +- **Activation scales.** The packed export drops ModelOpt's calibrated static + `input_scale` and quantizes activations with a unit global scale and dynamic + per-16 block scales, as FastVideo's NVFP4 linears do elsewhere. +- **Decode cost.** Video decode (9 s at 10 s/768p with the light INT8 VAE) is + the next largest cost after denoising. +- **Hardware.** Everything here targets sm_120. The GeForce RTX 5090 shares + the architecture but has 32 GB, which needs a reduced-AdaLN checkpoint and a + non-resident text encoder at this resolution. diff --git a/docs/inference/optimizations.md b/docs/inference/optimizations.md index da98a0313c..adcfc730ae 100644 --- a/docs/inference/optimizations.md +++ b/docs/inference/optimizations.md @@ -198,6 +198,11 @@ The `attn_qat_infer` kernel hard-gates on **sm_120 (consumer Blackwell / RTX 5090)**; on other GPUs the backend logs a notice and falls back to Flash Attention. See the [Attn-QAT paper](https://arxiv.org/abs/2603.00040). +For VSA-distilled MiniMax-H3 students, the same kernel has a block-sparse +forward that runs VSA's 64-token tile selection in FP4 +(`FASTVIDEO_H3_VSA_FP4=1`). See +[FastH3 NVFP4 on RTX PRO 6000](fasth3_rtx_pro_6000.md). + Enable both halves — attention via the env var, linear via `transformer_quant`: ```python diff --git a/mkdocs.yml b/mkdocs.yml index 81c0840163..f3e3e45f30 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -183,6 +183,7 @@ nav: - Configuration: inference/configuration.md - Offloading: inference/offloading.md - Optimizations: inference/optimizations.md + - FastH3 on RTX PRO 6000: inference/fasth3_rtx_pro_6000.md - ComfyUI: inference/comfyui.md - Support Matrix: inference/support_matrix.md - CLI: inference/cli.md From 3d727ff2020318c8856a46aed8480901bdda93cf Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 21:10:53 -0700 Subject: [PATCH 048/107] [bugfix]: validate sparse FP4 block lists before launch; restore the #44 sparse kernel tests A zero q2k_num or an out-of-range q2k_idx entry made fwd_sparse read outside its index row or the KV tensors. sageattn_blackwell_sparse{,_bshd} take validate=True to run check_sparse_block_lists (host sync) first. test_attn_qat_infer_sparse.py from #44 was dropped when the sm_120 work was folded into the release core; it is restored with a block-list check test. --- fastvideo-kernel/attn_qat_infer/api.py | 27 ++- .../tests/test_attn_qat_infer_sparse.py | 155 ++++++++++++++++++ 2 files changed, 180 insertions(+), 2 deletions(-) create mode 100644 fastvideo-kernel/tests/test_attn_qat_infer_sparse.py diff --git a/fastvideo-kernel/attn_qat_infer/api.py b/fastvideo-kernel/attn_qat_infer/api.py index 03ea3a6565..429d869c7d 100644 --- a/fastvideo-kernel/attn_qat_infer/api.py +++ b/fastvideo-kernel/attn_qat_infer/api.py @@ -277,6 +277,22 @@ def vsa_tile_mask_to_fp4_blocks( return q2k_idx, q2k_num, kv_valid, q2k_quad +def check_sparse_block_lists(q2k_idx: torch.Tensor, q2k_num: torch.Tensor, kv_len: int) -> None: + """Reject block lists the sparse kernel would read out of bounds (host sync). + + Each row needs ``1 <= q2k_num <= q2k_idx.size(-1)`` and every listed index + in ``[0, ceil(kv_len / BLOCK_N))``; a zero count or an out-of-range index + makes the kernel load outside its index row or the KV tensors. + """ + num_kv_blocks = -(-kv_len // BLOCK_N) + if int(q2k_num.min()) < 1 or int(q2k_num.max()) > q2k_idx.size(-1): + raise ValueError(f"q2k_num must be in [1, {q2k_idx.size(-1)}]") + listed = torch.arange(q2k_idx.size(-1), device=q2k_idx.device) < q2k_num.unsqueeze(-1) + idx = q2k_idx[listed] + if idx.numel() and (int(idx.min()) < 0 or int(idx.max()) >= num_kv_blocks): + raise ValueError(f"q2k_idx entries must be in [0, {num_kv_blocks})") + + def sageattn_blackwell_sparse(q, k, v, @@ -286,16 +302,20 @@ def sageattn_blackwell_sparse(q, q2k_quad: torch.Tensor | None = None, per_block_mean=True, single_level_p_quant=True, - sm_scale: float | None = None): + sm_scale: float | None = None, + validate: bool = False): """Block-sparse SageAttention3 FP4 forward (non-causal). Query block ``m`` (``BLOCK_M`` rows) of each (batch, head) attends only to the ``BLOCK_N``-token KV blocks in ``q2k_idx[b, h, m, :q2k_num[b, h, m]]``, restricted to the quadrants in ``q2k_quad`` when given; see :func:`vsa_tile_mask_to_fp4_blocks`. Q/K/V are ``[B, H, L, D]``. + ``validate=True`` runs :func:`check_sparse_block_lists` first. """ QL = q.size(2) KL = k.size(2) + if validate: + check_sparse_block_lists(q2k_idx, q2k_num, KL) is_bf16 = q.dtype == torch.bfloat16 q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean) per_block_mean = delta_s.shape[2] > 1 @@ -316,7 +336,8 @@ def sageattn_blackwell_sparse_bshd(q, kv_valid: torch.Tensor | None = None, q2k_quad: torch.Tensor | None = None, single_level_p_quant=True, - sm_scale: float | None = None) -> torch.Tensor: + sm_scale: float | None = None, + validate: bool = False) -> torch.Tensor: """:func:`sageattn_blackwell_sparse` for ``[B, L, H, D]`` inputs, without copies. The FP4 quantizers read strided input, so the sequence-major tensors a @@ -327,6 +348,8 @@ def sageattn_blackwell_sparse_bshd(q, batch, seq_len, heads, _ = q.shape if seq_len % BLOCK_M: raise ValueError(f"sequence length {seq_len} must be a multiple of {BLOCK_M}") + if validate: + check_sparse_block_lists(q2k_idx, q2k_num, seq_len) qh, kh, vh = (x.transpose(1, 2) for x in (q, k, v)) delta_s = _zero_delta_s(batch, heads, seq_len, q.device) return blockscaled_fp4_attn_sparse(scale_and_quant_fp4(qh), scale_and_quant_fp4_permute(kh), diff --git a/fastvideo-kernel/tests/test_attn_qat_infer_sparse.py b/fastvideo-kernel/tests/test_attn_qat_infer_sparse.py new file mode 100644 index 0000000000..3684b2156b --- /dev/null +++ b/fastvideo-kernel/tests/test_attn_qat_infer_sparse.py @@ -0,0 +1,155 @@ +#!/usr/bin/env python3 +"""Correctness tests for the block-sparse SageAttention3 FP4 forward (``fwd_sparse``). + +Covers the VSA tile layouts MiniMax-H3 uses: 64-token tiles carried by 64x64 +quadrant masks on the kernel's 128x128 blocks (odd and even tile counts), +256-token tiles, partially valid tiles, and an exempt prefix that every query +attends. The sparse kernel is compared with a token-masked fp32 reference and +must stay at the dense FP4 kernel's own error on the same data. Also checks +that full block lists reproduce the dense kernel bit for bit and that the +sequence-major entry point matches the head-major one bit for bit. + +Requires a Blackwell GPU (sm_120a) and the fp4attn_cuda / fp4quant_cuda +extensions built via ``cd fastvideo-kernel && ./build.sh``. + + pytest tests/test_attn_qat_infer_sparse.py -v +""" + +import math +import os +import sys + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +import pytest +import torch +import torch.nn.functional as F + +pytest.importorskip("fp4attn_cuda", reason="ATTN_QAT_INFER FP4 kernels require a sm_120a build") + +from attn_qat_infer.api import (BLOCK_M, BLOCK_N, check_sparse_block_lists, sageattn_blackwell, + sageattn_blackwell_sparse, sageattn_blackwell_sparse_bshd, + vsa_tile_mask_to_fp4_blocks) + +DEVICE = torch.device("cuda") +HEAD_DIM = 128 + + +def _tile_sizes(prefix: tuple[int, ...], video_tiles: int, tile: int, partial: dict[int, int]) -> torch.Tensor: + """Valid tokens per tile: segment-pure prefix chunks, then video tiles (some partial).""" + sizes = [] + for segment in prefix: + full, rem = divmod(segment, tile) + sizes += [tile] * full + ([rem] if rem else []) + sizes += [partial.get(index, tile) for index in range(video_tiles)] + return torch.tensor(sizes, dtype=torch.int32, device=DEVICE) + + +def _exempt_mask(n_prefix: int, n_tiles: int, sparsity: float, heads: int, gen: torch.Generator) -> torch.Tensor: + """VSA-H3 exempt selection: prefix rows dense, prefix keys always, top-k video keys.""" + n_video = n_tiles - n_prefix + k_video = max(1, math.ceil((1 - sparsity) * n_video)) + scores = torch.rand(1, heads, n_tiles, n_video, generator=gen, device=DEVICE) + mask = torch.zeros(1, heads, n_tiles, n_tiles, dtype=torch.bool, device=DEVICE) + mask.scatter_(-1, scores.topk(k_video, dim=-1).indices + n_prefix, True) + mask[..., :n_prefix] = True + mask[:, :, :n_prefix, :] = True + return mask + + +def _case(tile: int, prefix: tuple[int, ...], video_tiles: int, partial: dict[int, int], heads: int = 4): + gen = torch.Generator(device=DEVICE).manual_seed(tile * 1000 + video_tiles) + sizes = _tile_sizes(prefix, video_tiles, tile, partial) + n_tiles = sizes.numel() + n_prefix = n_tiles - video_tiles + rows = math.ceil(n_tiles * tile / BLOCK_M) * BLOCK_M + valid = torch.zeros(rows, dtype=torch.bool, device=DEVICE) + valid[:n_tiles * tile] = (torch.arange(tile, device=DEVICE)[None, :] < sizes[:, None]).reshape(-1) + q, k, v = (torch.randn(1, heads, rows, HEAD_DIM, device=DEVICE, dtype=torch.bfloat16, generator=gen) + for _ in range(3)) + for x in (q, k, v): + x[:, :, ~valid] = 0 + mask = _exempt_mask(n_prefix, n_tiles, 0.8, heads, gen) + return q, k, v, sizes, mask, valid + + +def _rel_l2(a: torch.Tensor, b: torch.Tensor) -> float: + return ((a.float() - b.float()).norm() / b.float().norm()).item() + + +@pytest.mark.parametrize( + ("tile", "prefix", "video_tiles", "partial"), + [ + (64, (77, 46), 19, { + 3: 40, + 11: 1 + }), # odd tile count, short prefix chunks + (64, (64, 40), 9, {}), # even tile count + (256, (100, 70), 4, { + 1: 200 + }), # 256-token tiles, partial tile + ], +) +def test_sparse_matches_masked_reference_at_fp4_floor(tile, prefix, video_tiles, partial) -> None: + q, k, v, sizes, mask, valid = _case(tile, prefix, video_tiles, partial) + q2k_idx, q2k_num, kv_valid, q2k_quad = vsa_tile_mask_to_fp4_blocks(mask, tile, sizes, validate=True) + assert (q2k_quad is not None) == (tile < BLOCK_N) + out = sageattn_blackwell_sparse(q.clone(), k.clone(), v.clone(), q2k_idx, q2k_num, kv_valid, q2k_quad) + + n_tok = sizes.numel() * tile + token_tile = torch.arange(sizes.numel(), device=DEVICE).repeat_interleave(tile) + token_mask = mask[:, :, token_tile][:, :, :, token_tile] & valid[None, None, None, :n_tok] + ref = F.scaled_dot_product_attention(q[:, :, :n_tok].float(), + k[:, :, :n_tok].float(), + v[:, :, :n_tok].float(), + attn_mask=token_mask) + rows = valid[:n_tok] + assert torch.isfinite(out[:, :, :n_tok][:, :, rows]).all() + sparse_err = _rel_l2(out[:, :, :n_tok][:, :, rows], ref[:, :, rows]) + + # FP4 floor: the dense kernel against dense fp32 attention on the valid tokens. + qv, kv, vv = (x[:, :, valid].contiguous() for x in (q, k, v)) + floor = _rel_l2(sageattn_blackwell(qv, kv, vv), F.scaled_dot_product_attention(qv.float(), kv.float(), vv.float())) + assert sparse_err <= 1.05 * floor + 1e-3, (sparse_err, floor) + + +def test_full_lists_reproduce_dense_bitwise() -> None: + gen = torch.Generator(device=DEVICE).manual_seed(0) + q, k, v = (torch.randn(1, 4, 2048, HEAD_DIM, device=DEVICE, dtype=torch.bfloat16, generator=gen) for _ in range(3)) + n_blocks = 2048 // BLOCK_N + q2k_idx = torch.arange(n_blocks, device=DEVICE, dtype=torch.int32).expand(1, 4, 2048 // BLOCK_M, + n_blocks).contiguous() + q2k_num = torch.full((1, 4, 2048 // BLOCK_M), n_blocks, device=DEVICE, dtype=torch.int32) + dense = sageattn_blackwell(q.clone(), k.clone(), v.clone()) + sparse = sageattn_blackwell_sparse(q.clone(), k.clone(), v.clone(), q2k_idx, q2k_num) + assert torch.equal(dense, sparse) + + +def test_bshd_entry_matches_bhsd() -> None: + q, k, v, sizes, mask, _ = _case(64, (77, 46), 19, {3: 40}) + lists = vsa_tile_mask_to_fp4_blocks(mask, 64, sizes) + bhsd = sageattn_blackwell_sparse(q.clone(), k.clone(), v.clone(), *lists) + bshd = sageattn_blackwell_sparse_bshd(*(x.transpose(1, 2).contiguous() for x in (q, k, v)), *lists) + assert torch.equal(bhsd, bshd) + + +def test_first_visited_block_check_rejects_unanchored_rows() -> None: + # Without the exempt prefix a query block may start from a block that masks + # one of its 64-row halves entirely; validate=True must refuse it. + mask = torch.zeros(1, 1, 4, 4, dtype=torch.bool, device=DEVICE) + mask[0, 0, 0, 3] = True # query tile 0 -> key tile 3 only + mask[0, 0, 1, 1] = True # query tile 1 -> key tile 1 only (block 0, other quadrant) + mask[0, 0, 2:, 2:] = True + with pytest.raises(ValueError, match="first block"): + vsa_tile_mask_to_fp4_blocks(mask, 64, torch.full((4, ), 64, dtype=torch.int32, device=DEVICE), validate=True) + + +def test_block_list_check_rejects_out_of_bounds_lists() -> None: + idx = torch.zeros((1, 1, 2, 2), dtype=torch.int32, device=DEVICE) + num = torch.ones((1, 1, 2), dtype=torch.int32, device=DEVICE) + idx[..., 1] = 99 # beyond q2k_num: never read, so not checked + check_sparse_block_lists(idx, num, 2 * BLOCK_N) + with pytest.raises(ValueError, match="q2k_num"): + check_sparse_block_lists(idx, torch.zeros_like(num), 2 * BLOCK_N) + with pytest.raises(ValueError, match="q2k_idx"): + check_sparse_block_lists(idx, num + 1, 2 * BLOCK_N) From 5efdd8b8ea68b24304f66985d42adf3c33cb659c Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 21:10:53 -0700 Subject: [PATCH 049/107] [bugfix]: keep ModelOpt-quantized projections off the dense re-quantize path; converter parity test With --quantize-attention, an attention projection ModelOpt already quantized matched both lists, and its packed uint8 weight was quantized again as if it were BF16. It now stays on the ModelOpt path. The dropped input_scale key is declared in DROPPED_MODELOPT_SUFFIXES with its reason. The new GB200 test converts a synthetic ModelOpt checkpoint, checks the ModelOpt bytes carry over bit for bit, loads the export and compares one forward per linear (calibrated static scales) against the reference. --- ...est_minimax_h3_modelopt_nvfp4_converter.py | 120 ++++++++++++++++++ .../convert_minimax_h3_modelopt_nvfp4_dit.py | 12 +- 2 files changed, 130 insertions(+), 2 deletions(-) create mode 100644 fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py diff --git a/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py new file mode 100644 index 0000000000..466731b117 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_nvfp4_converter.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Converter smoke test: a synthetic ModelOpt NVFP4 checkpoint, converted and loaded, matches its reference.""" + +from __future__ import annotations + +import importlib.util +import json +import subprocess +import sys +from pathlib import Path + +import pytest +import torch +import torch.nn as nn +from safetensors import safe_open +from safetensors.torch import save_file + +import fastvideo.layers.quantization.nvfp4_config as nv + +SCRIPT = Path(__file__).resolve().parents[4] / "scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py" +OUT, IN = 128, 128 +ACT_AMAX = 6.0 + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability() < (10, 0), + reason="NVFP4 mm_fp4 needs a Blackwell GPU") + + +def _converter(): + spec = importlib.util.spec_from_file_location("convert_minimax_h3_modelopt_nvfp4_dit", SCRIPT) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _modelopt_linear(prefix: str, gen: torch.Generator) -> dict[str, torch.Tensor]: + """Random E2M1 codes with positive E4M3 block scales: a valid ModelOpt unified-layout linear.""" + return { + f"{prefix}.weight": torch.randint(0, 256, (OUT, IN // 2), dtype=torch.uint8, generator=gen), + f"{prefix}.weight_scale": (torch.rand(OUT, IN // 16, generator=gen) + 0.5).to(torch.float8_e4m3fn), + f"{prefix}.weight_scale_2": torch.tensor(0.01, dtype=torch.float32), + f"{prefix}.input_scale": torch.tensor(1.0, dtype=torch.float32), + } + + +def _write_source(src: Path) -> dict[str, torch.Tensor]: + gen = torch.Generator().manual_seed(0) + tensors = { + **_modelopt_linear("transformer_blocks.0.ff.net.0.proj", gen), + # Already quantized by ModelOpt: --quantize-attention must keep it on the ModelOpt path. + **_modelopt_linear("transformer_blocks.0.attn.to_k", gen), + "transformer_blocks.0.attn.to_q.weight": (torch.randn(OUT, IN, generator=gen) * 0.05).to(torch.bfloat16), + "transformer_blocks.0.norm.weight": torch.ones(IN, dtype=torch.bfloat16), + } + src.mkdir() + save_file(tensors, str(src / "diffusion_pytorch_model-00001-of-00001.safetensors")) + weight_map = {k: "diffusion_pytorch_model-00001-of-00001.safetensors" for k in tensors} + (src / "diffusion_pytorch_model.safetensors.index.json").write_text(json.dumps({"weight_map": weight_map})) + (src / "config.json").write_text(json.dumps({"quantization_config": {"quant_algo": "NVFP4"}})) + return tensors + + +class _Linear(nn.Module): + + def __init__(self, prefix: str) -> None: + super().__init__() + self.weight = nn.Parameter(torch.empty(OUT, IN, dtype=torch.bfloat16), requires_grad=False) + self.quant_method = nv.NVFP4QuantizeMethod(layer_prefix=prefix) + + +def _model() -> nn.Module: + root = nn.Module() + block = nn.Module() + block.ff, block.attn = nn.Module(), nn.Module() + block.ff.fc_in = _Linear("transformer_blocks.0.ff.fc_in") + block.attn.to_q = _Linear("transformer_blocks.0.attn.to_q") + block.attn.to_k = _Linear("transformer_blocks.0.attn.to_k") + root.transformer_blocks = nn.ModuleList([block]) + return root + + +def test_converted_checkpoint_loads_and_matches_reference(tmp_path) -> None: + pytest.importorskip("flashinfer") + conv = _converter() + src, dst = tmp_path / "src", tmp_path / "dst" + source = _write_source(src) + amax = tmp_path / "amax.json" + amax.write_text(json.dumps({f"b0.{name}": {"all": ACT_AMAX} for name in ("ff.fc_in", "attn.to_q", "attn.to_k")})) + subprocess.run([sys.executable, str(SCRIPT), "--src", str(src), "--dst", str(dst), "--quantize-attention", + "--act-amax", str(amax)], check=True) + + export = dst / nv.H3_NVFP4_DIT_EXPORT_FILENAME + with safe_open(str(export), framework="pt", device="cpu") as reader: + keys = set(reader.keys()) + # ModelOpt bytes are carried over bit for bit, including the attention projection it already quantized. + for prefix, module in (("transformer_blocks.0.ff.net.0.proj", "transformer_blocks.0.ff.fc_in"), + ("transformer_blocks.0.attn.to_k", "transformer_blocks.0.attn.to_k")): + assert torch.equal(reader.get_tensor(f"{module}::_nvfp4_weight"), source[f"{prefix}.weight"]) + assert "transformer_blocks.0.attn.to_q::_nvfp4_weight" in keys + dense_index = json.loads((dst / "diffusion_pytorch_model.safetensors.index.json").read_text())["weight_map"] + assert set(dense_index) == {"transformer_blocks.0.norm.weight"} + assert "quantization_config" not in json.loads((dst / "config.json").read_text()) + + model = _model() + assert nv.load_minimax_h3_nvfp4_dit_export(model, str(export), device="cuda") == 3 + x = torch.randn(256, IN, device="cuda").clamp(-ACT_AMAX, ACT_AMAX).to(torch.bfloat16) + references = { + "ff.fc_in": conv.dequantize_modelopt(*(source[f"transformer_blocks.0.ff.net.0.proj.{s}"].cuda() + for s in ("weight", "weight_scale", "weight_scale_2"))), + "attn.to_k": conv.dequantize_modelopt(*(source[f"transformer_blocks.0.attn.to_k.{s}"].cuda() + for s in ("weight", "weight_scale", "weight_scale_2"))), + "attn.to_q": source["transformer_blocks.0.attn.to_q.weight"].cuda(), + } + block = model.transformer_blocks[0] + for name, reference in references.items(): + layer = block.get_submodule(name) + assert layer._nvfp4_input_global_sf.item() == pytest.approx(448.0 * 6.0 / ACT_AMAX) + out = layer.quant_method.apply(layer, x).float() + ref = x.float() @ reference.float().T + error = ((out - ref).norm() / ref.norm()).item() + assert error < 0.2, f"{name}: relative error {error:.3f}" diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py index 8b2ff1b448..380c52e993 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -58,6 +58,10 @@ from safetensors.torch import save_file EXPORT_FILENAME = "nvfp4_weights.safetensors" +# ModelOpt keys intentionally not carried into the export. ``input_scale`` is ModelOpt's static +# activation scale; FastVideo either quantizes activations per call with a unit global scale or, +# with ``--act-amax``, stores its own calibrated ``_nvfp4_input_global_sf`` per linear. +DROPPED_MODELOPT_SUFFIXES = ("input_scale",) _BLOCK_ATTN = re.compile(r"^transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out\.0)$") _BLOCK_FFN = re.compile(r"^transformer_blocks\.\d+\.ff\.net\.(?:0\.proj|2)$") _BLOCK_GATE = re.compile(r"^transformer_blocks\.\d+\.attn\.to_gate_compress$") @@ -160,14 +164,18 @@ def main() -> None: index = json.loads((args.src / "diffusion_pytorch_model.safetensors.index.json").read_text()) weight_map: dict[str, str] = index["weight_map"] modelopt = sorted(k[:-len(".weight_scale_2")] for k in weight_map if k.endswith(".weight_scale_2")) - modelopt_keys = {f"{p}.{s}" for p in modelopt for s in ("weight", "weight_scale", "weight_scale_2", "input_scale")} + modelopt_keys = {f"{p}.{s}" for p in modelopt for s in ("weight", "weight_scale", "weight_scale_2", *DROPPED_MODELOPT_SUFFIXES)} if args.quantize_gate and not args.quantize_attention: parser.error("--quantize-gate requires --quantize-attention") selected = ([_BLOCK_ATTN.pattern] if args.quantize_attention else []) + ( [_BLOCK_GATE.pattern] if args.quantize_gate else []) + ([_BLOCK_FFN.pattern] if args.quantize_ffn else []) dense_pattern = re.compile("|".join(selected)) if selected else None + # Projections ModelOpt already quantized stay on the ModelOpt path; re-quantizing their packed + # uint8 weight as if it were dense BF16 would corrupt them. + modelopt_set = set(modelopt) attention = sorted(k[:-len(".weight")] for k in weight_map - if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")])) if dense_pattern else [] + if k.endswith(".weight") and dense_pattern.match(k[:-len(".weight")]) + and k[:-len(".weight")] not in modelopt_set) if dense_pattern else [] if args.quantize_ffn and any(_BLOCK_FFN.match(p) for p in modelopt): parser.error("--quantize-ffn needs a bf16 source; this one already holds ModelOpt FFN linears") amax_table = None From 74a190f9c85cc24fe6a3bbda0d7549f6d01570a8 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 21:10:53 -0700 Subject: [PATCH 050/107] [bugfix]: headline benchmarks record their decode mode; 4090 bench never reports unmeasured memory as 0 GiB bench_headline.py records HEADLINE_* switches and the resolved engine and experimental config in results.json and W&B. A slash in a Modal --tag no longer creates a nested run directory the local writer cannot write. When cgroup sampling fails, bench_pod.py reports host peaks as null with the error instead of 0.0 GiB. --- scripts/benchmarks/minimax_h3_4090/bench_pod.py | 13 ++++++++++--- scripts/benchmarks/minimax_h3_pro6000/app.py | 3 ++- .../benchmarks/minimax_h3_pro6000/bench_headline.py | 6 ++++-- 3 files changed, 16 insertions(+), 6 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index a905e1989b..b1f4dde625 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -23,6 +23,7 @@ def __init__(self): self.peak_bytes = 0 self.peak_anon_bytes = 0 self.peak_gpu_bytes = None + self.host_error: str | None = None self._gpu_used = None self._nvml_shutdown = None self.thread = threading.Thread(target=self._sample, daemon=True) @@ -44,7 +45,9 @@ def _sample(self): 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): + except (OSError, KeyError, ValueError) as exc: + self.host_error = f"{type(exc).__name__}: {exc}" + print(f"host memory sampling stopped: {self.host_error}", flush=True) return self.stop.wait(0.1) @@ -164,8 +167,12 @@ def main(): "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)}) + # None when sampling failed: an unmeasured run must not read as 0 GiB. + "peak_host_cgroup_gib": (round(host_peak.peak_bytes / 2**30, 3) + if host_peak.host_error is None else None), + "peak_host_anon_gib": (round(host_peak.peak_anon_bytes / 2**30, 3) + if host_peak.host_error is None else None), + "host_memory_error": host_peak.host_error}) timed = [run["wall_s"] for run in results["runs"] if not run["warmup"]] if timed: results["median_e2e_s"] = statistics.median(timed) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index 0db7b113a7..3ece05c8b3 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -558,7 +558,8 @@ def headline_fetch(repo: str) -> str: def _headline(repo: str, gpus: int, profile: str, extra_env: dict | None, tag: str = "") -> dict: _install_kernel() model = f"/vol/models/{repo.split('/')[-1]}" - run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}{tag}" + # The run name is a single path component locally and on the volume. + run_name = f"pro6000x{gpus}-{repo.split('/')[-1]}{tag}".replace("/", "-") env = {**os.environ, **FAST_ENV, **(extra_env or {}), "HEADLINE_OUT": "/vol/outputs/headline", "HEADLINE_DEVICE": f"{gpus}x RTX PRO 6000", "PYTHONPATH": "/src/fastvideo"} proc = subprocess.run(["python", "/root/bench_headline.py", run_name, model, str(gpus), profile, diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py index 669581d804..cd25102905 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -63,7 +63,8 @@ def main(): engine[key] = value experimental.update(json.loads(os.environ.get("HEADLINE_EXPERIMENTAL_JSON", "{}"))) config = {"model_path": a.model_dir, "engine": engine, "pipeline": {"experimental": experimental}} - env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA"))} + # HEADLINE_* switches (VAE_PARALLEL, ENGINE_JSON, ...) change the run, so record them with the rest. + env = {k: v for k, v in os.environ.items() if k.startswith(("FASTVIDEO_", "PYTORCH_CUDA", "HEADLINE_"))} run = None if os.environ.get("WANDB_PROJECT"): import wandb @@ -72,6 +73,7 @@ def main(): dir=out_dir, config={"model_dir": a.model_dir, "num_gpus": a.num_gpus, "nvfp4_profile": a.nvfp4_profile, "dmd_steps": steps, "vsa_sparsity": sparsity, "settings": a.settings, "engine": engine, + "experimental": experimental, "env": env, "device": os.environ.get("HEADLINE_DEVICE", "")}) from fastvideo import VideoGenerator @@ -79,7 +81,7 @@ def main(): generator = VideoGenerator.from_config(config) results = {"run_name": a.run_name, "model_dir": a.model_dir, "num_gpus": a.num_gpus, "nvfp4_profile": a.nvfp4_profile, "load_s": round(time.perf_counter() - t0, 1), "env": env, - "settings": {}} + "engine": engine, "experimental": experimental, "settings": {}} try: for name in a.settings.split(","): width, height, frames = SETTINGS[name] From 4784019824748d20b8e09d3ec3685c0cfa740148 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 21:48:23 -0700 Subject: [PATCH 051/107] [bugfix]: streamed NVFP4 encoder layers pass their FP4 scalars on the GEMM's device A layerwise-streamed encoder layer is finalized on the host, so its alpha and activation global scale stayed on the CPU while the packed weights visited the GPU. FlashInfer 0.6.18 rejects a CPU globalScale; both test_streamed_encoder_matches_resident_and_releases_layers NVFP4 cases failed on GB200 and now pass. --- .../models/encoders/minimax_h3_checkpoint_nvfp4.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 2d5d2897e6..8ffca2bae9 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -444,8 +444,14 @@ def _apply_finalized(layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor if x.numel() == 0: # An empty prompt has nothing to quantize; the FP4 kernels are not defined for zero rows. return x.new_zeros(*original_shape[:-1], layer.output_size_per_partition, dtype=torch.bfloat16) - x_fp4, x_scale = _quantize_activation_nvfp4(x.reshape(-1, original_shape[-1]), layer._nvfp4_x_global_scale) - output = _nvfp4_linear(x_fp4, x_scale, layer.weight_packed, layer.weight_scale, layer._nvfp4_alpha) + x_global_scale, alpha = layer._nvfp4_x_global_scale, layer._nvfp4_alpha + if alpha.device != layer.weight_packed.device: + # A streamed encoder layer is finalized on the host and only its packed weights visit the GPU; + # FlashInfer requires the scalars on the GEMM's device. + x_global_scale = x_global_scale.to(layer.weight_packed.device, non_blocking=True) + alpha = alpha.to(layer.weight_packed.device, non_blocking=True) + x_fp4, x_scale = _quantize_activation_nvfp4(x.reshape(-1, original_shape[-1]), x_global_scale) + output = _nvfp4_linear(x_fp4, x_scale, layer.weight_packed, layer.weight_scale, alpha) if bias is not None: # The GEMM emits bf16; keep it that way whatever dtype the bias was built in. output = output + bias.to(output.dtype) From 7863e08651e166fa72bea63fec26f97b159d9ed6 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 21:26:55 -0700 Subject: [PATCH 052/107] [bugfix]: preserve ModelOpt H3 activation calibration in NVFP4 export --- ...st_minimax_h3_modelopt_activation_scale.py | 48 +++++++++++++++++++ .../convert_minimax_h3_modelopt_nvfp4_dit.py | 23 ++++++--- 2 files changed, 64 insertions(+), 7 deletions(-) create mode 100644 fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py diff --git a/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py new file mode 100644 index 0000000000..7d34e346f4 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_minimax_h3_modelopt_activation_scale.py @@ -0,0 +1,48 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU regression for the ModelOpt activation-scale export contract.""" +import importlib.util +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + + +@pytest.fixture +def converter(monkeypatch): + path = Path(__file__).resolve().parents[4] / 'scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py' + spec = importlib.util.spec_from_file_location('h3_modelopt_converter', path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + packed = torch.zeros((2, 8), dtype=torch.uint8) + scales = torch.ones((2, 1), dtype=torch.float8_e4m3fn) + def quantize(*args, **kwargs): + return packed, scales + monkeypatch.setattr(module, '_flashinfer', lambda: ( + SimpleNamespace(layout_128x4=0), None, quantize, lambda value: value)) + return module + + +def convert(converter, input_scale): + return converter.convert_modelopt_linear( + torch.zeros((2, 8), dtype=torch.uint8), + torch.ones((2, 1), dtype=torch.float8_e4m3fn), + torch.tensor(0.25), 'cpu', input_scale=input_scale)[0] + + +def test_modelopt_activation_scale_is_preserved_as_reciprocal(converter): + buffers = convert(converter, torch.tensor(2.0)) + assert buffers['_nvfp4_input_global_sf'].dtype == torch.float32 + assert buffers['_nvfp4_input_global_sf'].shape == () + assert buffers['_nvfp4_input_global_sf'].item() == 0.5 + assert buffers['_nvfp4_alpha'].item() == 0.25 + + +def test_missing_modelopt_activation_scale_keeps_legacy_export(converter): + assert '_nvfp4_input_global_sf' not in convert(converter, None) + + +@pytest.mark.parametrize('value', [0.0, -1.0, float('nan'), float('inf')]) +def test_invalid_modelopt_activation_scale_is_rejected(converter, value): + with pytest.raises(ValueError, match='finite positive scalar'): + convert(converter, torch.tensor(value)) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py index 380c52e993..fa8e06f9f6 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -13,10 +13,11 @@ ``load_minimax_h3_nvfp4_dit_export``) stores ``::_nvfp4_weight`` (same bytes), ``::_nvfp4_weight_scale`` (the same E4M3 bytes in FlashInfer's 128x4 swizzled layout), ``::_nvfp4_alpha`` (= ``weight_scale_2``) and -``::_weight_global_sf`` (= 1 / ``weight_scale_2``). The calibrated weights are -therefore carried over bit for bit; only the activation scale changes, because -FastVideo quantizes activations per call with a unit global scale and the -static ``input_scale`` is dropped. +``::_weight_global_sf`` (= 1 / ``weight_scale_2``), and +``::_nvfp4_input_global_sf`` (= 1 / ``input_scale``). The calibrated weight +bytes and activation scale are preserved. Dropping the activation scale would +replace its calibrated range with a unit global scale, clipping inputs above +2688. ``--quantize-attention`` additionally quantizes the dense BF16 attention projections (``attn.to_{q,k,v,out}``) of every main block exactly as @@ -109,7 +110,8 @@ def probe(buffers: dict[str, torch.Tensor], reference: torch.Tensor, rows: int = return ((out.float() - ref).norm() / ref.norm()).item() -def convert_modelopt_linear(weight, scale, scale_2, device) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: +def convert_modelopt_linear(weight, scale, scale_2, device, input_scale=None + ) -> tuple[dict[str, torch.Tensor], torch.Tensor, float]: """Carry the calibrated bytes over; return (buffers, dequantized weight, scale-byte agreement). The agreement compares the swizzled ModelOpt scales with the scales @@ -132,6 +134,11 @@ def convert_modelopt_linear(weight, scale, scale_2, device) -> tuple[dict[str, t "_nvfp4_alpha": scale_2.clone(), "_weight_global_sf": (1.0 / scale_2).to(torch.bfloat16), } + if input_scale is not None: + value = input_scale.to(dtype=torch.float32) + if value.numel() != 1 or not bool(torch.isfinite(value).all()) or not bool((value > 0).all()): + raise ValueError("ModelOpt input_scale must be a finite positive scalar") + buffers["_nvfp4_input_global_sf"] = value.reshape(()).reciprocal().to(device=device) return buffers, reference, agreement @@ -197,7 +204,9 @@ def main() -> None: else: buffers, reference, agreement = convert_modelopt_linear(tensor(f"{prefix}.weight"), tensor(f"{prefix}.weight_scale"), - tensor(f"{prefix}.weight_scale_2"), device) + tensor(f"{prefix}.weight_scale_2"), device, + input_scale=tensor(f"{prefix}.input_scale") + if f"{prefix}.input_scale" in weight_map else None) agreements.append(agreement) error = probe(buffers, reference) worst = max(worst, error) @@ -241,7 +250,7 @@ def main() -> None: shutil.copy2(extra, args.dst / extra.name) print(json.dumps({"exported_linears": len(modelopt) + len(attention), "modelopt_linears": len(modelopt), "quantized_dense_linears": len(attention), "worst_probe_error": round(worst, 4), - "static_activation_scales": len(attention) + len(modelopt) if amax_table else 0, + "static_activation_scales": sum(k.endswith("::_nvfp4_input_global_sf") for k in export), "gate_linears": sum(1 for p in attention if _BLOCK_GATE.match(p)), "min_scale_byte_agreement": round(min(agreements), 4) if agreements else None, "mean_scale_byte_agreement": round(sum(agreements) / len(agreements), 4) if agreements else None, From 9ca4eba8bbd461582d058bcae3bf171408c55cba Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:21:08 -0700 Subject: [PATCH 053/107] [bugfix]: declare NVFP4 static activation-scale cache attributes for mypy --- fastvideo/layers/quantization/nvfp4_config.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index ded48937fa..abe370bd59 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -380,6 +380,10 @@ def _load_amax_table(path: str) -> dict[str, float]: class NVFP4QuantizeMethod(QuantizeMethodBase): + # Lazily resolved by _static_activation_global_sf; class defaults also cover object.__new__ test doubles. + _static_sf_checked: bool = False + _static_sf: torch.Tensor | None = None + def __init__(self, layer_prefix: str = ""): super().__init__() self.weight_fp4 = None From 3c7e89c588a617f07b5940ec575df95686ce31c9 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:21:08 -0700 Subject: [PATCH 054/107] [bugfix]: clone batched VAE tile output before the next CUDA-graph replay Under the opt-in reduce-overhead decoder compile, each batch's output lives in the graph pool; retaining split views let the next batch overwrite earlier tiles before stitching. --- fastvideo/models/vaes/minimax_h3_video.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index 223aae259e..044dbecce2 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -843,7 +843,10 @@ def _decode_clip(self, z: torch.Tensor) -> torch.Tensor: per_call = _tile_batch_size() for start in range(0, len(latent_tiles), per_call): batch = torch.cat(latent_tiles[start:start + per_call], dim=0) - decoded.extend(self.decoder(self._project_decoder_tile(batch)).split(z.shape[0], dim=0)) + # Clone before retaining: under the reduce-overhead compile the decoder output lives in + # a CUDA-graph pool that the next batch's replay overwrites. + out = self.decoder(self._project_decoder_tile(batch)).clone() + decoded.extend(out.split(z.shape[0], dim=0)) columns = len(x_indices) rows = [decoded[index:index + columns] for index in range(0, len(decoded), columns)] # The eager tile driver owns NVTX so each marker remains outside From f4df26cc26e381c33a66c599fc36fa40eafd2406 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:21:08 -0700 Subject: [PATCH 055/107] [bugfix]: VSA zero-gate guard checks packed NVFP4 gates With layer_profile h3_dit_vsa the gate weight is purged and lives in _nvfp4_weight, so an all-zero packed gate skipped the guard. Codes that are +-0 in every nibble now count as a zero gate. --- fastvideo/pipelines/basic/minimax_h3/vsa_guard.py | 9 +++++++++ .../tests/stages/test_minimax_h3_vsa_guard.py | 14 ++++++++++++++ 2 files changed, 23 insertions(+) diff --git a/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py index bdb3f8e546..5820252acd 100644 --- a/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py +++ b/fastvideo/pipelines/basic/minimax_h3/vsa_guard.py @@ -21,6 +21,15 @@ def refuse_zero_initialized_h3_vsa(transformer: Any) -> None: if gate is None: continue weight = getattr(gate, "weight", None) + packed = getattr(gate, "_nvfp4_weight", None) + if (weight is None or weight.numel() == 0) and packed is not None and packed.numel() > 0: + # h3_dit_vsa: the gate is packed NVFP4 (weight is None). Two E2M1 codes per byte; the low three + # bits of each nibble are the magnitude, so a gate whose codes are all +-0 is a zero gate. + saw_weight = True + if bool(((packed & 0x77) != 0).any()): + any_trained = True + break + continue # Layerwise offload leaves a zero-element placeholder on the module; it says nothing about the gate. if weight is None or weight.numel() == 0: continue diff --git a/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py index c4c44a4931..655909f340 100644 --- a/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py +++ b/fastvideo/tests/stages/test_minimax_h3_vsa_guard.py @@ -72,3 +72,17 @@ def loader() -> nn.Module: refuse_zero_initialized_h3_vsa(LazyModule("transformer", loader)) assert loads == [] + + +def _packed_gate_transformer(codes: int) -> SimpleNamespace: + """An h3_dit_vsa gate: ``weight`` purged, two E2M1 codes per ``_nvfp4_weight`` byte.""" + gate = SimpleNamespace(weight=None, _nvfp4_weight=torch.full((4, 2), codes, dtype=torch.uint8)) + return SimpleNamespace(transformer_blocks=[SimpleNamespace(attn=SimpleNamespace(to_gate_compress=gate))]) + + +def test_refuse_zero_initialized_h3_vsa_checks_packed_nvfp4_gates() -> None: + # 0x88 is -0 in both nibbles: still a zero gate. + for zero_codes in (0x00, 0x88): + with pytest.raises(RuntimeError, match="to_gate_compress"): + refuse_zero_initialized_h3_vsa(_packed_gate_transformer(zero_codes)) + refuse_zero_initialized_h3_vsa(_packed_gate_transformer(0x12)) From 2964f151f3ce677e4f3c7ffd4b2ba7a0fc91ef9a Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:21:08 -0700 Subject: [PATCH 056/107] [bugfix]: sparse FP4 entry points validate block lists by default sageattn_blackwell_sparse{,_bshd} now check q2k_num / q2k_idx unless the caller passes validate=False. The H3 VSA path builds its lists with vsa_tile_mask_to_fp4_blocks and opts out to avoid the host sync. The converter comment now says input_scale is carried, not dropped. --- fastvideo-kernel/attn_qat_infer/api.py | 8 +++++--- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 7 +++++-- .../convert_minimax_h3_modelopt_nvfp4_dit.py | 6 +++--- 3 files changed, 13 insertions(+), 8 deletions(-) diff --git a/fastvideo-kernel/attn_qat_infer/api.py b/fastvideo-kernel/attn_qat_infer/api.py index 429d869c7d..e91196a82f 100644 --- a/fastvideo-kernel/attn_qat_infer/api.py +++ b/fastvideo-kernel/attn_qat_infer/api.py @@ -303,14 +303,16 @@ def sageattn_blackwell_sparse(q, per_block_mean=True, single_level_p_quant=True, sm_scale: float | None = None, - validate: bool = False): + validate: bool = True): """Block-sparse SageAttention3 FP4 forward (non-causal). Query block ``m`` (``BLOCK_M`` rows) of each (batch, head) attends only to the ``BLOCK_N``-token KV blocks in ``q2k_idx[b, h, m, :q2k_num[b, h, m]]``, restricted to the quadrants in ``q2k_quad`` when given; see :func:`vsa_tile_mask_to_fp4_blocks`. Q/K/V are ``[B, H, L, D]``. - ``validate=True`` runs :func:`check_sparse_block_lists` first. + Block lists are checked with :func:`check_sparse_block_lists` (a host + sync) unless ``validate=False``; pass that only for lists built by + :func:`vsa_tile_mask_to_fp4_blocks`, which are in range by construction. """ QL = q.size(2) KL = k.size(2) @@ -337,7 +339,7 @@ def sageattn_blackwell_sparse_bshd(q, q2k_quad: torch.Tensor | None = None, single_level_p_quant=True, sm_scale: float | None = None, - validate: bool = False) -> torch.Tensor: + validate: bool = True) -> torch.Tensor: """:func:`sageattn_blackwell_sparse` for ``[B, L, H, D]`` inputs, without copies. The FP4 quantizers read strided input, so the sequence-major tensors a diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index ee13941824..b3b7464526 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -234,7 +234,9 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ meta.span_sparsities) q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) with STAGES.span("fp4_attention"): - out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) + # Lists come from vsa_tile_mask_to_fp4_blocks and are in range; skip the per-call host sync. + out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad, + validate=False) with STAGES.span("out_untile"): out = out.transpose(1, 2).index_select(1, layout.untile) # [B, L, H, D], packed order if sim_fp8: @@ -401,7 +403,8 @@ def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tup meta.span_sparsities) q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) with STAGES.span("fp4_attention"): - out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad) + out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad, + validate=False) with STAGES.span("out_pack"): payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py index fa8e06f9f6..463a8371c8 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py @@ -59,9 +59,9 @@ from safetensors.torch import save_file EXPORT_FILENAME = "nvfp4_weights.safetensors" -# ModelOpt keys intentionally not carried into the export. ``input_scale`` is ModelOpt's static -# activation scale; FastVideo either quantizes activations per call with a unit global scale or, -# with ``--act-amax``, stores its own calibrated ``_nvfp4_input_global_sf`` per linear. +# ModelOpt keys not copied as dense tensors. ``input_scale`` (ModelOpt's static activation scale) +# is carried into the packed export as ``_nvfp4_input_global_sf`` = 1 / input_scale, unless +# ``--act-amax`` supplies a replacement calibration. DROPPED_MODELOPT_SUFFIXES = ("input_scale",) _BLOCK_ATTN = re.compile(r"^transformer_blocks\.\d+\.attn\.(?:to_q|to_k|to_v|to_out\.0)$") _BLOCK_FFN = re.compile(r"^transformer_blocks\.\d+\.ff\.net\.(?:0\.proj|2)$") From 826b5d0ddd98ea921b4fcc5f2c86985a48d7270e Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 08:26:44 -0700 Subject: [PATCH 057/107] [misc]: register FastH3 single-GPU switches in fastvideo.envs Upstream #1898 requires every FASTVIDEO_* variable to be declared in fastvideo/envs.py and read with envs.NAME.get(). The 29 H3, NVFP4, offload and memory-debug switches this branch added read os.environ directly and tests set them with monkeypatch.setenv, which fails test_env_access_follows_policy in the unit lane. Declare them with types and defaults matching the old parsing, read them through the registry, switch tests to envs.NAME.override(), and regenerate the env-var table. --- docs/contributing/env_vars.md | 29 +++++ .../backends/video_sparse_attn_h3.py | 3 +- fastvideo/envs.py | 115 ++++++++++++++++++ fastvideo/hooks/layerwise_offload.py | 10 +- fastvideo/layers/quantization/nvfp4_config.py | 12 +- fastvideo/models/dits/minimax_h3.py | 9 +- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 13 +- .../encoders/minimax_h3_checkpoint_nvfp4.py | 5 +- fastvideo/models/loader/component_loader.py | 12 +- fastvideo/models/loader/fsdp_load.py | 3 +- .../models/vaes/minimax_h3_int8_convrot.py | 6 +- fastvideo/models/vaes/minimax_h3_video.py | 7 +- .../basic/minimax_h3/minimax_h3_pipeline.py | 7 +- .../test_minimax_h3_encoder_layerwise.py | 9 +- fastvideo/tests/hooks/test_pinned_memory.py | 6 +- .../ssim/.reference_videos_download.lock | 0 .../test_minimax_h3_sequential_start.py | 10 +- .../test_minimax_h3_tile_first.py | 18 +-- .../tests/vaes/test_minimax_h3_int8_shared.py | 12 +- fastvideo/worker/gpu_worker.py | 9 +- 20 files changed, 229 insertions(+), 66 deletions(-) create mode 100644 fastvideo/tests/ssim/.reference_videos_download.lock diff --git a/docs/contributing/env_vars.md b/docs/contributing/env_vars.md index 4f47ddcb1f..883e2d8ac8 100644 --- a/docs/contributing/env_vars.md +++ b/docs/contributing/env_vars.md @@ -240,6 +240,35 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_VBENCH_FULL_INFO_JSON` | str | unset | eval | Path to VBench_full_info.json, used instead of the vendored copy. Deprecated names: `VBENCH_FULL_INFO_JSON`. | | `FASTVIDEO_FVD_REF_FEATURES` | str | unset | eval | Cached reference-feature file for the FVD metric. | | `FASTVIDEO_FAD_REF_FEATURES` | str | unset | eval | Cached reference-feature file for the audio Frechet distance metric. | +| `FASTVIDEO_H3_VSA_FP4` | bool | `0` | attention | Run MiniMax-H3 VSA attention on the block-sparse SageAttention3 FP4 kernel (sm_120, no-grad, single sequence-parallel rank). | +| `FASTVIDEO_H3_VSA_TILE_FIRST` | bool | `0` | attention | Single-rank MiniMax-H3 VSA with one tile gather of the block input instead of separate Q/K/V/gate scatters. | +| `FASTVIDEO_H3_VSA_SM89_KERNEL` | one of original, bf16, int8 | `original` | attention | Fine-attention kernel for MiniMax-H3 VSA on sm_89: original, bf16, or int8 (INT8 QK, BF16 PV). | +| `FASTVIDEO_H3_SIM_SP_FP8` | bool | `0` | debug | Simulate the FP8 sequence-parallel exchange of the MiniMax-H3 FP4 VSA path on one rank. | +| `FASTVIDEO_H3_FFN_CHUNK_TOKENS` | int | `0` | performance | Inference-only MiniMax-H3 FFN token chunk size; 0 runs the FFN unchunked. | +| `FASTVIDEO_H3_FP8_ATTENTION` | bool | `0` | performance | With NVFP4 layer_profile h3_dit_ffn, run MiniMax-H3 attention projections in FP8. | +| `FASTVIDEO_H3_FP8_GRANULARITY` | one of tensor, channel | `tensor` | performance | FP8 scaling granularity for FASTVIDEO_H3_FP8_ATTENTION. | +| `FASTVIDEO_NVFP4_MM_BACKEND` | str | `auto` | performance | FlashInfer mm_fp4 backend for NVFP4 linears, e.g. auto or cutlass. | +| `FASTVIDEO_NVFP4_ACT_AMAX` | path | unset | performance | JSON of calibrated NVFP4 input amax per linear, keyed b<block>.<sub> or full prefix; sets a static activation scale. | +| `FASTVIDEO_NVFP4_DYNAMIC_ACT` | str | `""` | performance | NVFP4 linears that derive the activation scale per call: all, or comma-separated layer-name suffixes such as ff.fc_out. | +| `FASTVIDEO_H3_ADALN_CACHE` | bool | `0` | performance | Cache MiniMax-H3 AdaLN modulation per timestep instead of keeping the projection weights resident. | +| `FASTVIDEO_H3_ADALN_TABLE` | path | unset | performance | Precomputed MiniMax-H3 AdaLN modulation table; enables the cache and skips loading the AdaLN projection weights. | +| `FASTVIDEO_H3_ADALN_DUMP` | path | unset | debug | Write the MiniMax-H3 AdaLN modulation table to this path while sampling. | +| `FASTVIDEO_H3_SPLICE_TRANSFORMER` | path | unset | eval | Second MiniMax-H3 transformer that runs the late DMD steps (checkpoint step-splice evaluation). | +| `FASTVIDEO_H3_SPLICE_FROM_STEP` | int | `4` | eval | First denoising step run by FASTVIDEO_H3_SPLICE_TRANSFORMER. | +| `FASTVIDEO_H3_ENCODER_LAYERWISE` | bool | `0` | performance | Stream MiniMax-H3 text-encoder language layers through exact-size pinned host memory (text-only prompts). | +| `FASTVIDEO_H3_ENCODER_FUSED_DEQUANT` | bool | `0` | performance | Expand the serialized NVFP4 MiniMax-H3 text encoder with one fused Triton pass on GPUs without FP4 GEMM. | +| `FASTVIDEO_H3_VAE_TILE_BATCH` | int | `1` | performance | Spatial tiles per MiniMax-H3 video VAE decoder call; 1 decodes per tile. | +| `FASTVIDEO_H3_VAE_INT8_SHARED_QKV` | bool | `0` | performance | Share the INT8 activation rotation and quantization across the MiniMax-H3 VAE Q/K/V projections. | +| `FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW` | bool | `0` | performance | Use transposed weight views in the MiniMax-H3 VAE INT8 projections. | +| `FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT` | bool | `0` | performance | Fused dequantization epilogue for the MiniMax-H3 VAE INT8 projections. | +| `FASTVIDEO_H3_PINNED_SWAP` | bool | `1` | performance | Swap offloaded MiniMax-H3 modules through exact-size pinned host arenas. | +| `FASTVIDEO_H3_PARK_MODULES` | str | unset | performance | Comma-separated MiniMax-H3 denoise modules parked on the host while the text encoder runs, e.g. vae,audio_vae. | +| `FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS` | bool | `0` | performance | Layerwise offload also streams large buffers such as packed FP4/FP8 weights. | +| `FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS` | int | `0` | performance | Keep the first N layerwise-offloaded blocks resident on the GPU. | +| `FASTVIDEO_H3_SP_PROFILE` | bool | `0` | profiling | CUDA-event spans per stage over one MiniMax-H3 FP4 VSA DiT forward. | +| `FASTVIDEO_H3_CAPTURE_QKV` | path | unset | debug | Directory for captured real MiniMax-H3 Q/K/V attention inputs. | +| `FASTVIDEO_CUDA_MEMORY_CAP_GIB` | float | `0.0` | debug | Cap this process's CUDA allocator at this many GiB to emulate a smaller GPU; 0 leaves it uncapped. | +| `FASTVIDEO_MEMORY_REPORT` | bool | `0` | debug | Log bytes held per pipeline component by device and dtype after loading. | | `FASTVIDEO_TEST_LTX2_OVERFIT_DATA_DIR` | str | `data/cats` | test | Raw data directory for preprocess_ltx2_overfit.py. Deprecated names: `LTX2_OVERFIT_DATA_DIR`. | | `FASTVIDEO_TEST_LTX2_OVERFIT_CAPTION_JSON` | str | `videos2caption_1_sample.json` | test | Caption file, relative to the raw data directory. Deprecated names: `LTX2_OVERFIT_CAPTION_JSON`. | | `FASTVIDEO_TEST_LTX2_OVERFIT_VIDEO_SUBDIR` | str | `video` | test | Video subdirectory, relative to the raw data directory. Deprecated names: `LTX2_OVERFIT_VIDEO_SUBDIR`. | diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index a4ea213206..0ed8e9a647 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -61,7 +61,6 @@ import functools import math -import os from dataclasses import dataclass from typing import Any @@ -548,7 +547,7 @@ 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") + self._sm89_kernel = envs.FASTVIDEO_H3_VSA_SM89_KERNEL.get() 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 diff --git a/fastvideo/envs.py b/fastvideo/envs.py index 8962a8e703..77ba871a67 100644 --- a/fastvideo/envs.py +++ b/fastvideo/envs.py @@ -516,6 +516,121 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="eval", doc="Cached reference-feature file for the audio Frechet distance metric.") +# ================== MiniMax-H3 single-GPU switches ================== + +FASTVIDEO_H3_VSA_FP4 = EnvBool(False, + category="attention", + doc="Run MiniMax-H3 VSA attention on the block-sparse SageAttention3 FP4 kernel " + "(sm_120, no-grad, single sequence-parallel rank).") +FASTVIDEO_H3_VSA_TILE_FIRST = EnvBool(False, + category="attention", + doc="Single-rank MiniMax-H3 VSA with one tile gather of the block input " + "instead of separate Q/K/V/gate scatters.") +FASTVIDEO_H3_VSA_SM89_KERNEL = EnvChoice("original", + choices=("original", "bf16", "int8"), + category="attention", + doc="Fine-attention kernel for MiniMax-H3 VSA on sm_89: original, bf16, " + "or int8 (INT8 QK, BF16 PV).") +FASTVIDEO_H3_SIM_SP_FP8 = EnvBool(False, + category="debug", + doc="Simulate the FP8 sequence-parallel exchange of the MiniMax-H3 FP4 VSA " + "path on one rank.") +FASTVIDEO_H3_FFN_CHUNK_TOKENS = EnvInt(0, + category="performance", + doc="Inference-only MiniMax-H3 FFN token chunk size; 0 runs the FFN " + "unchunked.") +FASTVIDEO_H3_FP8_ATTENTION = EnvBool(False, + category="performance", + doc="With NVFP4 layer_profile h3_dit_ffn, run MiniMax-H3 attention " + "projections in FP8.") +FASTVIDEO_H3_FP8_GRANULARITY = EnvChoice("tensor", + choices=("tensor", "channel"), + category="performance", + doc="FP8 scaling granularity for FASTVIDEO_H3_FP8_ATTENTION.") +FASTVIDEO_NVFP4_MM_BACKEND = EnvStr("auto", + category="performance", + doc="FlashInfer mm_fp4 backend for NVFP4 linears, e.g. auto or cutlass.") +FASTVIDEO_NVFP4_ACT_AMAX = EnvPath(None, + category="performance", + doc="JSON of calibrated NVFP4 input amax per linear, keyed b. or " + "full prefix; sets a static activation scale.") +FASTVIDEO_NVFP4_DYNAMIC_ACT = EnvStr("", + category="performance", + doc="NVFP4 linears that derive the activation scale per call: all, or " + "comma-separated layer-name suffixes such as ff.fc_out.") +FASTVIDEO_H3_ADALN_CACHE = EnvBool(False, + category="performance", + doc="Cache MiniMax-H3 AdaLN modulation per timestep instead of keeping the " + "projection weights resident.") +FASTVIDEO_H3_ADALN_TABLE = EnvPath(None, + category="performance", + doc="Precomputed MiniMax-H3 AdaLN modulation table; enables the cache and " + "skips loading the AdaLN projection weights.") +FASTVIDEO_H3_ADALN_DUMP = EnvPath(None, + category="debug", + doc="Write the MiniMax-H3 AdaLN modulation table to this path while sampling.") +FASTVIDEO_H3_SPLICE_TRANSFORMER = EnvPath(None, + category="eval", + doc="Second MiniMax-H3 transformer that runs the late DMD steps " + "(checkpoint step-splice evaluation).") +FASTVIDEO_H3_SPLICE_FROM_STEP = EnvInt(4, + category="eval", + doc="First denoising step run by FASTVIDEO_H3_SPLICE_TRANSFORMER.") +FASTVIDEO_H3_ENCODER_LAYERWISE = EnvBool(False, + category="performance", + doc="Stream MiniMax-H3 text-encoder language layers through exact-size " + "pinned host memory (text-only prompts).") +FASTVIDEO_H3_ENCODER_FUSED_DEQUANT = EnvBool(False, + category="performance", + doc="Expand the serialized NVFP4 MiniMax-H3 text encoder with one " + "fused Triton pass on GPUs without FP4 GEMM.") +FASTVIDEO_H3_VAE_TILE_BATCH = EnvInt(1, + category="performance", + doc="Spatial tiles per MiniMax-H3 video VAE decoder call; 1 decodes per " + "tile.") +FASTVIDEO_H3_VAE_INT8_SHARED_QKV = EnvBool(False, + category="performance", + doc="Share the INT8 activation rotation and quantization across the " + "MiniMax-H3 VAE Q/K/V projections.") +FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW = EnvBool(False, + category="performance", + doc="Use transposed weight views in the MiniMax-H3 VAE INT8 " + "projections.") +FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT = EnvBool(False, + category="performance", + doc="Fused dequantization epilogue for the MiniMax-H3 VAE INT8 " + "projections.") +FASTVIDEO_H3_PINNED_SWAP = EnvBool(True, + category="performance", + doc="Swap offloaded MiniMax-H3 modules through exact-size pinned host " + "arenas.") +FASTVIDEO_H3_PARK_MODULES = EnvStr(None, + category="performance", + doc="Comma-separated MiniMax-H3 denoise modules parked on the host while the " + "text encoder runs, e.g. vae,audio_vae.") +FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS = EnvBool(False, + category="performance", + doc="Layerwise offload also streams large buffers such as packed " + "FP4/FP8 weights.") +FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS = EnvInt(0, + category="performance", + doc="Keep the first N layerwise-offloaded blocks resident on the " + "GPU.") +FASTVIDEO_H3_SP_PROFILE = EnvBool(False, + category="profiling", + doc="CUDA-event spans per stage over one MiniMax-H3 FP4 VSA DiT forward.") +FASTVIDEO_H3_CAPTURE_QKV = EnvPath(None, + category="debug", + doc="Directory for captured real MiniMax-H3 Q/K/V attention inputs.") +FASTVIDEO_CUDA_MEMORY_CAP_GIB = EnvFloat(0.0, + category="debug", + doc="Cap this process's CUDA allocator at this many GiB to emulate a " + "smaller GPU; 0 leaves it uncapped.") +FASTVIDEO_MEMORY_REPORT = EnvBool(False, + category="debug", + doc="Log bytes held per pipeline component by device and dtype after " + "loading.") + # ================== Tests ================== FASTVIDEO_TEST_LTX2_OVERFIT_DATA_DIR = EnvStr("data/cats", diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 6d33327464..5a02af45b3 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 +import fastvideo.envs as envs from fastvideo.hooks.pinned_memory import PinnedTensorArena from fastvideo.logger import init_logger @@ -34,7 +35,7 @@ def _offload_tensors(module: nn.Module, names: dict[str, torch.Tensor] | None = yield name, tensor return yield from module.named_parameters() - if os.environ.get("FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS") == "1": + if envs.FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS.get(): for name, buf in module.named_buffers(): if buf is not None and buf.numel() * buf.element_size() >= _BUFFER_OFFLOAD_MIN_BYTES: yield name, buf @@ -191,10 +192,9 @@ def enable_layerwise_offload(model: nn.Module, 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 = max(0, envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.get()) + except ValueError as error: + logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS: %s", error) resident = 0 for name, submodule in model.named_children(): if isinstance(submodule, nn.ModuleList): diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index abe370bd59..3e683f749c 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -28,6 +28,8 @@ import torch.nn.functional as F from torch.nn.parameter import Parameter +import fastvideo.envs as envs + from fastvideo.layers.quantization.base_config import ( QuantizationConfig, QuantizeMethodBase, @@ -347,7 +349,7 @@ def _mm_fp4_backend() -> str: ``cudnn`` once activations reach tens of thousands of rows (measured at 73k rows on an RTX PRO 6000); short sequences are unaffected. """ - return os.environ.get("FASTVIDEO_NVFP4_MM_BACKEND", "auto") + return envs.FASTVIDEO_NVFP4_MM_BACKEND.get() def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor: @@ -414,7 +416,7 @@ def _static_activation_global_sf(self) -> torch.Tensor | None: if getattr(self, "_static_sf_checked", False): return self._static_sf self._static_sf_checked, self._static_sf = True, None - path = os.environ.get("FASTVIDEO_NVFP4_ACT_AMAX") + path = envs.FASTVIDEO_NVFP4_ACT_AMAX.get() if path: table = _load_amax_table(path) prefix = self.layer_prefix or "" @@ -429,7 +431,7 @@ def _dynamic_activation_scale(self) -> bool: """FASTVIDEO_NVFP4_DYNAMIC_ACT: "all", or comma-separated layer-name suffixes (e.g. "ff.fc_out").""" cached = getattr(self, "_dynamic_act_cached", None) if cached is None: - selected = os.environ.get("FASTVIDEO_NVFP4_DYNAMIC_ACT", "") + selected = envs.FASTVIDEO_NVFP4_DYNAMIC_ACT.get() suffixes = [part.strip() for part in selected.split(",") if part.strip()] prefix = self.layer_prefix or "" cached = "all" in suffixes or any(prefix.endswith(suffix) for suffix in suffixes) @@ -618,11 +620,11 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str): method = NVFP4QuantizeMethod(layer_prefix=prefix) method._retain_original_weights = self.retain_original_weights return method - if (self.layer_profile == "h3_dit_ffn" and os.environ.get("FASTVIDEO_H3_FP8_ATTENTION") == "1" + if (self.layer_profile == "h3_dit_ffn" and envs.FASTVIDEO_H3_FP8_ATTENTION.get() and _H3_BLOCK_ATTN_PROJ.search(prefix) is not None): # Mixed precision: NVFP4 MLPs, FP8 (per-tensor weight, dynamic per-tensor activation) attention. from fastvideo.layers.quantization.fp8_config import FP8QuantizeMethod - return FP8QuantizeMethod(granularity=os.environ.get("FASTVIDEO_H3_FP8_GRANULARITY", "tensor")) + return FP8QuantizeMethod(granularity=envs.FASTVIDEO_H3_FP8_GRANULARITY.get()) return None diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index bf8b226be1..dd9a645a1e 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -5,7 +5,6 @@ import contextlib import json -import os import math from collections.abc import Iterator, Mapping from typing import Any @@ -142,7 +141,7 @@ def __init__( self.fc_out.quant_method, MXFP8QuantizeMethod) # Inference-only token chunking: the 2 * ffn_dim intermediate is ~5.3x the block input # (4.5 GiB at 78k tokens), so chunks bound the activation peak on 24-32 GB GPUs. - self.chunk_tokens = int(os.environ.get("FASTVIDEO_H3_FFN_CHUNK_TOKENS", "0")) + self.chunk_tokens = envs.FASTVIDEO_H3_FFN_CHUNK_TOKENS.get() def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: tokens = hidden_states.shape[-2] if hidden_states.dim() > 1 else 0 @@ -242,7 +241,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._vsa_tile_first = use_vsa and envs.FASTVIDEO_H3_VSA_TILE_FIRST.get() 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. @@ -536,7 +535,7 @@ def forward(self, temb: torch.Tensor) -> tuple[torch.Tensor, ...]: out = F.linear(x.to(weight.dtype), weight, bias) if self._cache_key is not None: cache[self._cache_key] = out - if os.environ.get("FASTVIDEO_H3_ADALN_DUMP"): + if envs.FASTVIDEO_H3_ADALN_DUMP.get(): # Projection inputs, kept only when dumping, for offline low-rank fits. self.__dict__.setdefault("_modulation_inputs", {})[self._cache_key] = x.detach() return out.view(-1, 6 * self.hidden_size).chunk(6, dim=-1) @@ -1049,7 +1048,7 @@ def _move_adaln_tables(self, device: torch.device) -> None: cache[key] = value.to(device) def _maybe_dump_adaln_tables(self) -> None: - path = os.environ.get("FASTVIDEO_H3_ADALN_DUMP") + path = envs.FASTVIDEO_H3_ADALN_DUMP.get() if not path: return entries = sum(len(b.adaln_proj._modulation_cache) for b in self.transformer_blocks) diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index b3b7464526..3f3a48ef5b 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -19,14 +19,14 @@ from __future__ import annotations import math -import os from typing import Any import torch +import fastvideo.envs as envs + from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAMetadata, _build_block_mask, _pool_tiles) -VSA_FP4_ENV = "FASTVIDEO_H3_VSA_FP4" _BLOCK = 128 @@ -38,7 +38,7 @@ class _StageTimer: """ def __init__(self) -> None: - self.enabled = os.environ.get("FASTVIDEO_H3_SP_PROFILE", "0") == "1" + self.enabled = envs.FASTVIDEO_H3_SP_PROFILE.get() self._spans: list[tuple[str, torch.cuda.Event, torch.cuda.Event]] = [] def span(self, name: str): @@ -75,7 +75,7 @@ def flush(self) -> dict[str, float]: def vsa_fp4_requested() -> bool: - return os.environ.get(VSA_FP4_ENV, "0") == "1" + return envs.FASTVIDEO_H3_VSA_FP4.get() def _api() -> Any: @@ -172,7 +172,7 @@ def vsa_tile_first_attention(attn: Any, hidden_states: torch.Tensor, 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") + capture_root = envs.FASTVIDEO_H3_CAPTURE_QKV.get() if capture_root and attn._layer_idx in (0, 20, 41): from pathlib import Path root = Path(capture_root) @@ -219,7 +219,7 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ query = attn._apply_rotary_emb(attn.norm_q(query), rope) key = attn._apply_rotary_emb(attn.norm_k(key), rope) - sim_fp8 = os.environ.get(SIM_SP_FP8_ENV, "0") == "1" + sim_fp8 = envs.FASTVIDEO_H3_SIM_SP_FP8.get() if sim_fp8: query, key, value = (_fp8_roundtrip(t) for t in (query, key, value)) @@ -270,7 +270,6 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ # Debug: apply the SP path's FP8 rounding on one GPU (q/k/v after RoPE and the # attention output, one scale per token and head), to separate exchange # rounding from sharding errors when comparing SP>1 against SP=1. -SIM_SP_FP8_ENV = "FASTVIDEO_H3_SIM_SP_FP8" @torch.compile(dynamic=True, fullgraph=True) diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 8ffca2bae9..933041e88d 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -37,13 +37,14 @@ content checks a copied tensor can still fail. """ -import os from typing import Any import torch from torch import nn from torch.nn.parameter import Parameter +import fastvideo.envs as envs + from fastvideo.distributed import get_tp_world_size from fastvideo.layers.linear import LinearBase, LinearMethodBase from fastvideo.layers.quantization.base_config import QuantizationConfig @@ -422,7 +423,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._nvfp4_fused_dequant = envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.get() 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), diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index ab7bf07715..66a1d2dceb 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -13,6 +13,8 @@ import torch import torch.distributed as dist + +import fastvideo.envs as envs import torch.nn as nn from safetensors.torch import load_file as safetensors_load_file, safe_open from torch.distributed import init_device_mesh @@ -468,7 +470,7 @@ def load_model( 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 envs.FASTVIDEO_H3_ENCODER_LAYERWISE.get() and callable(prepare_layerwise): if target_device.type != "cpu": raise ValueError("Layerwise H3 encoder requires text_encoder_cpu_offload=True") prepare_layerwise(runtime_device) @@ -1153,9 +1155,9 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): layerwise_load = (fastvideo_args.inference_mode and fastvideo_args.dit_layerwise_offload and not fastvideo_args.use_fsdp_inference) # The AdaLN host cache also needs the projection weights to stay off the device from the start. - adaln_table = os.environ.get("FASTVIDEO_H3_ADALN_TABLE") or None + adaln_table = envs.FASTVIDEO_H3_ADALN_TABLE.get() or None adaln_host_cache = (fastvideo_args.inference_mode and not fastvideo_args.use_fsdp_inference - and (os.environ.get("FASTVIDEO_H3_ADALN_CACHE") == "1" or adaln_table is not None)) + and (envs.FASTVIDEO_H3_ADALN_CACHE.get() or adaln_table is not None)) layerwise_load = layerwise_load or adaln_host_cache model = maybe_load_fsdp_model( model_cls=model_cls, @@ -1221,7 +1223,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # FASTVIDEO_H3_SPLICE_TRANSFORMER=: a second checkpoint of the same architecture # runs denoising steps FASTVIDEO_H3_SPLICE_FROM_STEP (default 4) onward. # Only the primary ``transformer`` component splices; the spliced load itself never does. - splice_path = os.environ.get("FASTVIDEO_H3_SPLICE_TRANSFORMER") + splice_path = envs.FASTVIDEO_H3_SPLICE_TRANSFORMER.get() if (splice_path and not getattr(self, "_loading_splice", False) and hasattr(model, "attach_step_splice") and os.path.basename(os.path.normpath(model_path)) == "transformer"): self._loading_splice = True @@ -1229,7 +1231,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): late = self.load(splice_path, fastvideo_args) finally: self._loading_splice = False - from_step = int(os.environ.get("FASTVIDEO_H3_SPLICE_FROM_STEP", "4")) + from_step = envs.FASTVIDEO_H3_SPLICE_FROM_STEP.get() model.attach_step_splice(late, from_step) logger.info("Step splice: steps >= %d run the transformer from %s", from_step, splice_path) return model diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 01ccf3a75a..2e1d5831a3 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -6,7 +6,6 @@ from __future__ import annotations import contextlib -import os import re from collections.abc import Callable, Generator from itertools import chain @@ -387,7 +386,7 @@ def maybe_load_fsdp_model( skip_param_names |= set(fp8_targets.values()) fp8_keys = {f"{prefix}.{suffix}" for prefix in fp8_prefixes for suffix in ("weight", "weight_scale")} weight_iterator = ((name, tensor) for name, tensor in weight_iterator if name not in fp8_keys) - if os.environ.get("FASTVIDEO_H3_ADALN_TABLE"): + if envs.FASTVIDEO_H3_ADALN_TABLE.get(): # Precomputed AdaLN modulation replaces the per-block projections; never read their weights. skip_param_names |= {name for name, _ in model.named_parameters() if re.fullmatch(r"transformer_blocks\.\d+\.adaln_proj\.linear\.(weight|bias)", name)} diff --git a/fastvideo/models/vaes/minimax_h3_int8_convrot.py b/fastvideo/models/vaes/minimax_h3_int8_convrot.py index 86c0b0d2d7..d9a12a01f1 100644 --- a/fastvideo/models/vaes/minimax_h3_int8_convrot.py +++ b/fastvideo/models/vaes/minimax_h3_int8_convrot.py @@ -28,6 +28,8 @@ import torch.nn.functional as F from safetensors.torch import load_file as safetensors_load_file +import fastvideo.envs as envs + from fastvideo.logger import init_logger logger = init_logger(__name__) @@ -104,8 +106,8 @@ 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._fused_dequant = os.environ.get("FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT", "0") == "1" + self._transpose_view = envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.get() + self._fused_dequant = envs.FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT.get() 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: diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index 044dbecce2..fca78c244a 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -7,7 +7,6 @@ """ import math -import os from collections.abc import Iterator from dataclasses import dataclass @@ -16,6 +15,8 @@ import torch.nn.functional as F from torch.utils.checkpoint import checkpoint +import fastvideo.envs as envs + from fastvideo.attention import get_attn_backend from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig from fastvideo.platforms import AttentionBackendEnum @@ -301,7 +302,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" + self._share_int8_qkv = envs.FASTVIDEO_H3_VAE_INT8_SHARED_QKV.get() 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) @@ -530,7 +531,7 @@ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: def _tile_batch_size() -> int: """Spatial tiles decoded per decoder call (``FASTVIDEO_H3_VAE_TILE_BATCH``, default 1 = per tile).""" - return max(1, int(os.environ.get("FASTVIDEO_H3_VAE_TILE_BATCH", "1"))) + return max(1, envs.FASTVIDEO_H3_VAE_TILE_BATCH.get()) def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool: diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 9b2cb09433..f56a926af9 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -3,7 +3,6 @@ from __future__ import annotations import gc -import os import json import math from dataclasses import dataclass @@ -13,6 +12,8 @@ import torch from torch.distributed.tensor import DTensor +import fastvideo.envs as envs + from fastvideo.attention.selector import (_active_component_attention_backend_scope, coerce_attn_backend, get_env_variable_attn_backend) from fastvideo.configs.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAEArchConfig @@ -470,7 +471,7 @@ def _move_module(self, module: Any, device: str | torch.device) -> bool: return True if _module_has_dtensor_params(module): return False - if os.environ.get("FASTVIDEO_H3_PINNED_SWAP", "1") == "1": + if envs.FASTVIDEO_H3_PINNED_SWAP.get(): _pinned_swap(module, torch.device(device)) else: module.to(device) @@ -482,7 +483,7 @@ def _parked_module_names() -> tuple[str, ...]: Cards with room for the DiT next to the encoder park only the VAEs and keep the DiT resident. """ - requested = os.environ.get("FASTVIDEO_H3_PARK_MODULES") + requested = envs.FASTVIDEO_H3_PARK_MODULES.get() if not requested: return _DENOISE_MODULE_NAMES return tuple(name for name in requested.split(",") if name in _DENOISE_MODULE_NAMES) diff --git a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py index a13712a7b3..b96d0042f5 100644 --- a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py +++ b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py @@ -6,6 +6,8 @@ import pytest import torch +import fastvideo.envs as envs + 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 @@ -18,10 +20,11 @@ @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for encoder streaming") @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): +def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, env_overrides, 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") + env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.override(6)) + env_overrides.enter_context(envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.override(False)) config = MiniMaxH3Qwen3VLConfig() config.arch_config = MiniMaxH3Qwen3VLArchConfig( vocab_size=64, hidden_size=128, intermediate_size=256, diff --git a/fastvideo/tests/hooks/test_pinned_memory.py b/fastvideo/tests/hooks/test_pinned_memory.py index 7c7e76d2a7..9402397a8e 100644 --- a/fastvideo/tests/hooks/test_pinned_memory.py +++ b/fastvideo/tests/hooks/test_pinned_memory.py @@ -6,6 +6,8 @@ import pytest import torch + +import fastvideo.envs as envs from torch import nn from fastvideo.hooks.hooks import ModuleHookManager @@ -62,8 +64,8 @@ def cudaHostRegister(self, *_args): arena.close() -def test_offload_mutation_and_prefetched_detach(monkeypatch): - monkeypatch.setenv("FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS", "1") +def test_offload_mutation_and_prefetched_detach(env_overrides): + env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_OFFLOAD_BUFFERS.override(True)) 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())} diff --git a/fastvideo/tests/ssim/.reference_videos_download.lock b/fastvideo/tests/ssim/.reference_videos_download.lock new file mode 100644 index 0000000000..e69de29bb2 diff --git a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index 3763e9159a..44138298a8 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -8,6 +8,8 @@ import pytest import torch +import fastvideo.envs as envs + import fastvideo.pipelines.composed_pipeline_base as composed_pipeline_base from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.utils import FlexibleArgumentParser @@ -51,9 +53,13 @@ def to(device): return module -def _patch_pipeline_construction(monkeypatch, events: list, *, unified_memory: bool = False) -> None: +@pytest.fixture(autouse=True) +def _no_pinned_swap(env_overrides): # These contract tests use lightweight objects, not tensor-bearing modules. - monkeypatch.setenv("FASTVIDEO_H3_PINNED_SWAP", "0") + env_overrides.enter_context(envs.FASTVIDEO_H3_PINNED_SWAP.override(False)) + + +def _patch_pipeline_construction(monkeypatch, events: list, *, unified_memory: bool = False) -> None: monkeypatch.setattr( composed_pipeline_base, "maybe_init_distributed_environment_and_model_parallel", diff --git a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py index 752a709ce6..fa2ac25fcd 100644 --- a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py +++ b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py @@ -7,6 +7,8 @@ import pytest import torch +import fastvideo.envs as envs + 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 @@ -52,7 +54,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, tmp_path, +def test_tile_first_matches_generic_vsa_with_partial_tiles(env_overrides, 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") @@ -62,15 +64,15 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(monkeypatch, distribu 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") - monkeypatch.setenv("FASTVIDEO_H3_VSA_SM89_KERNEL", "original") + env_overrides.enter_context(envs.FASTVIDEO_ATTENTION_BACKEND.override("VIDEO_SPARSE_ATTN_H3")) + env_overrides.enter_context(envs.override_external("FASTVIDEO_VSA_TRITON", "1")) + env_overrides.enter_context(envs.FASTVIDEO_VSA_SM100A.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_FP4.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_TILE_FIRST.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_SM89_KERNEL.override("original")) capture = kernel == "int8" and fp8 and gate_active and fused_rope if capture: - monkeypatch.setenv("FASTVIDEO_H3_CAPTURE_QKV", str(tmp_path)) + env_overrides.enter_context(envs.FASTVIDEO_H3_CAPTURE_QKV.override(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) diff --git a/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py index 3874ab3eec..87caab3ee3 100644 --- a/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py +++ b/fastvideo/tests/vaes/test_minimax_h3_int8_shared.py @@ -5,6 +5,8 @@ import pytest import torch +import fastvideo.envs as envs + from fastvideo.models.vaes.minimax_h3_int8_convrot import Int8ConvRotLinear, shared_int8_projections @@ -12,8 +14,8 @@ @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") +def test_shared_int8_and_transpose_views_are_exact(rows, dtype, convrot, env_overrides): + env_overrides.enter_context(envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.override(False)) 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])) @@ -51,11 +53,11 @@ def test_shared_int8_keeps_cpu_fallback_exact(): @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): +def test_vae_attention_shares_quantized_projections_exactly(distributed_setup, env_overrides): 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") + env_overrides.enter_context(envs.FASTVIDEO_H3_VAE_INT8_SHARED_QKV.override(False)) + env_overrides.enter_context(envs.FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW.override(False)) torch.manual_seed(49) attention = MiniMaxH3VideoAttention(256, 2, 128).to("cuda").eval() for name in ("to_q", "to_k", "to_v"): diff --git a/fastvideo/worker/gpu_worker.py b/fastvideo/worker/gpu_worker.py index beb2af9385..f4de6a30f0 100644 --- a/fastvideo/worker/gpu_worker.py +++ b/fastvideo/worker/gpu_worker.py @@ -1,5 +1,4 @@ # SPDX-License-Identifier: Apache-2.0 -import os from typing import Any, cast import torch @@ -105,10 +104,10 @@ def init_device(self) -> None: if current_platform.is_cuda_alike(): torch.cuda.set_device(self.device) # Debug: FASTVIDEO_CUDA_MEMORY_CAP_GIB emulates a smaller card by capping this process's allocator. - cap_gib = os.environ.get("FASTVIDEO_CUDA_MEMORY_CAP_GIB") - if cap_gib: + cap_gib = envs.FASTVIDEO_CUDA_MEMORY_CAP_GIB.get() + if cap_gib > 0: total = torch.cuda.get_device_properties(self.device).total_memory - torch.cuda.set_per_process_memory_fraction(min(1.0, float(cap_gib) * 1024**3 / total), self.device) + torch.cuda.set_per_process_memory_fraction(min(1.0, cap_gib * 1024**3 / total), self.device) logger.info("Capped CUDA allocator at %s GiB of %.1f GiB", cap_gib, total / 1024**3) self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0] if current_platform.is_cuda(): @@ -129,7 +128,7 @@ def init_device(self) -> None: self.distributed_init_method) self.pipeline = build_pipeline(self.fastvideo_args) - if os.environ.get("FASTVIDEO_MEMORY_REPORT") == "1" and self.rank == 0: + if envs.FASTVIDEO_MEMORY_REPORT.get() and self.rank == 0: _log_pipeline_memory(self.pipeline) def execute_forward(self, forward_batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: From a59e84275e53a8cdbfc65b0f9895f0714128fbb4 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 08:39:44 -0700 Subject: [PATCH 058/107] [test]: NVFP4 H3 encoder tests follow the pre-Blackwell de-quantized fallback validate_runtime now accepts sm80-sm99 with a per-call bf16 de-quantization and checks the tensor-parallel size first, and _apply_finalized takes the fallback on any device without FP4 GEMM. Stub the TP size and assert the sm75 refusal / sm89 acceptance, pin the FP4 path in the apply test, and cover the fallback against the de-quantized reference linear. These two tests failed in the encoder lane on non-Blackwell CI GPUs. --- ...st_minimax_h3_qwen3_vl_checkpoint_nvfp4.py | 26 +++++++++++++++++-- 1 file changed, 24 insertions(+), 2 deletions(-) diff --git a/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py b/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py index a92bb7eb73..e3d2e4921f 100644 --- a/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py +++ b/fastvideo/tests/encoders/test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py @@ -260,9 +260,13 @@ def test_runtime_preflight_gates_device_capability_parallelism_and_flashinfer(mo with pytest.raises(RuntimeError, match="requires a CUDA device"): config.validate_runtime(torch.device("cpu")) - monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 9)) - with pytest.raises(RuntimeError, match="sm100 or newer"): + monkeypatch.setattr(h3_nvfp4, "get_tp_world_size", lambda: 1) + # Below sm100 there is no FP4 GEMM: sm80+ de-quantizes each linear to bf16, older GPUs are refused. + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (7, 5)) + with pytest.raises(RuntimeError, match="sm80"): config.validate_runtime(torch.device("cuda")) + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 9)) + config.validate_runtime(torch.device("cuda")) monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (12, 1)) monkeypatch.setattr(h3_nvfp4, "get_tp_world_size", lambda: 2) @@ -388,6 +392,7 @@ def fake_linear(x_fp4, x_scale, weight_packed, weight_scale, alpha): receipt.update(alpha=alpha, weight_packed=weight_packed, weight_scale=weight_scale) return torch.zeros(x_fp4.shape[0], weight_packed.shape[0], dtype=torch.bfloat16) + monkeypatch.setattr(h3_nvfp4, "_fp4_gemm_supported", lambda device: True) monkeypatch.setattr(h3_nvfp4, "_quantize_activation_nvfp4", fake_quantize) monkeypatch.setattr(h3_nvfp4, "_nvfp4_linear", fake_linear) output = MiniMaxH3SerializedNVFP4LinearMethod._apply_finalized(layer, torch.ones(1, 5, 128), layer.bias) @@ -402,6 +407,23 @@ def fake_linear(x_fp4, x_scale, weight_packed, weight_scale, alpha): assert torch.equal(output, torch.full((1, 5, 256), 3.0, dtype=torch.bfloat16)) +def test_apply_without_fp4_gemm_matches_the_dequantized_linear(distributed_setup, monkeypatch) -> None: + """Pre-Blackwell GPUs expand the packed weight to bf16 for the call; the result is x @ W_deq.T + b.""" + config = MiniMaxH3SerializedNVFP4Config.from_config(_checkpoint_quantization_config()) + layer = _language_linear(config, bias=True) + _fill_loaded(layer, global_scale=4.0) + layer.quant_method.process_weights_after_loading(layer) + layer.bias.data.fill_(3.0) + monkeypatch.setattr(h3_nvfp4, "_fp4_gemm_supported", lambda device: False) + x = torch.randn(1, 5, 128, dtype=torch.bfloat16) + output = MiniMaxH3SerializedNVFP4LinearMethod._apply_finalized(layer, x, layer.bias) + + weight = h3_nvfp4.dequantize_serialized_nvfp4(layer.weight_packed, layer.weight_scale, 4.0, torch.bfloat16) + # 0x11 packs two E2M1 codes of 0.5; E4M3 block scale 1; global scale 4 -> every weight is 0.125. + assert torch.equal(weight, torch.full_like(weight, 0.125)) + torch.testing.assert_close(output, torch.nn.functional.linear(x, weight, layer.bias.to(torch.bfloat16))) + + def _tiny_conditioner_config(keep_bf16: tuple[str, ...] = ("mlp.down_proj", )) -> MiniMaxH3Qwen3VLConfig: """One language layer whose linears satisfy the FlashInfer tile geometry, and a small vision tower.""" config = MiniMaxH3Qwen3VLConfig() From e9a35b2318c68c8e760567bf524cb5887a22669c Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 08:56:53 -0700 Subject: [PATCH 059/107] [misc]: drop the os imports layerwise offload no longer uses --- fastvideo/hooks/layerwise_offload.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 5a02af45b3..211ea0f3e1 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -28,7 +28,6 @@ def _offload_tensors(module: nn.Module, names: dict[str, torch.Tensor] | None = ``names`` restricts the walk to the tensors chosen at init: an offloaded buffer is a zero-element placeholder afterwards and would fail the size test. """ - import os if names is not None: for name, tensor in chain(module.named_parameters(), module.named_buffers()): if name in names: @@ -187,7 +186,6 @@ def enable_layerwise_offload(model: nn.Module, async_stream = torch.cuda.Stream() # 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 if resident_blocks is not None: resident = max(0, resident_blocks) else: From 84eb65d4fc6f75bb33a8a1d5b98fd5158c5f88b0 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 12:10:24 -0700 Subject: [PATCH 060/107] [bugfix]: FP4 VSA path stays compatible with fastvideo-kernel releases that predate validate= Passing validate=False to an older sageattn_blackwell_sparse_bshd raised TypeError on the first forward. Fall back to the positional call there. Also: headline benchmark adds a 768p 5 s setting and a showcase mode that renders a prompt set per seed at 480p 5 s. --- fastvideo/models/dits/minimax_h3_vsa_fp4.py | 20 +++++++++++--- scripts/benchmarks/minimax_h3_pro6000/app.py | 3 ++- .../minimax_h3_pro6000/bench_headline.py | 27 +++++++++++++++++-- .../minimax_h3_pro6000/showcase_prompts.json | 14 ++++++++++ 4 files changed, 57 insertions(+), 7 deletions(-) create mode 100644 scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json diff --git a/fastvideo/models/dits/minimax_h3_vsa_fp4.py b/fastvideo/models/dits/minimax_h3_vsa_fp4.py index 3f3a48ef5b..308e11c73d 100644 --- a/fastvideo/models/dits/minimax_h3_vsa_fp4.py +++ b/fastvideo/models/dits/minimax_h3_vsa_fp4.py @@ -27,6 +27,20 @@ from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAMetadata, _build_block_mask, _pool_tiles) + +def _sparse_fp4_attention(api: Any, *args: Any) -> torch.Tensor: + """Run the block-sparse FP4 kernel on lists built by vsa_tile_mask_to_fp4_blocks. + + Those lists are in range by construction, so skip the host-sync validation where the installed + fastvideo-kernel supports turning it off; older releases have no ``validate`` argument. + """ + try: + return api.sageattn_blackwell_sparse_bshd(*args, validate=False) + except TypeError as error: + if "validate" not in str(error): + raise + return api.sageattn_blackwell_sparse_bshd(*args) + _BLOCK = 128 @@ -235,8 +249,7 @@ def vsa_fp4_attention(attn: Any, hidden_states: torch.Tensor, rotary_emb: tuple[ q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) with STAGES.span("fp4_attention"): # Lists come from vsa_tile_mask_to_fp4_blocks and are in range; skip the per-call host sync. - out = api.sageattn_blackwell_sparse_bshd(query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad, - validate=False) + out = _sparse_fp4_attention(api, query, key, value, q2k_idx, q2k_num, kv_valid, q2k_quad) with STAGES.span("out_untile"): out = out.transpose(1, 2).index_select(1, layout.untile) # [B, L, H, D], packed order if sim_fp8: @@ -402,8 +415,7 @@ def vsa_fp4_attention_sp(attn: Any, hidden_states: torch.Tensor, rotary_emb: tup meta.span_sparsities) q2k_idx, q2k_num, kv_valid, q2k_quad = api.vsa_tile_mask_to_fp4_blocks(mask, layout.tile, vbs) with STAGES.span("fp4_attention"): - out_bhsd = api.sageattn_blackwell_sparse_bshd(q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad, - validate=False) + out_bhsd = _sparse_fp4_attention(api, q_t, k_t, v_t, q2k_idx, q2k_num, kv_valid, q2k_quad) with STAGES.span("out_pack"): payload, scale = _pack_seq_fp8(out_bhsd, layout.untile, world, local_rows) diff --git a/scripts/benchmarks/minimax_h3_pro6000/app.py b/scripts/benchmarks/minimax_h3_pro6000/app.py index 3ece05c8b3..d5e69bcc76 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/app.py +++ b/scripts/benchmarks/minimax_h3_pro6000/app.py @@ -542,7 +542,8 @@ def main(step: str = "all", ladder: str = "base"): # MODAL_PROFILE=aryan5v modal run --detach app.py::headline --repo FastVideo/ --profile h3_dit_ffn --gpus 1,4,8 HERE = pathlib.Path(__file__).resolve().parent headline_image = (image.add_local_file(HERE / "bench_headline.py", "/root/bench_headline.py") - .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json")) + .add_local_file(HERE / "headline_prompts.json", "/root/headline_prompts.json") + .add_local_file(HERE / "showcase_prompts.json", "/root/showcase_prompts.json")) SECRETS = [modal.Secret.from_name("hf-fastvideo")] diff --git a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py index cd25102905..0ed366cbbc 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py +++ b/scripts/benchmarks/minimax_h3_pro6000/bench_headline.py @@ -13,7 +13,7 @@ import statistics import time -SETTINGS = {"480p5s": (832, 480, 124), "768p10s": (1344, 768, 243)} # 17n+5 frames at 24 fps +SETTINGS = {"480p5s": (832, 480, 124), "768p5s": (1344, 768, 124), "768p10s": (1344, 768, 243)} # 17n+5 frames at 24 fps PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") @@ -23,7 +23,11 @@ def main(): ap.add_argument("model_dir") ap.add_argument("num_gpus", type=int) ap.add_argument("nvfp4_profile") - ap.add_argument("--settings", default="480p5s,768p10s") + ap.add_argument("--settings", default=os.environ.get("HEADLINE_SETTINGS", "480p5s,768p10s")) + ap.add_argument("--showcase", default=os.environ.get("HEADLINE_SHOWCASE"), + help="JSON {key: prompt}: after timing, render each prompt at 480p5s for every seed") + ap.add_argument("--showcase-seeds", default=os.environ.get("HEADLINE_SHOWCASE_SEEDS", "1234,42")) + ap.add_argument("--clip-prefix", default=os.environ.get("HEADLINE_CLIP_PREFIX", "model")) ap.add_argument("--prompts", default=os.path.join(os.path.dirname(os.path.abspath(__file__)), "headline_prompts.json")) ap.add_argument("--out", default=os.environ.get("HEADLINE_OUT", "headline")) ap.add_argument("--timed", type=int, default=2) @@ -106,6 +110,25 @@ def main(): if run is not None: run.summary[f"{name}_e2e_median_s"] = statistics.median(timed) json.dump(results, open(os.path.join(out_dir, "results.json"), "w"), indent=1) + if a.showcase: + # Gallery candidates: one clip per (prompt, seed) at 480p 5 s, named -[-s].mp4. + showcase = json.load(open(a.showcase)) + width, height, frames = SETTINGS["480p5s"] + seeds = [int(x) for x in a.showcase_seeds.split(",")] + os.makedirs(os.path.join(out_dir, "showcase"), exist_ok=True) + results["showcase"] = [] + for seed in seeds: + for key, text in showcase.items(): + suffix = "" if seed == 1234 else f"-s{seed}" + path = os.path.join(out_dir, "showcase", f"{a.clip_prefix}-{key}{suffix}.mp4") + t = time.perf_counter() + generator.generate_video(prompt=text, height=height, width=width, num_frames=frames, fps=24, + guidance_scale=1.0, num_inference_steps=len(steps) + 1, seed=seed, + output_path=path, save_video=True) + wall = round(time.perf_counter() - t, 2) + results["showcase"].append({"prompt": key, "seed": seed, "e2e_s": wall, "path": path}) + print("SHOWCASE", key, seed, wall, flush=True) + json.dump(results, open(os.path.join(out_dir, "results.json"), "w"), indent=1) finally: generator.shutdown() if run is not None: diff --git a/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json b/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json new file mode 100644 index 0000000000..c34ac48024 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json @@ -0,0 +1,14 @@ +{ + "violinist-archway": "A street violinist plays a short, bright melody under a stone archway at dusk. Medium shot from a fixed camera at chest height. Warm string lights hang overhead and a few pedestrians pass softly out of focus in the background. The bow moves in long, even strokes, and the violinist's fingers press precisely on the strings. The sound is a clear solo violin with light echo from the stone, distant footsteps and a faint murmur of the street. No other instruments. The violinist finishes the phrase, lowers the bow slightly and smiles.", + "chef-tasting": "In a bright home kitchen, a chef in a white jacket stands behind a wooden counter and looks straight at the camera. She lifts a wooden spoon from a simmering pot, tastes the sauce and says, \"A little more lemon, and it's perfect.\" Steam rises from the pot. The camera is fixed at eye level in a medium close-up. Soft bubbling, the light tap of the spoon on the pot rim and her clear, warm voice; no music.", + "fox-snow": "A red fox stands in a quiet snowy meadow at sunrise, listening with its head tilted. It crouches, then leaps high and dives nose-first into the deep snow, sending up a spray of powder, and pops back up with snow on its face. Wide shot from a low fixed camera, soft golden light, breath visible in the cold air. Sounds: gentle wind, the crunch of snow, a soft thump as the fox lands, a distant bird call. No music.", + "surfer-wave": "A single surfer rides a glassy turquoise wave at midday, carving a smooth turn along the face as the lip curls behind them. The camera follows from the water at board level, keeping the surfer centered. Spray catches the sunlight. Sounds: the roar of the breaking wave, rushing water and wind, a short exhilarated shout from the surfer. No music.", + "robot-windowsill": "3D animated short. A small round robot with big glowing eyes waters a single potted sunflower on a sunny windowsill. It tilts a tiny watering can carefully, the flower perks up, and the robot does a happy little wiggle. Fixed medium shot, soft cinematic lighting, gentle depth of field. Sounds: water trickling, soft mechanical whirs and a cheerful two-note beep. Light playful music box melody.", + "jazz-drummer": "A jazz drummer in a dim club plays a crisp brush pattern on the snare, then a short fill across the toms ending on a soft cymbal crash. Medium shot from a fixed camera slightly to the side, warm amber stage light, a hint of haze. Only drums are heard: the swish of brushes, tom hits and the shimmering cymbal, with a faint murmur from the audience.", + "rain-cafe": "A young woman sits by a rain-streaked café window, writing in a leather notebook with a fountain pen. She pauses, looks out at the street, then back at the camera and says softly, \"Some days are made for staying in.\" Close-up from a fixed camera, soft gray daylight with warm interior lamps. Sounds: steady rain on the glass, the scratch of the pen, a cup set down gently in the background. No music.", + "horse-beach": "A dark brown horse gallops along the edge of a wide beach at sunrise, splashing through shallow water. The camera tracks alongside at a steady distance, keeping the horse in profile. Low golden light, long shadows, spray from the hooves. Sounds: rhythmic hoofbeats on wet sand, splashing water, waves breaking, the horse's breathing. No music.", + "anime-rooftop": "2D anime style. A girl with short black hair and large headphones around her neck stands on a city rooftop at sunset, her jacket moving in the wind. She looks over the skyline, then turns to the camera and says, \"Tomorrow, we try again.\" Fixed medium shot, vivid orange and purple sky, clean cel shading. Sounds: wind, distant city traffic, a soft piano note under her line.", + "potter-hands": "Close-up of a potter's hands shaping a spinning clay bowl on a wheel. Wet fingers press gently, and the rim rises and smooths with each rotation. The camera is fixed just above wheel height, warm window light from the left. Sounds: the steady hum of the wheel, wet clay sliding under fingers, a soft splash. No music, no speech.", + "astronaut-radio": "Inside a spacecraft cockpit, an astronaut in a white suit looks out the window at Earth rising over the horizon. She presses a button on the console and says into her headset, \"Houston, the view is incredible today.\" Medium close-up from a fixed camera, soft blue light from the window, small floating particles. Sounds: a quiet cabin hum, a short radio click before and after her line.", + "grandma-bread": "In a cozy kitchen, a grandmother kneads bread dough on a floured wooden table, folding and pressing it with practiced hands. She looks up toward someone off camera, smiles and says, \"Patience makes the best bread.\" Medium shot from a fixed camera, warm afternoon light, flour dust drifting in the air. Sounds: dough thumping on the table, a ticking wall clock, her gentle voice. No music." +} From d60e209fe1ee62b2e9eb106abc091b9538fcc4b9 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:12:01 -0700 Subject: [PATCH 061/107] [feat]: add resident FastH3 V2 Spark recipe --- docs/getting_started/installation/spark.md | 3 + .../installation/spark_performance.md | 96 +++++++++++++++++-- .../basic/basic_fasth3_spark_v2_nvfp4.yaml | 60 ++++++++++++ 3 files changed, 152 insertions(+), 7 deletions(-) create mode 100644 examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml diff --git a/docs/getting_started/installation/spark.md b/docs/getting_started/installation/spark.md index 33888a853a..570cf71f2c 100644 --- a/docs/getting_started/installation/spark.md +++ b/docs/getting_started/installation/spark.md @@ -144,6 +144,9 @@ for which models are practical on the GB10, what makes them faster, and what won't help on this hardware (and why) — so you don't spend a night tuning knobs that can't move here. +For the eight-forward FastH3 V2 NVFP4 stack with a trimmed encoder and light +VAE, use the [one-Spark resident recipe](spark_performance.md#fasth3-v2-nvfp4-on-one-spark). + Two Sparks with QSFP cables: [Pair two NVIDIA DGX Sparks](spark_pair.md) for one FastH3 clip across both GPUs (`sp_size=2` over Ray). Copy-paste commands for one or two Sparks also live on the diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index e0f5cb4d13..bea284642b 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -161,14 +161,16 @@ is power-cycled. To avoid it: on: "CPU" offload uses the same unified RAM. Multi-GPU FSDP sharding remains available because it partitions weights without parking them in a separate host pool. -- **MiniMax H3 / FastH3** still needs deferred loading on one GB10. The Qwen3-VL - conditioner is tens of gigabytes of BF16. If the DiT and VAEs load while that - encoder is still resident, the process is a typical `earlyoom` kill (Python is - preferred). On unified memory, `lazy_module_load` auto-enables and owns that +- **Older MiniMax H3 / FastH3 bf16 weights** need deferred loading on one GB10. + The full Qwen3-VL conditioner is tens of gigabytes of BF16. If the DiT and + VAEs load while that encoder is still resident, the process can be killed by + `earlyoom`. On unified memory, `lazy_module_load` auto-enables and owns that split (encoder, then DiT, then VAE; DiT can drop before decode). Sequential - load is the H3-only fallback when lazy is off; do not pass - `--no-lazy-module-load` here. Geometry scalars come from checkpoint - `config.json`, not live weights. See [Offloading](../../inference/offloading.md). + load is the H3-only fallback when lazy is off. Keep deferred loading for + those older checkpoints. The trimmed NVFP4 encoder and light VAE in the + [V2 resident recipe](#fasth3-v2-nvfp4-on-one-spark) are a different memory + profile. Geometry scalars come from checkpoint `config.json`, not live + weights. See [Offloading](../../inference/offloading.md). - **FastH3 TAEH3** (`--video-decode-backend taeh3`) is an opt-in preview decoder. T2VA never materializes the 9.7 GiB video VAE (DiT still loads after Qwen via sequential start). On this box, alpine 768×1344×124 decoded in **2.4 s** versus @@ -207,6 +209,86 @@ A few things that surprise people on this box (beyond the memory notes above): `Released MiniMax-H3 text encoder after conditioning` before `Loading MiniMax-H3 denoise modules`). +## FastH3 V2 NVFP4 on one Spark + +This recipe uses the full V2 eight-forward transformer, the 50-layer NVFP4 +Qwen3-VL encoder, and the light H3 video VAE. Its configuration keeps all +three resident on one GB10. Runtime, memory fit, and quality still need a run +on that device. The earlier bf16 H3 memory guidance above concerns a larger +checkpoint. + +Install FastVideo from a checkout that includes the ModelOpt converter and +FlashInfer FP4 support, following [the Spark install guide](spark.md). Sign in +to Hugging Face with access to the FastVideo model repositories. Download the +V2 scheduler and audio components, the compact encoder and VAE from the pruned +repo, and the ModelOpt V2 transformer. The pruned model's encoder and VAE are +the same components used by V2. + +```bash +SPARK_STACK=./FastH3-V2-Spark-NVFP4 +V2_FP4_SRC=./FastH3-V2-ModelOpt-NVFP4 + +hf download FastVideo/FastVideo-FastH3-8-Step-V2 \ + --local-dir "$SPARK_STACK" \ + --exclude 'transformer/*' --exclude 'text_encoder/*' --exclude 'vae/*' +hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ + --local-dir "$SPARK_STACK" \ + --include 'text_encoder/*' --include 'vae/*' +hf download FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4 \ + --local-dir "$V2_FP4_SRC" --include 'transformer/*' + +nice -n 19 python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \ + --src "$V2_FP4_SRC/transformer" --dst "$SPARK_STACK/transformer" \ + --quantize-attention --quantize-gate + +test -f "$SPARK_STACK/transformer/nvfp4_weights.safetensors" +test -f "$SPARK_STACK/text_encoder/config.json" +test -f "$SPARK_STACK/vae/config.json" +test -f "$SPARK_STACK/fastvideo_inference.json" +python -m json.tool "$SPARK_STACK/fastvideo_inference.json" >/dev/null +``` + +The converter probes each packed linear through FlashInfer `mm_fp4`. If that +probe fails on `sm_121`, convert the transformer on another Blackwell GPU and +copy the resulting `transformer/` directory to the Spark. Do not omit +`fastvideo_inference.json`: it supplies V2's trained denoising ladder. The +recipe's `num_inference_steps: 9` means nine sigma points and eight DiT +forwards. + +Run `examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml` from the +repository root. It uses 832x480, 243 frames, VSA sparsity 0.8 with +64-token tiles, and the full H3 VAE. It does not use frame dropping or spatial +upscaling. + +```bash +FASTVIDEO_MINIMAX_H3_FUSIONS=all \ +FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +FASTVIDEO_H3_VAE_TILE_BATCH=1 \ +FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 \ +FASTVIDEO_STAGE_LOGGING=1 \ +nice -n 19 fastvideo generate \ + --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml +``` + +For a roughly five-second clip, set `--request.sampling.num_frames 124` and +write to a separate output path. H3 permits frame counts of `17n+5`; 124 is +the closest legal count above five seconds at 24 fps. For the secondary +10-second setting, set `--request.sampling.width 1344` and +`--request.sampling.height 768`, keeping 243 frames. Use the two prompts in +`handoff_spark_mac/benchmark_prompts.json` from the local release handoff. +Warm up once, then time at least two `generate_video` calls per prompt in one +process. Record the median wall time, denoise and decode stage times, peak +memory, exact command and commit, and retain every MP4 for review. The V2 +configuration must pass a visual and audio quality check before publication. + +After the V2 baseline works, sweep `FASTVIDEO_H3_VAE_TILE_BATCH` and +`FASTVIDEO_NVFP4_MM_BACKEND` on the same prompts. Compare the optional AdaLN +table and VAE compile only with the same frame count, schedule, and VSA +sparsity. The V2 converter packs VSA gates, so its `h3_dit_vsa` profile must +match the recipe. A later pruned NVFP4 transformer uses the separate +`h3_dit_ffn` profile, with attention and VSA gates left dense. + ## Reproduce these numbers Two scripts under `examples/inference/optimizations/` reproduce the claims on diff --git a/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml new file mode 100644 index 0000000000..c3986fabd0 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml @@ -0,0 +1,60 @@ +# FastH3 V2 eight-forward video+audio on one DGX Spark. +# Assemble ./FastH3-V2-Spark-NVFP4 as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer must include NVFP4 attention +# and VSA gates (layer_profile: h3_dit_vsa). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml +generator: + model_path: ./FastH3-V2-Spark-NVFP4 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_vsa + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_v2_nvfp4/ + save_video: true + return_frames: false From 487468266bbf5354049e8cc047fbe06cfe906293 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:29:56 -0700 Subject: [PATCH 062/107] [feat]: benchmark resident FastH3 Spark recipe --- .../installation/spark_performance.md | 19 ++- .../basic/benchmark_fasth3_spark_nvfp4.py | 135 ++++++++++++++++++ 2 files changed, 150 insertions(+), 4 deletions(-) create mode 100644 examples/inference/basic/benchmark_fasth3_spark_nvfp4.py diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index bea284642b..cd616a16ca 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -277,10 +277,21 @@ the closest legal count above five seconds at 24 fps. For the secondary 10-second setting, set `--request.sampling.width 1344` and `--request.sampling.height 768`, keeping 243 frames. Use the two prompts in `handoff_spark_mac/benchmark_prompts.json` from the local release handoff. -Warm up once, then time at least two `generate_video` calls per prompt in one -process. Record the median wall time, denoise and decode stage times, peak -memory, exact command and commit, and retain every MP4 for review. The V2 -configuration must pass a visual and audio quality check before publication. +The benchmark script runs one warmup and at least two timed generations for +each prompt in one process. It saves the MP4s and prints the wall time, stage +times, peak memory, and median. Use the same environment variables as above: + +```bash +nice -n 19 python examples/inference/basic/benchmark_fasth3_spark_nvfp4.py \ + --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml \ + --prompts /path/to/fasth3-local-release/handoff_spark_mac/benchmark_prompts.json \ + --output-dir outputs/fasth3_spark_v2_nvfp4/benchmark-243 --frames 243 + +# Repeat with --frames 124 and a different output directory for the five-second check. +``` + +Record the exact command and commit with the measurements. Review every clip's +video and audio before publishing a quality or speed claim. After the V2 baseline works, sweep `FASTVIDEO_H3_VAE_TILE_BATCH` and `FASTVIDEO_NVFP4_MM_BACKEND` on the same prompts. Compare the optional AdaLN diff --git a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py new file mode 100644 index 0000000000..9a82beeb99 --- /dev/null +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -0,0 +1,135 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Time the resident FastH3 V2 Spark recipe with the release prompts. + +One process loads the model, then each prompt gets one excluded warmup and at +least two timed calls. Every call writes a video. This script does not alter +the V2 schedule, VSA sparsity, or video resolution. +""" + +from __future__ import annotations + +import argparse +import json +import statistics +import time +from copy import deepcopy +from pathlib import Path + +from fastvideo import VideoGenerator +from fastvideo.api.parser import load_raw_config, parse_config +from fastvideo.api.schema import RunConfig + +PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") + + +def _stage_seconds(result: object) -> dict[str, float]: + logging_info = getattr(result, "logging_info", None) + stages = getattr(logging_info, "stages", None) + if isinstance(logging_info, dict): + stages = logging_info.get("stages", stages) + if not isinstance(stages, dict): + return {} + return { + name: float(metrics["execution_time"]) + for name, metrics in stages.items() + if isinstance(metrics, dict) and metrics.get("execution_time") is not None + } + + +def _stage_total(stages: dict[str, float], fragment: str) -> float | None: + matches = [seconds for name, seconds in stages.items() if fragment in name.lower()] + return sum(matches) if matches else None + + +def _request(base: RunConfig, prompt: str, frames: int, width: int, height: int, + output: Path): + request = deepcopy(base.request) + request.prompt = prompt + request.inputs.prompt_path = None + request.sampling.num_frames = frames + request.sampling.width = width + request.sampling.height = height + request.output.output_path = str(output) + request.output.save_video = True + request.output.return_frames = False + return request + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", type=Path, required=True) + parser.add_argument("--prompts", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--model-path", type=Path) + parser.add_argument("--frames", type=int, default=243) + parser.add_argument("--width", type=int, default=832) + parser.add_argument("--height", type=int, default=480) + parser.add_argument("--repeats", type=int, default=2) + args = parser.parse_args() + + if args.frames not in (124, 243): + parser.error("use 124 frames for roughly five seconds or 243 for the ten-second headline") + if args.repeats < 2: + parser.error("the release protocol requires at least two timed calls") + + config = parse_config(RunConfig, load_raw_config(args.config)) + if args.model_path: + config.generator.model_path = str(args.model_path) + if config.request.sampling.num_inference_steps != 9: + parser.error("the V2 contract requires nine sigma points for eight DiT forwards") + if config.generator.engine.offload.lazy_module_load is not False: + parser.error("the resident recipe requires lazy_module_load: false") + contract = Path(config.generator.model_path) / "fastvideo_inference.json" + if not contract.is_file(): + parser.error(f"missing trained V2 schedule: {contract}") + inference = json.loads(contract.read_text()) + if inference.get("num_inference_steps") != 9 or inference.get("transformer_forwards") != 8: + parser.error("the checkpoint is not the trained V2 eight-forward schedule") + + prompts = json.loads(args.prompts.read_text()) + if any(prompt_id not in prompts for prompt_id in PROMPT_IDS): + parser.error(f"prompt JSON must contain {', '.join(PROMPT_IDS)}") + args.output_dir.mkdir(parents=True, exist_ok=True) + + generator = VideoGenerator.from_config(config.generator) + try: + for prompt_id in PROMPT_IDS: + times = [] + for index in range(args.repeats + 1): + warmup = index == 0 + label = "warmup" if warmup else f"run-{index:02d}" + requested_path = args.output_dir / f"{prompt_id}-{args.width}x{args.height}-{args.frames}-{label}.mp4" + request = _request(config, prompts[prompt_id], args.frames, args.width, args.height, + requested_path) + started = time.perf_counter() + result = generator.generate(request) + wall = time.perf_counter() - started + output = Path(result.video_path) if result.video_path else requested_path + if not output.is_file(): + raise RuntimeError(f"generation returned without an MP4: {output}") + stages = _stage_seconds(result) + row = { + "prompt_id": prompt_id, + "warmup": warmup, + "frames": args.frames, + "width": args.width, + "height": args.height, + "e2e_seconds": round(wall, 3), + "denoise_seconds": _stage_total(stages, "denois"), + "decode_seconds": _stage_total(stages, "decod"), + "peak_memory_mb": result.peak_memory_mb, + "stages": stages, + "mp4": str(output), + } + print(json.dumps(row, sort_keys=True), flush=True) + if not warmup: + times.append(wall) + print(json.dumps({"prompt_id": prompt_id, "timed_runs": len(times), + "median_e2e_seconds": round(statistics.median(times), 3)}, + sort_keys=True), flush=True) + finally: + generator.shutdown() + + +if __name__ == "__main__": + main() From 8c01b5a48585c76c88d2a121593bd6fe1e37d369 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:30:57 -0700 Subject: [PATCH 063/107] [docs]: pin Spark benchmark environment --- docs/getting_started/installation/spark_performance.md | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index cd616a16ca..00c9080d2c 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -279,9 +279,13 @@ the closest legal count above five seconds at 24 fps. For the secondary `handoff_spark_mac/benchmark_prompts.json` from the local release handoff. The benchmark script runs one warmup and at least two timed generations for each prompt in one process. It saves the MP4s and prints the wall time, stage -times, peak memory, and median. Use the same environment variables as above: +times, peak memory, and median. Set the Spark environment before running it: ```bash +export FASTVIDEO_MINIMAX_H3_FUSIONS=all +export FASTVIDEO_NVFP4_MM_BACKEND=cutlass FASTVIDEO_H3_VAE_TILE_BATCH=1 +export FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 +export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 nice -n 19 python examples/inference/basic/benchmark_fasth3_spark_nvfp4.py \ --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml \ --prompts /path/to/fasth3-local-release/handoff_spark_mac/benchmark_prompts.json \ From 6ab83c33fd828657ab7a2b7cffdb4bfca5d5e4c2 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:01:22 -0700 Subject: [PATCH 064/107] [bugfix]: retain per-stage CUDA peaks in Spark benchmarks --- .../basic/benchmark_fasth3_spark_nvfp4.py | 27 ++++++++++++------- fastvideo/pipelines/stages/base.py | 4 +++ 2 files changed, 22 insertions(+), 9 deletions(-) diff --git a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py index 9a82beeb99..f8e3ee7fcd 100644 --- a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Time the resident FastH3 V2 Spark recipe with the release prompts. +"""Time a resident FastH3 eight-forward Spark recipe with the release prompts. One process loads the model, then each prompt gets one excluded warmup and at least two timed calls. Every call writes a video. This script does not alter @@ -22,18 +22,24 @@ PROMPT_IDS = ("latency-ceramics-005", "latency-harbor-005") -def _stage_seconds(result: object) -> dict[str, float]: +def _stage_metrics(result: object) -> dict[str, dict]: logging_info = getattr(result, "logging_info", None) stages = getattr(logging_info, "stages", None) if isinstance(logging_info, dict): stages = logging_info.get("stages", stages) if not isinstance(stages, dict): return {} - return { - name: float(metrics["execution_time"]) - for name, metrics in stages.items() - if isinstance(metrics, dict) and metrics.get("execution_time") is not None - } + return {name: metrics for name, metrics in stages.items() if isinstance(metrics, dict)} + + +def _stage_seconds(metrics: dict[str, dict]) -> dict[str, float]: + return {name: float(stage["execution_time"]) for name, stage in metrics.items() + if stage.get("execution_time") is not None} + + +def _peak_mb(metrics: dict[str, dict], key: str) -> float | None: + values = [float(stage[key]) for stage in metrics.values() if stage.get(key) is not None] + return max(values) if values else None def _stage_total(stages: dict[str, float], fragment: str) -> float | None: @@ -107,7 +113,8 @@ def main() -> None: output = Path(result.video_path) if result.video_path else requested_path if not output.is_file(): raise RuntimeError(f"generation returned without an MP4: {output}") - stages = _stage_seconds(result) + metrics = _stage_metrics(result) + stages = _stage_seconds(metrics) row = { "prompt_id": prompt_id, "warmup": warmup, @@ -117,7 +124,9 @@ def main() -> None: "e2e_seconds": round(wall, 3), "denoise_seconds": _stage_total(stages, "denois"), "decode_seconds": _stage_total(stages, "decod"), - "peak_memory_mb": result.peak_memory_mb, + "peak_memory_mb": _peak_mb(metrics, "peak_allocated_mb"), + "peak_reserved_mb": _peak_mb(metrics, "peak_reserved_mb"), + "result_peak_memory_mb": result.peak_memory_mb, "stages": stages, "mp4": str(output), } diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index aa086a7d8f..71d9dd0133 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -195,6 +195,10 @@ def _execute( torch.cuda.max_memory_allocated() / gib, torch.cuda.memory_reserved() / gib, torch.cuda.memory_allocated() / gib) + batch.logging_info.add_stage_metric(stage_key, "peak_allocated_mb", + torch.cuda.max_memory_allocated() / 1024**2) + batch.logging_info.add_stage_metric(stage_key, "peak_reserved_mb", + torch.cuda.max_memory_reserved() / 1024**2) torch.cuda.reset_peak_memory_stats() batch.logging_info.add_stage_execution_time(stage_key, execution_time) batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name) From b1c3495faec60efe864dfe746453cc6d4ad566b4 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:01:22 -0700 Subject: [PATCH 065/107] [feat]: add resident pruned and paired Spark recipes --- .../basic_fasth3_spark_pair_pruned_nvfp4.yaml | 63 +++++++++++++++++++ .../basic_fasth3_spark_pair_v2_nvfp4.yaml | 63 +++++++++++++++++++ .../basic_fasth3_spark_pruned_nvfp4.yaml | 60 ++++++++++++++++++ 3 files changed, 186 insertions(+) create mode 100644 examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml create mode 100644 examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml create mode 100644 examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml diff --git a/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml new file mode 100644 index 0000000000..32852ad514 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml @@ -0,0 +1,63 @@ +# Start Ray on both Sparks and source spark_pair_env.sh first; see spark_pair.md. +# FastH3 pruned ckpt300 eight-forward video+audio on two DGX Sparks over QSFP RoCE. +# Download the complete checkpoint, including its trained schedule, as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer uses NVFP4 FFN weights with bf16 attention +# and VSA gates (layer_profile: h3_dit_ffn). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml +generator: + model_path: FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 + engine: + num_gpus: 2 + execution_backend: ray + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_ffn + parallelism: + tp_size: 1 + sp_size: 2 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: true + vae_parallel_decode_strategy: gather + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pair_pruned_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml new file mode 100644 index 0000000000..89bef8ed8a --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml @@ -0,0 +1,63 @@ +# Start Ray on both Sparks and source spark_pair_env.sh first; see spark_pair.md. +# FastH3 V2 eight-forward video+audio on two DGX Sparks over QSFP RoCE. +# Assemble ./FastH3-V2-Spark-NVFP4 as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer must include NVFP4 attention +# and VSA gates (layer_profile: h3_dit_vsa). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml +generator: + model_path: ./FastH3-V2-Spark-NVFP4 + engine: + num_gpus: 2 + execution_backend: ray + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_vsa + parallelism: + tp_size: 1 + sp_size: 2 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: true + vae_parallel_decode_strategy: gather + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pair_v2_nvfp4/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml new file mode 100644 index 0000000000..a5d1af4c47 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml @@ -0,0 +1,60 @@ +# FastH3 pruned ckpt300 eight-forward video+audio on one DGX Spark. +# Download the complete checkpoint, including its trained schedule, as described in +# docs/getting_started/installation/spark_performance.md. +# +# GB10 uses Triton VSA. The packed transformer uses NVFP4 FFN weights with bf16 attention +# and VSA gates (layer_profile: h3_dit_ffn). No temporal or spatial fast mode. +# +# FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +# FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml +generator: + model_path: FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 + engine: + num_gpus: 1 + use_fsdp_inference: false + quantization: + transformer_quant: NVFP4 + layer_profile: h3_dit_ffn + parallelism: + tp_size: 1 + sp_size: 1 + offload: + dit: false + dit_layerwise: false + text_encoder: false + image_encoder: false + vae: false + pin_cpu_memory: false + lazy_module_load: false + compile: + enabled: false + vae_enabled: false + pipeline: + workload_type: t2v + vae_tiling: true + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.8 + VSA_tile_size: 64 + h3_sequential_load: false + inference_torch_compile: false + vae_parallel_decode: false + video_decode_backend: h3-vae +request: + prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. + negative_prompt: "" + sampling: + seed: 2026 + height: 480 + width: 832 + num_frames: 243 + fps: 24 + num_inference_steps: 9 # nine sigma points, eight DiT forwards + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pruned_nvfp4/ + save_video: true + return_frames: false From 4ce033aebd2098fe7644bc005639404a871f9636 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:09:41 -0700 Subject: [PATCH 066/107] [bugfix]: preserve H3 release helpers after main integration --- fastvideo/models/loader/fsdp_load.py | 2 ++ fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py | 3 +++ 2 files changed, 5 insertions(+) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 2e1d5831a3..a325595c9f 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -5,6 +5,8 @@ # Copyright 2025 The FastVideo Authors. from __future__ import annotations + +import os import contextlib import re from collections.abc import Callable, Generator diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index f56a926af9..661b8ef8c6 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -471,6 +471,9 @@ def _move_module(self, module: Any, device: str | torch.device) -> bool: return True if _module_has_dtensor_params(module): return False + if not callable(getattr(module, "named_parameters", None)): + module.to(device) + return True if envs.FASTVIDEO_H3_PINNED_SWAP.get(): _pinned_swap(module, torch.device(device)) else: From 40f6c3c77c2665202d3d7aef45638ee5edc74ac8 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:13:49 -0700 Subject: [PATCH 067/107] [feat]: support pruned FastH3 AdaLN in MLX --- fastvideo/mlx_runtime/minimax_h3.py | 40 +++++++++++++++++++---------- 1 file changed, 27 insertions(+), 13 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 58b677bd80..5b5e546cc6 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -766,12 +766,12 @@ def _feed_forward(weights: dict[str, Any], x): return linear(value * silu(gate), weights["ff.net.2.weight"]) -def _adaln_tables(weights: dict[str, Any], temb): +def _adaln_tables(weights: dict[str, Any], temb, *, apply_silu: bool = True): """Six (n_t * 3, hidden) modulation tables from (n_t, time_embed_dim).""" import mlx.core as mx projected = linear( - silu(temb).astype(weight_dtype(weights["adaln_proj.linear.weight"])), + (silu(temb) if apply_silu else temb).astype(weight_dtype(weights["adaln_proj.linear.weight"])), weights["adaln_proj.linear.weight"], weights["adaln_proj.linear.bias"], ) @@ -914,6 +914,7 @@ def __init__( self.qk_norm_eps = float(config["qk_norm_eps"]) self.final_norm_eps = float(config["final_norm_eps"]) self.patch_dim = self.in_channels * math.prod(self.patch_size) + self.adaln_rank = config.get("adaln_rank") self._adaln_cache: MiniMaxH3StepCache | None = None self.vsa_config = MiniMaxH3VSAConfig() self._vsa_geometry: MiniMaxH3VSAGeometry | None = None @@ -932,11 +933,15 @@ def compute_temb(self, timesteps): self.weights["time_embedder.linear_1.weight"], self.weights["time_embedder.linear_1.bias"], ) - return linear( + temb = linear( silu(temb), self.weights["time_embedder.linear_2.weight"], self.weights["time_embedder.linear_2.bias"], ) + if self.adaln_rank is not None: + temb = linear( + silu(temb).astype(weight_dtype(self.weights["adaln_basis.weight"])), self.weights["adaln_basis.weight"]) + return temb def refine_text(self, text_rows): hidden = linear( @@ -969,9 +974,10 @@ def precompute_adaln(self, timesteps: np.ndarray, *, drop_weights: bool = True) timesteps = np.unique(np.asarray(timesteps, dtype=np.float32)) temb = self.compute_temb(mx.array(timesteps)) - block_tables = [_adaln_tables(block, temb) for block in self.blocks] + block_tables = [_adaln_tables(block, temb, apply_silu=self.adaln_rank is None) for block in self.blocks] shift_scale = linear( - silu(temb).astype(weight_dtype(self.weights["norm_out.linear.weight"])), + (silu(temb) if self.adaln_rank is None else temb).astype( + weight_dtype(self.weights["norm_out.linear.weight"])), self.weights["norm_out.linear.weight"], self.weights["norm_out.linear.bias"], ) @@ -1107,7 +1113,7 @@ def forward( adaln_indices = (timestep_indices * MINIMAX_H3_MODALITY_NUM + token_tags).astype(mx.int32) for block_index, block in enumerate(self.blocks): - tables = _adaln_tables(block, temb) + tables = _adaln_tables(block, temb, apply_silu=self.adaln_rank is None) packed = _transformer_block( block, packed, @@ -1124,7 +1130,8 @@ def forward( mx.eval(packed) # per-block sync: see forward_with_cache note shift_scale = linear( - silu(temb).astype(weight_dtype(self.weights["norm_out.linear.weight"])), + (silu(temb) if self.adaln_rank is None else temb).astype( + weight_dtype(self.weights["norm_out.linear.weight"])), self.weights["norm_out.linear.weight"], self.weights["norm_out.linear.bias"], ) @@ -1353,10 +1360,10 @@ def assign(key: str, value) -> None: for shard in _safetensors_shards(transformer_path): shard_arrays = mx.load(str(shard)) for key, source in shard_arrays.items(): - if not key.startswith("time_embedder."): + if not (key.startswith("time_embedder.") or key == "adaln_basis.weight"): continue keep_fp32 = key.split(".", 1)[0] in FP32_MODULE_PREFIXES - target_dtype = mx.float32 if keep_fp32 else cast_dtype + target_dtype = mx.float32 if keep_fp32 else (mx.float16 if key == "adaln_basis.weight" else cast_dtype) assign(key, _load_array(source, target_dtype)) del shard_arrays required_time_keys = { @@ -1373,6 +1380,10 @@ def assign(key: str, value) -> None: weight_dtype(weights["time_embedder.linear_1.weight"])) temb = linear(t_freq, weights["time_embedder.linear_1.weight"], weights["time_embedder.linear_1.bias"]) temb = linear(silu(temb), weights["time_embedder.linear_2.weight"], weights["time_embedder.linear_2.bias"]) + if config.get("adaln_rank") is not None: + if "adaln_basis.weight" not in weights: + raise KeyError("Rank-reduced AdaLN checkpoint is missing adaln_basis.weight") + temb = linear(silu(temb).astype(weight_dtype(weights["adaln_basis.weight"])), weights["adaln_basis.weight"]) mx.eval(temb) cached_block_tables = [None] * num_blocks @@ -1381,7 +1392,7 @@ def assign(key: str, value) -> None: for key, source in shard_arrays.items(): if _is_ignored_dense_key(key, include_vsa=include_vsa): continue - if temb is not None and key.startswith("time_embedder."): + if temb is not None and (key.startswith("time_embedder.") or key == "adaln_basis.weight"): continue if key.startswith("transformer_blocks."): index = int(key.split(".")[1]) @@ -1390,7 +1401,9 @@ def assign(key: str, value) -> None: if key.startswith("rope."): continue # non-persistent analytic buffer, rebuilt on the fly keep_fp32 = key.split(".", 1)[0] in FP32_MODULE_PREFIXES - target_dtype = mx.float32 if keep_fp32 else cast_dtype + factorized_adaln = config.get("adaln_rank") is not None and (".adaln_proj." in key or key.startswith( + ("norm_out.linear.", "adaln_basis."))) + target_dtype = mx.float32 if keep_fp32 else (mx.float16 if factorized_adaln else cast_dtype) array = _load_array(source, target_dtype) if temb is not None and ".adaln_proj.linear." in key: _, index_str, sub = key.split(".", 2) @@ -1398,7 +1411,7 @@ def assign(key: str, value) -> None: block_pending = pending_adaln.setdefault(index, {}) block_pending[sub] = array if {"adaln_proj.linear.weight", "adaln_proj.linear.bias"} <= block_pending.keys(): - tables = _adaln_tables(block_pending, temb) + tables = _adaln_tables(block_pending, temb, apply_silu=config.get("adaln_rank") is None) mx.eval(tables) assert cached_block_tables is not None cached_block_tables[index] = tables @@ -1428,7 +1441,8 @@ def assign(key: str, value) -> None: if missing_cache_blocks: raise KeyError(f"Missing AdaLN cache tables for blocks {missing_cache_blocks}") shift_scale = linear( - silu(temb).astype(weight_dtype(weights["norm_out.linear.weight"])), + (silu(temb) if config.get("adaln_rank") is None else temb).astype( + weight_dtype(weights["norm_out.linear.weight"])), weights["norm_out.linear.weight"], weights["norm_out.linear.bias"], ) From 159168dd1dd535bf0d045d41c9100d1c5130ac02 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:17:23 -0700 Subject: [PATCH 068/107] [feat]: trim MLX H3 conditioner to required layers --- fastvideo/mlx_runtime/minimax_h3_conditioner.py | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index 464250a9c2..b395a9dab1 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -79,7 +79,17 @@ def __init__(self, component_dir: Path): self._header_cache: dict[str, tuple[dict, int]] = {} if index_path.exists(): weight_map = json.loads(index_path.read_text())["weight_map"] - self.key_to_shard = {k: str(component_dir / s) for k, s in weight_map.items()} + + def needed(key: str) -> bool: + if key == "model.language_model.embed_tokens.weight": + return True + prefix = "model.language_model.layers." + if not key.startswith(prefix): + return False + layer = key[len(prefix):].split(".", 1)[0] + return layer.isdigit() and int(layer) < TEXT_ENCODER_LAYER + + self.key_to_shard = {k: str(component_dir / s) for k, s in weight_map.items() if needed(k)} else: single = component_dir / "model.safetensors" if not single.exists(): @@ -294,8 +304,8 @@ def encode_tokens(self, token_ids: list[int]) -> tuple[np.ndarray, np.ndarray]: del rows gc.collect() - if cfg.num_layers <= TEXT_ENCODER_LAYER: - raise ValueError(f"Conditioner needs > {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") + if cfg.num_layers < TEXT_ENCODER_LAYER: + raise ValueError(f"Conditioner needs at least {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") for layer in range(TEXT_ENCODER_LAYER): hidden = self._decoder_layer(layer, hidden, cos, sin) # Per-layer sync: without this the whole 50-layer graph accumulates From f4f903da40027537555d8269827aca11cc3fe002 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 10:22:43 -0700 Subject: [PATCH 069/107] [feat]: allow FP8 source for MLX H3 conversion --- fastvideo/mlx_runtime/minimax_h3.py | 14 +++++++++++++- .../convert_minimax_h3_mlx.py | 5 +++++ 2 files changed, 18 insertions(+), 1 deletion(-) diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 5b5e546cc6..7cdf1c3dbf 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -1390,6 +1390,8 @@ def assign(key: str, value) -> None: for shard in _safetensors_shards(transformer_path): shard_arrays = mx.load(str(shard)) for key, source in shard_arrays.items(): + if key.endswith(".weight_scale"): + continue # paired with its FP8 weight if _is_ignored_dense_key(key, include_vsa=include_vsa): continue if temb is not None and (key.startswith("time_embedder.") or key == "adaln_basis.weight"): @@ -1404,7 +1406,17 @@ def assign(key: str, value) -> None: factorized_adaln = config.get("adaln_rank") is not None and (".adaln_proj." in key or key.startswith( ("norm_out.linear.", "adaln_basis."))) target_dtype = mx.float32 if keep_fp32 else (mx.float16 if factorized_adaln else cast_dtype) - array = _load_array(source, target_dtype) + if source.dtype == mx.uint8 and key.endswith(".weight"): + scale_key = key + "_scale" + if scale_key not in shard_arrays: + raise KeyError(f"FP8 weight {key} needs {scale_key} in the same safetensors shard") + scale = shard_arrays[scale_key].astype(mx.float32) + if scale.size != source.shape[0]: + raise ValueError(f"FP8 scale for {key} has {scale.size} entries, expected {source.shape[0]}") + array = (mx.from_fp8(source, dtype=mx.float16) * scale.reshape(-1, 1)).astype(target_dtype) + mx.eval(array) + else: + array = _load_array(source, target_dtype) if temb is not None and ".adaln_proj.linear." in key: _, index_str, sub = key.split(".", 2) index = int(index_str) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index db96f413c2..1e2813e6ab 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -18,6 +18,11 @@ INT8/INT6/INT4 grid, and record ``vsa.capable`` in the manifest. Write VSA checkpoints to a new directory — do not overwrite an existing dense export. +An FP8 transformer with per-channel ``weight_scale`` can also be a source. +The loader dequantizes each FP8 matrix before applying the requested MLX +quantization. This saves download bytes but quantizes twice; compare its clips +with the BF16-sourced export before using it for release. + python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ --out ~/models/FastH3-MLX-vsa \\ From c5d233e86cfa76ceb35882761e1e04773d3d920d Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 15:56:41 -0700 Subject: [PATCH 070/107] [docs]: add pruned FastH3 MLX conversion recipe --- docs/cookbook/minimax-h3.md | 5 +++ docs/getting_started/installation/mlx.md | 41 ++++++++++++++++++++++++ 2 files changed, 46 insertions(+) diff --git a/docs/cookbook/minimax-h3.md b/docs/cookbook/minimax-h3.md index e01ee210a4..0ccb0b3bbe 100644 --- a/docs/cookbook/minimax-h3.md +++ b/docs/cookbook/minimax-h3.md @@ -13,6 +13,11 @@ is why V2 is the higher-quality FastH3. The V2 schedule contract is in **CompactH3** is the 42-block 20B NVFP4 H3 checkpoint for one Blackwell GPU (RTX 5090 or RTX PRO 6000). The FastH3 V1 and V2 recipes are unchanged. +The 42-block pruned checkpoint has an [MLX INT8/INT6 conversion and +eight-forward T2VA command](../getting_started/installation/mlx.md#pruned-eight-forward-checkpoint). +It reads `fastvideo_inference.json` for the trained schedule. The command +uses native 832x480 resolution and all requested frames. +
      All model families diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index 359218e583..36d20791d9 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -46,6 +46,47 @@ is the higher-quality FastH3. Recorded shapes and evidence live in the [support matrix](../../inference/support_matrix.md#apple-silicon-native-runtime). +## Pruned eight-forward checkpoint + +The pruned FastH3 checkpoint has 42 transformer blocks and rank-16 AdaLN. +Its `fastvideo_inference.json` fixes eight denoising forwards, video/audio +shifts of 10/3, and VSA sparsity 0.8. Keep that file beside the transformer +when converting. The converter reads its schedule to build the AdaLN cache. + +```bash +hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ + --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ + --exclude 'text_encoder/*' + +# The pruned repo's packed NVFP4 encoder is for Blackwell. MLX reads the +# official BF16 encoder through layer 50. The last three shards are unused. +hf download MiniMaxAI/MiniMax-H3 \ + --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ + --include 'text_encoder/model-0000[1-9]-of-00014.safetensors' \ + --include 'text_encoder/model-0001[0-1]-of-00014.safetensors' \ + --include 'text_encoder/model.safetensors.index.json' \ + --include 'text_encoder/config.json' + +python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \ + --model-root ./FastH3-Pruned-8Step-BF16-ckpt300/transformer \ + --out ./FastH3-Pruned-MLX-vsa \ + --formats "int8 int6" --include-vsa + +python examples/inference/basic/mlx_fasth3.py \ + --model-root ./FastH3-Pruned-8Step-BF16-ckpt300 \ + --mlx-checkpoint ./FastH3-Pruned-MLX-vsa/int8 \ + --prompt "(S1) A potter asks [English] Is the rim ready?" \ + --height 480 --width 832 --num-frames 243 --steps 8 \ + --vsa --vsa-sparsity 0.8 --vsa-tile-size 64 \ + --output-path ./outputs/fasth3_pruned_int8_480p.mp4 +``` + +At 24 fps, 124 frames is the legal H3 count for a roughly five-second clip. +Use `--num-frames 124` and a separate output path for that run. The `--fast` +and `--fast-spatial` options change the workload and are not part of the +native-resolution benchmark. A 36 GB Mac may need INT6 and phased loading; +measure memory before claiming all-resident operation. + ## Hardware - FastMetal 1.3B and 5B: 16 GB unified memory and up From 89d3986ff60a41952a6d4cbf260a7aa44c8744ee Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:17:50 -0700 Subject: [PATCH 071/107] [feat]: add resident NVFP4 conditioning for MLX H3 --- docs/getting_started/installation/mlx.md | 45 +++++- .../mlx_runtime/minimax_h3_conditioner.py | 134 +++++++++++++++-- fastvideo/mlx_runtime/minimax_h3_pipeline.py | 95 +++++++++++- .../tests/test_minimax_h3_resident_nvfp4.py | 136 ++++++++++++++++++ 4 files changed, 388 insertions(+), 22 deletions(-) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index 36d20791d9..1f4974b98e 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -58,8 +58,8 @@ hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ --exclude 'text_encoder/*' -# The pruned repo's packed NVFP4 encoder is for Blackwell. MLX reads the -# official BF16 encoder through layer 50. The last three shards are unused. +# Optional BF16 encoder fallback: stream the first 50 language layers. +# The last three shards are unused. The packed NVFP4 option is described below. hf download MiniMaxAI/MiniMax-H3 \ --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ --include 'text_encoder/model-0000[1-9]-of-00014.safetensors' \ @@ -87,6 +87,47 @@ and `--fast-spatial` options change the workload and are not part of the native-resolution benchmark. A 36 GB Mac may need INT6 and phased loading; measure memory before claiming all-resident operation. +### Packed encoder and resident loading + +The experimental MLX conditioner can read the released FastVideo NVFP4 +text encoder directly, using native `nvfp4` matrix multiplication. It keeps +the packed weights and BF16 embedding table in memory, with FP32 +activations. CUDA uses quantized activations, so the two encoders are not +bit-exact. Validate generated video and audio before publishing a timing. +MLX 0.32.2 supports the required operator on Apple Silicon. + +Pass the packed encoder directory as `conditioner_dir`; `conditioner_mode="auto"` +selects it from `config.json`. The BF16 fallback continues to stream layers. +To request all-resident generation through the Python API: + +```python +from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline + +pipeline = MiniMaxH3MLXPipeline( + model_root="./FastH3-Pruned-8Step-BF16-ckpt300", + mlx_dit_checkpoint="./FastH3-Pruned-MLX-vsa/int6", + conditioner_dir="./FastH3-NVFP4-encoder", + conditioner_mode="nvfp4", + resident=True, + vae_dtype="fp16", +) +try: + pipeline.prepare_resident() # Load encoder, DiT, video VAE and audio VAE. + result = pipeline.generate( + "(S1) A potter asks [English] Is the rim ready?", + output_path="./outputs/fasth3_pruned_resident.mp4", + height=480, width=832, num_frames=243, num_steps=8, + vsa=True, vsa_sparsity=0.8, vsa_tile_size=64, + ) +finally: + pipeline.close() +``` + +Resident placement requires space for activations as well as all four +components. On a 36 GiB Mac, try INT6 first and measure peak allocation. +If loading or inference runs out of memory, use phased loading by leaving +`resident=False`. Changing placement does not change frames or resolution. + ## Hardware - FastMetal 1.3B and 5B: 16 GB unified memory and up diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index b395a9dab1..d5fe8b34b8 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -22,6 +22,7 @@ import json from dataclasses import dataclass from pathlib import Path +from typing import Any import numpy as np import mlx.core as mx @@ -146,7 +147,7 @@ def close(self) -> None: gc.collect() -_DTYPES = {"F32": np.float32, "F16": np.float16, "I64": np.int64, "I32": np.int32} +_DTYPES = {"F32": np.float32, "F16": np.float16, "I64": np.int64, "I32": np.int32, "U8": np.uint8} def _read_bf16_words(path: str, key: str, header: dict, data_start: int) -> np.ndarray: @@ -215,8 +216,20 @@ def _rms_norm(x, weight, eps: float): return x / mx.sqrt(mx.mean(x * x, axis=-1, keepdims=True) + eps) * weight +@dataclass(frozen=True) +class NVFP4Matrix: + """MLX row-major E2M1/E4M3 weights and the export's inverse global scale.""" + + weight: mx.array + scales: mx.array + global_scale: float + + def matmul(self, x): + return mx.quantized_matmul(x, self.weight, self.scales, mode="nvfp4") / self.global_scale + + def _linear(x, weight, bias=None): - y = x @ weight.T + y = weight.matmul(x) if isinstance(weight, NVFP4Matrix) else x @ weight.T if bias is not None: y = y + bias return y @@ -248,7 +261,7 @@ class StreamedMiniMaxH3TextConditioner: def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): self.component_dir = Path(component_dir) self.config = ConditionerConfig.from_config_json(self.component_dir / "config.json") - self.index = _ShardIndex(self.component_dir) + self.index: Any = _ShardIndex(self.component_dir) self.tokenizer = self._load_tokenizer(tokenizer_dir) def _load_tokenizer(self, tokenizer_dir: str | Path | None): @@ -294,28 +307,31 @@ def encode_tokens(self, token_ids: list[int]) -> tuple[np.ndarray, np.ndarray]: ]) cos, sin = _mrope_cos_sin(positions, cfg) - # Embedding rows gathered individually; the (151936, 5120) table is - # never fully materialized. - rows = [] - for token in token_ids: - key = "model.language_model.embed_tokens.weight" - rows.append(self.index.get_row(key, token)) - hidden = mx.array(np.stack(rows).astype(np.float32)) - del rows - gc.collect() + hidden = self._embed_tokens(token_ids) if cfg.num_layers < TEXT_ENCODER_LAYER: raise ValueError(f"Conditioner needs at least {TEXT_ENCODER_LAYER} layers, has {cfg.num_layers}.") for layer in range(TEXT_ENCODER_LAYER): hidden = self._decoder_layer(layer, hidden, cos, sin) - # Per-layer sync: without this the whole 50-layer graph accumulates - # and the machine runs out of memory (same failure mode as the DiT). + # Per-layer sync keeps the 50-layer activation graph bounded. mx.eval(hidden) gc.collect() tags = np.full((seq_len, ), 1, dtype=np.int64) # MINIMAX_H3_TEXT_TAG return np.asarray(hidden).astype(np.float32), tags + def _embed_tokens(self, token_ids: list[int]): + # Embedding rows gathered individually; the full table is never + # materialized by the BF16 streaming path. + rows = [] + for token in token_ids: + key = "model.language_model.embed_tokens.weight" + rows.append(self.index.get_row(key, token)) + hidden = mx.array(np.stack(rows).astype(np.float32)) + del rows + gc.collect() + return hidden + # -- layers ---------------------------------------------------------- def _decoder_layer(self, index: int, hidden, cos, sin): @@ -378,6 +394,96 @@ def close(self) -> None: self.index.close() +def unswizzle_nvfp4_scales(scale: np.ndarray, rows: int, cols: int) -> np.ndarray: + """FlashInfer 128x4 scale bytes -> MLX row-major group-16 scale bytes.""" + pad_rows, pad_cols = -(-rows // 128) * 128, -(-cols // 4) * 4 + if scale.size != pad_rows * pad_cols: + raise ValueError(f"NVFP4 scales need {pad_rows * pad_cols} bytes, got {scale.size}.") + tiles = scale.reshape(pad_rows // 128, pad_cols // 4, 32, 4, 4) + return np.ascontiguousarray(tiles.transpose(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols]) + + +class _ResidentNVFP4Index: + """Load the released 50-layer encoder without expanding packed matrices.""" + + def __init__(self, source: _ShardIndex): + self.weights: dict[str, mx.array | NVFP4Matrix] = {} + for key in sorted(source.key_to_shard): + if key.endswith(".weight_packed"): + prefix = key.removesuffix(".weight_packed") + packed = np.array(source.get(key), copy=True) + if packed.dtype != np.uint8 or packed.ndim != 2 or packed.shape[1] % 4: + raise ValueError(f"Invalid packed NVFP4 matrix {key}: {packed.shape}, {packed.dtype}") + rows, cols = packed.shape[0], packed.shape[1] * 2 + if cols % 16: + raise ValueError(f"NVFP4 input width must be divisible by 16: {key}") + scales = unswizzle_nvfp4_scales(source.get(prefix + ".weight_scale"), rows, cols // 16) + global_scale = float(source.get(prefix + ".weight_global_scale").reshape(-1)[0]) + if not np.isfinite(global_scale) or global_scale <= 0: + raise ValueError(f"Invalid NVFP4 global scale for {prefix}: {global_scale}") + weight = mx.array(packed).view(mx.uint32) + scale_bytes = mx.array(scales) + if not bool(mx.all(mx.isfinite(mx.from_fp8(scale_bytes, dtype=mx.float32)))): + raise ValueError(f"Non-finite NVFP4 block scales for {prefix}") + mx.eval(weight, scale_bytes) + self.weights[prefix + ".weight"] = NVFP4Matrix(weight, scale_bytes, global_scale) + elif key.endswith(".weight"): + if key == "model.language_model.embed_tokens.weight": + shard = source.key_to_shard[key] + header, data_start = source._cache_header(shard) + if header[key]["dtype"] == "BF16": + value = mx.array(_read_bf16_words(shard, key, header, data_start)).view(mx.bfloat16) + else: + value = mx.array(source.get(key)) + else: + value = source.get_mlx(key) + mx.eval(value) + self.weights[key] = value + source.close() + + def get_mlx(self, key: str): + return self.weights[key] + + def close(self) -> None: + self.weights.clear() + gc.collect() + + +class ResidentNVFP4MiniMaxH3TextConditioner(StreamedMiniMaxH3TextConditioner): + """Released NVFP4 encoder weights with floating-point MLX activations. + + The packed weights and embedding table stay resident. CUDA quantizes + activations to FP4; this path keeps FP32 activations, so hidden states are + not expected to be bit-exact with the CUDA encoder. + """ + + def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): + raw = json.loads((Path(component_dir) / "config.json").read_text()) + quant = raw.get("quantization_config", {}) + expected = { + "quant_method": "nvfp4", + "fmt": "e2m1", + "group_size": 16, + "scale_fmt": "e4m3", + "scale_layout": "128x4", + "activation_scheme": "dynamic" + } + if any(quant.get(key) != value for key, value in expected.items()): + raise ValueError("MLX NVFP4 conditioning requires the FastVideo group-16, 128x4 encoder export.") + # Fail on an older MLX before reading the encoder's large shards. + try: + packed, scales = mx.quantize(mx.ones((1, 64)), mode="nvfp4") + mx.eval(mx.quantized_matmul(mx.ones((1, 64)), packed, scales, mode="nvfp4")) + except (ValueError, RuntimeError) as error: + raise RuntimeError("Native NVFP4 conditioning requires an MLX build with nvfp4 matmul support.") from error + super().__init__(component_dir, tokenizer_dir) + self.index = _ResidentNVFP4Index(self.index) + + def _embed_tokens(self, token_ids: list[int]): + table = self.index.get_mlx("model.language_model.embed_tokens.weight") + return table[mx.array(token_ids, dtype=mx.int32)].astype(mx.float32) + + def _apply_mrope(q_or_k, cos, sin): """q_or_k: (S, H, D); cos/sin: (S, 1, D).""" half = q_or_k.shape[-1] // 2 diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index c6ca112f32..d55131e49c 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -49,6 +49,7 @@ MINIMAX_H3_MIN_DURATION, MINIMAX_H3_VIDEO_SHIFT, MiniMaxH3SchedulerState, + _eval_value, adaln_timestep_union, align_num_frames, audio_latent_num_frames, @@ -328,6 +329,8 @@ def __init__( video_decode_backend: str = "h3-vae", taeh3_checkpoint: str | Path | None = None, taeh3_chunk_size: int = 5, + conditioner_mode: str = "auto", + resident: bool = False, ) -> None: import mlx.core as mx @@ -343,6 +346,13 @@ def __init__( except Exception as error: # noqa: BLE001 - best effort on older MLX logger.info("Could not raise the Metal wired limit: %s", error) self.model_root = Path(model_root) + if conditioner_mode not in ("auto", "streamed", "nvfp4"): + raise ValueError(f"Unknown H3 conditioner mode: {conditioner_mode}") + if resident and video_decode_backend != "h3-vae": + raise ValueError("Resident H3 generation requires the H3 video VAE.") + self.conditioner_mode = conditioner_mode + self.resident = resident + self._resident_components: dict[str, Any] = {} self.dit_checkpoint = Path(mlx_dit_checkpoint) self.vae_dtype = vae_dtype if video_decode_backend not in ("h3-vae", "taeh3"): @@ -414,20 +424,70 @@ def resolve_geometry( # -- phase 1: conditioning ------------------------------------------- + def prepare_resident(self) -> None: + """Load the encoder, DiT, and both decoders once, before timed requests.""" + if not self.resident or self._resident_components: + return + from fastvideo.mlx_runtime.minimax_h3_conditioner import ResidentNVFP4MiniMaxH3TextConditioner + from fastvideo.mlx_runtime.minimax_h3_audio_vae import mlx_h3_audio_vae_from_dir + from fastvideo.mlx_runtime.minimax_h3_video_vae import mlx_h3_video_vae_from_dir + + import mlx.core as mx + + try: + conditioner = self._load_conditioner() + if not isinstance(conditioner, ResidentNVFP4MiniMaxH3TextConditioner): + conditioner.close() + raise ValueError("All-resident generation requires the packed NVFP4 text encoder.") + self._resident_components["conditioner"] = conditioner + dit = load_mlx_h3_checkpoint(self.dit_checkpoint) + self._resident_components["dit"] = dit + for group in [dit.weights, *dit.blocks, *dit.refiner]: + for value in group.values(): + _eval_value(value) + cache = dit._adaln_cache + if cache is not None: + mx.eval(cache.block_tables, cache.norm_out_shift, cache.norm_out_scale) + self._resident_components["video_vae"] = mlx_h3_video_vae_from_dir(self.model_root / "vae", + include_encoder=False, + storage_dtype=self.vae_dtype) + self._resident_components["audio_vae"] = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", + include_encoder=False) + mx.eval(list(self._resident_components["audio_vae"].weights.values())) + logger.info("H3 components resident: %.2f GiB active MLX memory", mx.get_active_memory() / 2**30) + except Exception: + self.close() + raise + + def close(self) -> None: + conditioner = self._resident_components.get("conditioner") + if conditioner is not None: + conditioner.close() + self._resident_components.clear() + _cleanup_mlx() + def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: """Returns (hidden states (S, hidden), token tags). Uses the cache or the streamed conditioner.""" cache_key = None if self.prompt_cache_dir is not None: - cache_key = prompt_cache_path(self.prompt_cache_dir, self.model_root, prompt) + identity = (f"{self.model_root}::conditioner=" + f"{getattr(self, 'conditioner_dir', self.model_root / 'text_encoder')}::" + f"{getattr(self, 'conditioner_mode', 'auto')}") + cache_key = prompt_cache_path(self.prompt_cache_dir, identity, prompt) if cache_key.exists(): data = np.load(cache_key) logger.info("Loaded prompt embeddings from cache %s", cache_key) return data["hidden_states"], data["token_tags"] - conditioner = self._load_conditioner() + if getattr(self, "resident", False): + self.prepare_resident() + conditioner = self._resident_components["conditioner"] + else: + conditioner = self._load_conditioner() hidden, tags = conditioner.encode_prompt(prompt) - conditioner.close() + if not getattr(self, "resident", False): + conditioner.close() _cleanup_mlx() if cache_key is not None: cache_key.parent.mkdir(parents=True, exist_ok=True) @@ -450,8 +510,17 @@ def has_conditioner_weights(self) -> bool: return marker.exists() or single.exists() def _load_conditioner(self): - from fastvideo.mlx_runtime.minimax_h3_conditioner import StreamedMiniMaxH3TextConditioner + from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + ResidentNVFP4MiniMaxH3TextConditioner, + StreamedMiniMaxH3TextConditioner, + ) + config = json.loads((self.conditioner_dir / "config.json").read_text()) + packed = config.get("quantization_config", {}).get("quant_method") == "nvfp4" + if self.conditioner_mode == "nvfp4" or (self.conditioner_mode == "auto" and packed): + return ResidentNVFP4MiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + if packed: + raise ValueError("The streamed conditioner requires BF16 weights; use conditioner_mode='nvfp4'.") return StreamedMiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) # -- phase 2: denoise -------------------------------------------------- @@ -478,6 +547,10 @@ def denoise( geometry = self.resolve_geometry(height, width, num_frames, enforce_duration=audio_num_frames is None) audio_frames = geometry["num_frames"] if audio_num_frames is None else align_num_frames(audio_num_frames) + if dit is None and getattr(self, "resident", False): + _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) + self.prepare_resident() + dit = self._resident_components["dit"] owned_dit = dit is None if owned_dit: _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) @@ -624,7 +697,13 @@ def decode_video(self, raise RuntimeError(f"TAEH3 produced unexpected frame shape: {frames.shape}") _cleanup_mlx() return frames - vae = mlx_h3_video_vae_from_dir(self.model_root / "vae", include_encoder=False, storage_dtype=self.vae_dtype) + if getattr(self, "resident", False): + self.prepare_resident() + vae = self._resident_components["video_vae"] + else: + vae = mlx_h3_video_vae_from_dir(self.model_root / "vae", + include_encoder=False, + storage_dtype=self.vae_dtype) expected_height = height // vae.spatial_compression_ratio expected_width = width // vae.spatial_compression_ratio if (geometry["latent_height"], geometry["latent_width"]) != (expected_height, expected_width): @@ -666,7 +745,11 @@ def decode_audio(self, audio_rows: np.ndarray, *, num_frames: int) -> np.ndarray num_audio_latents = audio_latent_num_frames(align_num_frames(num_frames)) latents = unpack_audio_tokens(audio_rows, num_audio_latents) - vae = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", include_encoder=False) + if getattr(self, "resident", False): + self.prepare_resident() + vae = self._resident_components["audio_vae"] + else: + vae = mlx_h3_audio_vae_from_dir(self.model_root / "audio_vae", include_encoder=False) z = vae.denormalize_latents(mx.array(latents)) waveform = np.asarray(vae.decode(z))[:, 0, :] # (B, 1, S) -> (B, S) del vae, z diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py new file mode 100644 index 0000000000..1b9f7e2566 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py @@ -0,0 +1,136 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Native MLX NVFP4 encoder storage and residency regression checks.""" +from types import SimpleNamespace + +import numpy as np +import pytest + +mx = pytest.importorskip("mlx.core") +from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + NVFP4Matrix, ResidentNVFP4MiniMaxH3TextConditioner, _ResidentNVFP4Index, + _ShardIndex, unswizzle_nvfp4_scales, +) +from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline + + +def _swizzle(values): + # Independent coordinate mapping of FlashInfer's 128-row, four-column tiles. + rows, cols = values.shape + padded = np.zeros((-(-rows // 128) * 128, -(-cols // 4) * 4), np.uint8) + padded[:rows, :cols] = values + output = np.empty(padded.size, np.uint8) + for r in range(padded.shape[0]): + for c in range(padded.shape[1]): + address = ((((r // 128) * (padded.shape[1] // 4) + c // 4) * 32 + + r % 32) * 4 + (r % 128) // 32) * 4 + c % 4 + output[address] = padded[r, c] + return output + + +def _decode_e4m3(values): + sign = np.where(values & 128, -1.0, 1.0) + exponent = (values >> 3) & 15 + fraction = values & 7 + return sign * np.where(exponent == 0, fraction * 2.0**-9, + (1.0 + fraction / 8.0) * 2.0**(exponent.astype(int) - 7)) + + +def test_padded_scale_layout_round_trip(): + rng = np.random.default_rng(11) + values = rng.integers(0, 127, (140, 7), dtype=np.uint8) + np.testing.assert_array_equal(unswizzle_nvfp4_scales(_swizzle(values), 140, 7), values) + with pytest.raises(ValueError, match="bytes"): + unswizzle_nvfp4_scales(np.zeros(1, np.uint8), 140, 7) + + +@pytest.mark.parametrize("global_scale", [0.5, 4.0]) +def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, global_scale): + from safetensors.numpy import save_file + + rng = np.random.default_rng(3) + packed = rng.integers(0, 256, (128, 64), dtype=np.uint8) + scales = rng.integers(24, 96, (128, 8), dtype=np.uint8) + prefix = "model.language_model.layers.0.self_attn.q_proj" + save_file({prefix + ".weight_packed": packed, + prefix + ".weight_scale": _swizzle(scales), + prefix + ".weight_global_scale": np.array([global_scale], np.float32)}, + tmp_path / "model.safetensors") + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + weight = index.get_mlx(prefix + ".weight") + assert isinstance(weight, NVFP4Matrix) + assert weight.weight.dtype == mx.uint32 + assert weight.scales.dtype == mx.uint8 + lut = np.array([0, .5, 1, 1.5, 2, 3, 4, 6, 0, -.5, -1, -1.5, -2, -3, -4, -6], np.float32) + dense = np.stack((lut[packed & 15], lut[packed >> 4]), axis=-1).reshape(128, 128) + dense *= np.repeat(_decode_e4m3(scales), 16, axis=1) / global_scale + x = rng.standard_normal((3, 128)).astype(np.float32) + np.testing.assert_allclose(np.array(weight.matmul(mx.array(x))), x @ dense.T, rtol=3e-5, atol=1e-3) + index.close() + assert not index.weights + + +def test_resident_embedding_keeps_bf16_storage(tmp_path): + torch = pytest.importorskip("torch") + from safetensors.torch import save_file + + key = "model.language_model.embed_tokens.weight" + table = torch.arange(60).reshape(10, 6).to(torch.bfloat16) + save_file({key: table}, tmp_path / "model.safetensors") + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + assert index.get_mlx(key).dtype == mx.bfloat16 + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + conditioner.index = index + np.testing.assert_array_equal(np.array(conditioner._embed_tokens([7, 1])), table[[7, 1]].float().numpy()) + conditioner.close() + + +def _pipeline(): + pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) + pipeline.resident = True + pipeline._resident_components = {} + pipeline.dit_checkpoint = "tiny" + pipeline.model_root = __import__("pathlib").Path("tiny") + pipeline.vae_dtype = "fp16" + return pipeline + + +def test_resident_preload_reuses_models_and_evaluates_audio(monkeypatch): + import fastvideo.mlx_runtime.minimax_h3_pipeline as module + import fastvideo.mlx_runtime.minimax_h3_audio_vae as audio + import fastvideo.mlx_runtime.minimax_h3_video_vae as video + + pipeline = _pipeline() + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + conditioner.index = SimpleNamespace(close=lambda: None) + monkeypatch.setattr(pipeline, "_load_conditioner", lambda: conditioner) + calls = [] + def load_dit(path): + calls.append(path) + return SimpleNamespace(weights={"x": mx.ones((1,))}, blocks=[], refiner=[], _adaln_cache=None) + monkeypatch.setattr(module, "load_mlx_h3_checkpoint", load_dit) + monkeypatch.setattr(video, "mlx_h3_video_vae_from_dir", lambda *a, **k: object()) + decoder = SimpleNamespace(weights={"x": mx.ones((2,))}) + monkeypatch.setattr(audio, "mlx_h3_audio_vae_from_dir", lambda *a, **k: decoder) + pipeline.prepare_resident() + pipeline.prepare_resident() + assert calls == ["tiny"] + assert set(pipeline._resident_components) == {"conditioner", "dit", "video_vae", "audio_vae"} + pipeline.close() + assert not pipeline._resident_components + + +def test_failed_preload_releases_encoder(monkeypatch): + import fastvideo.mlx_runtime.minimax_h3_pipeline as module + + pipeline = _pipeline() + conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) + closed = [] + conditioner.index = SimpleNamespace(close=lambda: closed.append(True)) + monkeypatch.setattr(pipeline, "_load_conditioner", lambda: conditioner) + def fail(path): + raise RuntimeError("out of memory") + monkeypatch.setattr(module, "load_mlx_h3_checkpoint", fail) + with pytest.raises(RuntimeError, match="out of memory"): + pipeline.prepare_resident() + assert closed == [True] + assert not pipeline._resident_components From cfe9d5f41ae90a759448d9d79330f57480546ee2 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:21:39 -0700 Subject: [PATCH 072/107] [feat]: enable experimental MXFP8 H3 conversion --- scripts/checkpoint_conversion/convert_minimax_h3_mlx.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index 1e2813e6ab..4abcac2805 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -23,6 +23,10 @@ quantization. This saves download bytes but quantizes twice; compare its clips with the BF16-sourced export before using it for release. +``--formats mxfp8`` tries native MLX FP8 weight storage and matrix +multiplication. It is experimental and requires operator support from the +installed MLX build; the default formats remain affine INT8/INT6/INT4. + python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ --out ~/models/FastH3-MLX-vsa \\ @@ -58,8 +62,8 @@ logger = init_logger(__name__) -SUPPORTED_FORMATS = ("int8", "int6", "int4") -DEFAULT_FORMATS = " ".join(SUPPORTED_FORMATS) +SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8") +DEFAULT_FORMATS = "int8 int6 int4" def _adaln_cache_timesteps(model_root: str | Path | None = None) -> np.ndarray: From aef7cabf3130eab9b1d6307e3d9d858d45c80bf3 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:24:28 -0700 Subject: [PATCH 073/107] [test]: check MXFP8 H3 checkpoint round trip --- .../tests/test_minimax_h3_fp8_checkpoint.py | 27 +++++++++++++++++++ 1 file changed, 27 insertions(+) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py new file mode 100644 index 0000000000..a0dc7b3216 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py @@ -0,0 +1,27 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exercise native MXFP8 H3 checkpoint storage on supported MLX builds.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +from fastvideo.mlx_runtime.fastwan import MLXQuantizationSpec, ensure_quantization_supported, linear, quantize_matrix +from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, save_mlx_h3_checkpoint + + +def test_mxfp8_checkpoint_preserves_quantized_matrix(tmp_path): + spec = MLXQuantizationSpec.from_name('mxfp8') + ensure_quantization_supported(spec) + weight = quantize_matrix(mx.random.normal((64, 64)).astype(mx.bfloat16), spec) + x = mx.random.normal((3, 64)).astype(mx.bfloat16) + config = dict(hidden_size=64, num_attention_heads=1, attention_head_dim=64, ffn_dim=128, + in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=64, + freq_dim=64, time_embed_dim=64, rope_freq_dim=4, rope_theta=10000., + norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5) + dit = MLXMiniMaxH3DiT({'test.weight': weight}, [], [], config) + save_mlx_h3_checkpoint(dit, tmp_path) + loaded = load_mlx_h3_checkpoint(tmp_path) + actual = linear(x, loaded.weights['test.weight']).astype(mx.float32) + expected = linear(x, weight).astype(mx.float32) + np.testing.assert_array_equal(np.array(actual), np.array(expected)) + assert loaded.weights['test.weight'].biases is None + assert loaded.weights['test.weight'].spec == spec From d14f6663ab9b7461c9173120afe968ba009661c0 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 16:35:38 -0700 Subject: [PATCH 074/107] [bugfix]: omit unused layers from single-shard MLX encoder --- .../mlx_runtime/minimax_h3_conditioner.py | 22 +++++++++---------- .../tests/test_minimax_h3_resident_nvfp4.py | 12 ++++++++++ 2 files changed, 23 insertions(+), 11 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index d5fe8b34b8..f9b1fe5fca 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -78,18 +78,18 @@ def __init__(self, component_dir: Path): index_path = component_dir / "model.safetensors.index.json" self.key_to_shard: dict[str, str] = {} self._header_cache: dict[str, tuple[dict, int]] = {} - if index_path.exists(): - weight_map = json.loads(index_path.read_text())["weight_map"] - def needed(key: str) -> bool: - if key == "model.language_model.embed_tokens.weight": - return True - prefix = "model.language_model.layers." - if not key.startswith(prefix): - return False - layer = key[len(prefix):].split(".", 1)[0] - return layer.isdigit() and int(layer) < TEXT_ENCODER_LAYER + def needed(key: str) -> bool: + if key == "model.language_model.embed_tokens.weight": + return True + prefix = "model.language_model.layers." + if not key.startswith(prefix): + return False + layer = key[len(prefix):].split(".", 1)[0] + return layer.isdigit() and int(layer) < TEXT_ENCODER_LAYER + if index_path.exists(): + weight_map = json.loads(index_path.read_text())["weight_map"] self.key_to_shard = {k: str(component_dir / s) for k, s in weight_map.items() if needed(k)} else: single = component_dir / "model.safetensors" @@ -99,7 +99,7 @@ def needed(key: str) -> bool: with open(single, "rb") as handle: (header_len, ) = struct.unpack(" Date: Sat, 3 Oct 2026 21:12:36 -0700 Subject: [PATCH 075/107] [test]: verify rank-16 H3 modulation and cache reload --- .../tests/test_minimax_h3_rank16_adaln.py | 72 +++++++++++++++++++ 1 file changed, 72 insertions(+) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py new file mode 100644 index 0000000000..cd954981b6 --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_rank16_adaln.py @@ -0,0 +1,72 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Check the pruned model's shared rank-16 modulation against NumPy math.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +from fastvideo.mlx_runtime.fastwan import timestep_embedding +from fastvideo.mlx_runtime.minimax_h3 import ( + MLXMiniMaxH3DiT, + MiniMaxH3SchedulerState, + adaln_timestep_union, + load_mlx_h3_checkpoint, + save_mlx_h3_checkpoint, +) + + +def test_rank16_cache_has_one_silu_before_shared_basis(tmp_path): + rng = np.random.default_rng(2026) + hidden, rank = 32, 16 + + def array(shape): + return mx.array(rng.normal(0, 0.1, shape).astype(np.float32)) + + config = dict(hidden_size=hidden, num_attention_heads=1, attention_head_dim=hidden, ffn_dim=64, + in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=hidden, + freq_dim=hidden, time_embed_dim=hidden, rope_freq_dim=4, rope_theta=10000., + norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5, adaln_rank=rank, num_layers=42) + weights = { + 'time_embedder.linear_1.weight': array((hidden, hidden)), + 'time_embedder.linear_1.bias': array((hidden,)), + 'time_embedder.linear_2.weight': array((hidden, hidden)), + 'time_embedder.linear_2.bias': array((hidden,)), + 'adaln_basis.weight': array((rank, hidden)), + 'norm_out.linear.weight': array((2 * hidden, rank)), + 'norm_out.linear.bias': array((2 * hidden,)), + } + blocks = [{'attn.to_q.weight': array((hidden, hidden)), + 'adaln_proj.linear.weight': array((18 * hidden, rank)), + 'adaln_proj.linear.bias': array((18 * hidden,))} for _ in range(42)] + dit = MLXMiniMaxH3DiT(weights, blocks, [], config) + rungs = [999, 874, 749, 624, 500, 375, 250, 125] + timesteps = adaln_timestep_union(MiniMaxH3SchedulerState.from_dmd_steps(10, rungs), + MiniMaxH3SchedulerState.from_dmd_steps(3, rungs)) + + def project(x, weight, bias=None): + result = x @ np.array(weight).T + return result if bias is None else result + np.array(bias) + + def silu(x): + return x / (1 + np.exp(-x)) + + features = np.array(timestep_embedding(mx.array(timesteps), hidden)) + first = project(features, weights['time_embedder.linear_1.weight'], weights['time_embedder.linear_1.bias']) + second = project(silu(first), weights['time_embedder.linear_2.weight'], weights['time_embedder.linear_2.bias']) + expected_basis = project(silu(second), weights['adaln_basis.weight']) + np.testing.assert_allclose(np.array(dit.compute_temb(mx.array(timesteps))), expected_basis, atol=1e-6) + expected_blocks = [project(expected_basis, block['adaln_proj.linear.weight'], + block['adaln_proj.linear.bias']).reshape(-1, 6 * hidden) for block in blocks] + expected_out = project(expected_basis, weights['norm_out.linear.weight'], weights['norm_out.linear.bias']) + cache = dit.precompute_adaln(timesteps) + for tables, expected in zip(cache.block_tables, expected_blocks, strict=True): + np.testing.assert_allclose(np.concatenate([np.array(t) for t in tables], axis=-1), expected, atol=1e-6) + np.testing.assert_allclose(np.array(cache.norm_out_shift), expected_out[:, :hidden], atol=1e-6) + np.testing.assert_allclose(np.array(cache.norm_out_scale), expected_out[:, hidden:], atol=1e-6) + assert all(block['adaln_proj.linear.weight'] is None for block in blocks) + + save_mlx_h3_checkpoint(dit, tmp_path) + loaded = load_mlx_h3_checkpoint(tmp_path) + assert loaded.adaln_rank == rank + assert len(loaded.blocks) == 42 + np.testing.assert_array_equal(loaded._adaln_cache.timesteps, timesteps) + np.testing.assert_array_equal(np.array(loaded._adaln_cache.norm_out_scale), np.array(cache.norm_out_scale)) From 7735e125cdfcbb04d1e8b444d99a3f710f9d7555 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 21:32:08 -0700 Subject: [PATCH 076/107] [bugfix]: fence NVFP4 activation quantization on DGX Spark --- fastvideo/layers/quantization/nvfp4_config.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 3e683f749c..36171b1f47 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -303,6 +303,11 @@ def _nvfp4_quantize( x_for_quant = F.pad(x, (0, 0, 0, pad_rows)) quantized, scales = torch.ops.fastvideo_fp4.nvfp4_quantize(x_for_quant, global_sf, sf_layout, do_shuffle) + if x.is_cuda and torch.cuda.get_device_capability(x.device) == (12, 1): + # On GB10 with FlashInfer 0.6.18, identical native H3 requests can + # diverge unless quantization completes before its padded input is + # released. Fence this boundary rather than every CUDA launch. + torch.cuda.current_stream(x.device).synchronize() if sf_layout != SfLayout.layout_linear.value: quantized = quantized.narrow(0, 0, logical_rows) return quantized, scales From 030100897e3b33831cf5bb1032c13da8e384bc25 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 21:37:56 -0700 Subject: [PATCH 077/107] [docs]: explain calibrated and reproducible Spark NVFP4 inference --- .../installation/spark_performance.md | 20 ++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index 00c9080d2c..8a54e867ef 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -213,9 +213,9 @@ A few things that surprise people on this box (beyond the memory notes above): This recipe uses the full V2 eight-forward transformer, the 50-layer NVFP4 Qwen3-VL encoder, and the light H3 video VAE. Its configuration keeps all -three resident on one GB10. Runtime, memory fit, and quality still need a run -on that device. The earlier bf16 H3 memory guidance above concerns a larger -checkpoint. +three resident on one GB10. This stack fits in the Spark's unified memory; +benchmark your installed runtime and review the clips before publishing a +speed claim. The earlier bf16 H3 memory guidance above concerns a larger checkpoint. Install FastVideo from a checkout that includes the ModelOpt converter and FlashInfer FP4 support, following [the Spark install guide](spark.md). Sign in @@ -255,10 +255,20 @@ copy the resulting `transformer/` directory to the Spark. Do not omit recipe's `num_inference_steps: 9` means nine sigma points and eight DiT forwards. +The converter preserves ModelOpt's calibrated `input_scale` as the reciprocal +`_nvfp4_input_global_sf`. Reconvert older exports that discarded this scale: +unit activation scaling clips inputs above 2688. An explicit `--act-amax` +table overrides the source calibration. + +On GB10 with FlashInfer 0.6.18, FastVideo fences activation quantization before +releasing its padded input. Without this completion fence, identical H3 +requests produced different DiT latents and occasionally corrupt video. +The fence applies to `sm_121`; other architectures retain asynchronous execution. + Run `examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml` from the repository root. It uses 832x480, 243 frames, VSA sparsity 0.8 with -64-token tiles, and the full H3 VAE. It does not use frame dropping or spatial -upscaling. +64-token tiles, and the light H3 VAE through the `h3-vae` decode backend. +It does not use frame dropping or spatial upscaling. ```bash FASTVIDEO_MINIMAX_H3_FUSIONS=all \ From 58308798dd7a2c8cb23ae7609703bf75e82cd88c Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 21:46:41 -0700 Subject: [PATCH 078/107] [bugfix]: keep Spark NVFP4 fence opaque to Dynamo --- fastvideo/layers/quantization/nvfp4_config.py | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 36171b1f47..02763066c7 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -193,8 +193,18 @@ def _nvfp4_quantize_op( sf_layout: int, do_shuffle: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: + spark = torch.cuda.get_device_capability(x.device) == (12, 1) + if spark and torch.cuda.is_current_stream_capturing(): + raise RuntimeError("NVFP4 activation quantization on DGX Spark requires a completion fence; " + "disable CUDA graph capture.") SfLayout, _, nvfp4_quantize = _require_flashinfer() - return nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) + quantized, scales = nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) + if spark: + # With FlashInfer 0.6.18 on GB10, queued activation quantization + # plus GEMM can diverge. Completing quantization while its padded + # input is alive restores identical native H3 requests. + torch.cuda.current_stream(x.device).synchronize() + return quantized, scales @_nvfp4_quantize_op.register_fake def _nvfp4_quantize_op_fake( @@ -303,11 +313,6 @@ def _nvfp4_quantize( x_for_quant = F.pad(x, (0, 0, 0, pad_rows)) quantized, scales = torch.ops.fastvideo_fp4.nvfp4_quantize(x_for_quant, global_sf, sf_layout, do_shuffle) - if x.is_cuda and torch.cuda.get_device_capability(x.device) == (12, 1): - # On GB10 with FlashInfer 0.6.18, identical native H3 requests can - # diverge unless quantization completes before its padded input is - # released. Fence this boundary rather than every CUDA launch. - torch.cuda.current_stream(x.device).synchronize() if sf_layout != SfLayout.layout_linear.value: quantized = quantized.narrow(0, 0, logical_rows) return quantized, scales From ad06a38464c3a08e6653525cf196dea34a01cfaa Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 22:01:45 -0700 Subject: [PATCH 079/107] [bugfix]: order fresh NVFP4 global scales on DGX Spark --- fastvideo/layers/quantization/nvfp4_config.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 02763066c7..43d2fbd35c 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -198,7 +198,16 @@ def _nvfp4_quantize_op( raise RuntimeError("NVFP4 activation quantization on DGX Spark requires a completion fence; " "disable CUDA graph capture.") SfLayout, _, nvfp4_quantize = _require_flashinfer() - quantized, scales = nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) + if spark: + # FlashInfer's PDL kernel reads the global scale before its + # dependency wait. Fresh dynamic scales require normal ordering. + quantized, scales = nvfp4_quantize(x, + global_sf, + sfLayout=SfLayout(sf_layout), + do_shuffle=do_shuffle, + enable_pdl=False) + else: + quantized, scales = nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) if spark: # With FlashInfer 0.6.18 on GB10, queued activation quantization # plus GEMM can diverge. Completing quantization while its padded From c7969ea7753d5753fab8eb7b242f55a90549d953 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 3 Oct 2026 22:52:13 -0700 Subject: [PATCH 080/107] [bugfix]: carry H3 performance switches to Ray workers --- docs/contributing/env_vars.md | 3 +++ fastvideo/envs.py | 10 ++++++++++ fastvideo/layers/quantization/nvfp4_config.py | 3 +-- fastvideo/models/vaes/minimax_h3_video.py | 1 - .../tests/worker/test_ray_distributed_executor.py | 14 ++++++++++++++ 5 files changed, 28 insertions(+), 3 deletions(-) diff --git a/docs/contributing/env_vars.md b/docs/contributing/env_vars.md index 883e2d8ac8..535ad581a9 100644 --- a/docs/contributing/env_vars.md +++ b/docs/contributing/env_vars.md @@ -192,6 +192,7 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_FA4` | bool | `0` | attention | The FLASH_ATTN backend uses FlashAttention-4 (flash_attn.cute) instead of FA3 or FA2. | | `FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN` | bool | `0` | attention | MiniMax-H3 dense DiT self-attention uses the FlashAttention-4 packed-varlen entry point. This changes the floating-point reduction order, so it is an inference-only opt-in. | | `FASTVIDEO_VSA_SM100A` | bool | `0` | attention | VIDEO_SPARSE_ATTN_H3 sends no-grad tile-64 forwards to the data-center Blackwell (sm_100a) kernel. fastvideo-kernel reads the same variable with the same rule. | +| `FASTVIDEO_VSA_TRITON` | bool | `0` | attention | Force the Triton MiniMax-H3 sparse attention kernel. fastvideo-kernel reads the same variable. | | `FASTVIDEO_NVFP4_FA4` | bool | `0` | attention | FlashAttention-4 quantizes Q and K to NVFP4. An explicit nvfp4_fa4 attention implementation argument takes precedence. | | `FASTVIDEO_DISABLE_ATTENTION_COMPILE` | bool | `1` | attention | Keep attention forward out of torch.compile graphs (torch.compiler.disable). Set it to 0 to let attention constructed under that setting be traced. Setting it explicitly to true also blocks regional compile. | | `FASTVIDEO_MLX_WINDOW` | int | `0` | attention | MLX FastWan windowed attention size in tokens. 0 uses full attention. | @@ -201,6 +202,8 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_VAE_PARALLEL_ENCODE` | bool | `0` | performance | MiniMax-H3 reference-video VAE encode splits its temporal chunks across the sequence-parallel ranks. Same as FastVideoArgs.vae_parallel_encode=True. | | `FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY` | str | unset | performance | Collective that moves chunks in parallel VAE decode: gather (used when unset) or all_gather. | | `FASTVIDEO_MINIMAX_H3_FUSIONS` | str | `""` | performance | MiniMax-H3 inference-only Triton fusions: all, 1, or a comma-separated subset of modulate,qknorm_rope,swiglu. Empty, 0, or none keeps the eager implementation. | +| `FASTVIDEO_H3_VAE_TILE_BATCH` | int | `1` | performance | Spatial tiles per MiniMax-H3 light-VAE decoder call. Values below 1 use one tile. | +| `FASTVIDEO_NVFP4_MM_BACKEND` | str | `auto` | performance | FlashInfer NVFP4 matrix multiplication backend: auto, cutlass, cudnn, trtllm, or b12x. | | `FASTVIDEO_FSDP2_AUTOWRAP` | bool | `0` | performance | FSDP2 shards modules by parameter count instead of the model's shard conditions. Not supported by self-forcing distillation. | | `FASTVIDEO_FSDP2_MIN_PARAMS` | int | `10000000` | performance | Minimum parameter count of a module that FASTVIDEO_FSDP2_AUTOWRAP shards. | | `FASTVIDEO_MLX_COMPILE` | bool | `0` | performance | Compile the MLX DiT forward with mx.compile. | diff --git a/fastvideo/envs.py b/fastvideo/envs.py index 77ba871a67..c6ea13306c 100644 --- a/fastvideo/envs.py +++ b/fastvideo/envs.py @@ -331,6 +331,10 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="attention", doc="VIDEO_SPARSE_ATTN_H3 sends no-grad tile-64 forwards to the data-center Blackwell (sm_100a) kernel. " "fastvideo-kernel reads the same variable with the same rule.") +FASTVIDEO_VSA_TRITON = EnvBool( + False, + category="attention", + doc="Force the Triton MiniMax-H3 sparse attention kernel. fastvideo-kernel reads the same variable.") FASTVIDEO_NVFP4_FA4 = EnvBool( False, category="attention", @@ -378,6 +382,12 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="performance", doc="MiniMax-H3 inference-only Triton fusions: all, 1, or a comma-separated subset of " "modulate,qknorm_rope,swiglu. Empty, 0, or none keeps the eager implementation.") +FASTVIDEO_H3_VAE_TILE_BATCH = EnvInt( + 1, category="performance", doc="Spatial tiles per MiniMax-H3 light-VAE decoder call. Values below 1 use one tile.") +FASTVIDEO_NVFP4_MM_BACKEND = EnvStr( + "auto", + category="performance", + doc="FlashInfer NVFP4 matrix multiplication backend: auto, cutlass, cudnn, trtllm, or b12x.") FASTVIDEO_FSDP2_AUTOWRAP = EnvBool(False, category="performance", doc="FSDP2 shards modules by parameter count instead of the model's shard " diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 43d2fbd35c..837554c7e4 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -29,7 +29,6 @@ from torch.nn.parameter import Parameter import fastvideo.envs as envs - from fastvideo.layers.quantization.base_config import ( QuantizationConfig, QuantizeMethodBase, @@ -211,7 +210,7 @@ def _nvfp4_quantize_op( if spark: # With FlashInfer 0.6.18 on GB10, queued activation quantization # plus GEMM can diverge. Completing quantization while its padded - # input is alive restores identical native H3 requests. + # input is alive prevents the observed intermittent corruption. torch.cuda.current_stream(x.device).synchronize() return quantized, scales diff --git a/fastvideo/models/vaes/minimax_h3_video.py b/fastvideo/models/vaes/minimax_h3_video.py index fca78c244a..d2f37ce21f 100644 --- a/fastvideo/models/vaes/minimax_h3_video.py +++ b/fastvideo/models/vaes/minimax_h3_video.py @@ -16,7 +16,6 @@ from torch.utils.checkpoint import checkpoint import fastvideo.envs as envs - from fastvideo.attention import get_attn_backend from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig from fastvideo.platforms import AttentionBackendEnum diff --git a/fastvideo/tests/worker/test_ray_distributed_executor.py b/fastvideo/tests/worker/test_ray_distributed_executor.py index 0bf8b4d1e3..ffb3e48421 100644 --- a/fastvideo/tests/worker/test_ray_distributed_executor.py +++ b/fastvideo/tests/worker/test_ray_distributed_executor.py @@ -36,3 +36,17 @@ def test_ray_log_queue_stays_on_the_driver() -> None: executor.clear_log_queue() assert executor._log_queue is None assert "log_queue" in signature(Executor.set_log_queue).parameters + + +def test_ray_carries_h3_performance_switches() -> None: + import fastvideo.envs as envs + from fastvideo.worker.ray_env import get_env_vars_to_copy + + with (envs.FASTVIDEO_H3_VAE_TILE_BATCH.override(8), + envs.FASTVIDEO_NVFP4_MM_BACKEND.override("cutlass"), + envs.FASTVIDEO_VSA_TRITON.override(True)): + copied = get_env_vars_to_copy() + assert {"FASTVIDEO_H3_VAE_TILE_BATCH", "FASTVIDEO_NVFP4_MM_BACKEND", "FASTVIDEO_VSA_TRITON"} <= copied + assert envs.FASTVIDEO_H3_VAE_TILE_BATCH.get() == 8 + assert envs.FASTVIDEO_NVFP4_MM_BACKEND.get() == "cutlass" + assert envs.FASTVIDEO_VSA_TRITON.get() is True From fd8b213aa34100806f0e62461b31c17c5ae6dfb2 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 08:44:15 -0700 Subject: [PATCH 081/107] [docs]: record verified Track C Spark release timings --- .../installation/spark_performance.md | 36 +++++++++++++++++++ 1 file changed, 36 insertions(+) diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index 8a54e867ef..581f10f9d8 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -314,6 +314,42 @@ sparsity. The V2 converter packs VSA gates, so its `h3_dit_vsa` profile must match the recipe. A later pruned NVFP4 transformer uses the separate `h3_dit_ffn` profile, with attention and VSA gates left dense. +### Native FastH3 release measurements + +Measured on October 4, 2026, with the trained eight-forward ladder, VSA 0.8, +832x480 native video, the 50-layer NVFP4 encoder and light video/audio VAEs +resident. Each cell is the median of two timed calls after one warmup, in +seconds, for `latency-ceramics-005` / `latency-harbor-005`, seed 2026. + +| Model | Frames | One Spark | Two Sparks, SP2 | +|---|---:|---:|---:| +| Pruned ckpt300, FFN NVFP4 | 124 | 134.361 / 134.675 | **78.245 / 78.272** | +| Pruned ckpt300, FFN NVFP4 | 243 | 284.188 / 280.593 | **165.987 / 163.001** | +| Full V2, NVFP4 FFN/attention/gates | 124 | 142.460 / 140.425 | **88.356 / 86.101** | +| Full V2, NVFP4 FFN/attention/gates | 243 | 308.204 / 306.141 | **179.876 / 180.213** | + +The base integrates upstream main `0cc41a22` with experimental NVFP4 support. +One-Spark tested commits are `6d6b57fe` (pruned 124), `3f24557a` (pruned 243, +with `CUDA_LAUNCH_BLOCKING=1`), `f5126f78` (V2 124) and `6e9d7a0b` (V2 243). +The final pair uses `715d4a5f`, with matching actual-worker code fingerprints, +CUTLASS FP4 GEMMs and Triton VSA. Light-VAE tile batch is 8 on one Spark and 1 +on the pair; pair batch 8 was slower at 124 frames (79.993 / 80.022 s pruned). +All offload/deferred-loading and compile options are disabled. The pair uses +QSFP RoCE, SP2/TP1 and parallel VAE gathering. See the pair configs +`basic_fasth3_spark_pair_pruned_nvfp4.yaml` and +`basic_fasth3_spark_pair_v2_nvfp4.yaml` beside the benchmark script. + +Every final pair warmup and repeat has correct dimensions/frame count, coherent +sampled frames and identical full decoded-video hashes within its prompt. +The V2 one-Spark 124-frame harbor warmup differs from the timed clips but remains +coherent. These checks establish repeat reliability for the tested recipes; +BF16 reference parity, speech accuracy and lip sync need separate review. + +The [older H3 local blog](https://haoailab.com/blogs/fasth3-local/) reports +243 s on one Spark and 209 s on two Sparks at 124 frames. It uses the four-step +Preview checkpoint and full VAE, so the old and new values are context, not a +matched optimization comparison. It has no matching 243-frame baseline. + ## Reproduce these numbers Two scripts under `examples/inference/optimizations/` reproduce the claims on From 532fd6b2a5808fed2b559260f84daff2434bd550 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 08:46:32 -0700 Subject: [PATCH 082/107] [perf]: bound MLX H3 sparse attention gather memory --- fastvideo/mlx_runtime/minimax_h3_vsa.py | 4 ++-- .../test_minimax_h3_vsa_gather_budget.py | 24 +++++++++++++++++++ 2 files changed, 26 insertions(+), 2 deletions(-) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa.py b/fastvideo/mlx_runtime/minimax_h3_vsa.py index 3c49969e5c..48ca15fa32 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa.py @@ -544,11 +544,11 @@ def _key_valid_mask(block_idx, variable_block_sizes, tile_elems: int): return offsets[None, None, None, :] < selected_sizes[:, :, :, None] -_REFERENCE_GATHER_TARGET_BYTES = 2 * 1024**3 +_REFERENCE_GATHER_TARGET_BYTES = 256 * 1024**2 def _reference_gather_query_chunk(heads: int, dim: int, k_sel: int, tile_elems: int, n_q: int) -> int: - """Batch as many query tiles as fit in ~2 GiB of gathered BF16 K/V.""" + """Bound gathered BF16 K/V to 256 MiB, leaving space for resident weights.""" bytes_per_query = 4 * heads * max(k_sel, 1) * tile_elems * dim chunk = min(n_q, max(1, _REFERENCE_GATHER_TARGET_BYTES // max(bytes_per_query, 1))) return int(chunk) diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py new file mode 100644 index 0000000000..376478624a --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py @@ -0,0 +1,24 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Changing the gather memory budget must preserve selected-tile attention.""" +import numpy as np +import pytest + +mx = pytest.importorskip('mlx.core') +import fastvideo.mlx_runtime.minimax_h3_vsa as vsa + + +def test_gather_budget_preserves_attention_with_partial_tiles(monkeypatch): + geom = vsa.build_h3_tile_geometry((7, 67), (6, 8, 8), 64) + rng = np.random.default_rng(2026) + shape = (geom.padded_length, 2, 128) + q, k, value = [mx.array(rng.normal(size=shape).astype(np.float32)).astype(mx.bfloat16) for _ in range(3)] + # Reverse video order also exercises the selected-key order, not a dense mask. + selected = np.array([0, 1, geom.num_tiles - 1, geom.num_prefix_tiles], dtype=np.int32) + idx = mx.array(np.broadcast_to(selected, (2, geom.num_video_tiles, selected.size)).copy()) + monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 2 * 1024**3) + expected = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) + mx.eval(expected) + monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 1) + actual = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) + mx.eval(actual) + np.testing.assert_array_equal(np.array(actual.astype(mx.float32)), np.array(expected.astype(mx.float32))) From e8bcd3127f8af35908b7f09edc5897bfca16b648 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 09:06:23 -0700 Subject: [PATCH 083/107] [perf]: batch 32 keys in opt-in MLX H3 sparse attention --- fastvideo/mlx_runtime/minimax_h3_vsa_simd.py | 41 +++++++++++-------- .../test_mlx_minimax_h3_vsa_regressions.py | 20 +++++++++ 2 files changed, 43 insertions(+), 18 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index bb21b1912a..972c1b118f 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -37,16 +37,17 @@ """ # One threadgroup = one (head, video query tile). 8 SIMD-groups x 32 = 256 -# threads cover 64 query rows. Q is half smem (16 KiB). K/V stage 8 keys. +# threads cover 64 query rows. K/V stage 32 keys, with 28.25 KiB total smem. +# Updating online softmax once per 32 keys reduces accumulator rescaling barriers. _SIMD_SOURCE = """ const int TILE = 64; const int D = 128; const int SG = 32; const int N_SG = 8; const int ROWS = 8; - const int KCHUNK = 8; - threadgroup float kvsmem[8 * 128]; - threadgroup float score_smem[8 * 64]; + const int KCHUNK = 32; + threadgroup float kvsmem[KCHUNK * D]; + threadgroup float score_smem[N_SG * ROWS * KCHUNK]; threadgroup float scale_tmp[8 * 64]; threadgroup float qtile[8 * 64]; threadgroup float row_alpha[8 * 8]; @@ -65,7 +66,7 @@ int qt = n_prefix + (int)q_tile; int q_valid = active ? vbs[qt] : 0; int q_base_tile = (((int)head * S) + qt * TILE) * D; - threadgroup float *sg_scores = score_smem + sid * 64; + threadgroup float *sg_scores = score_smem + sid * ROWS * KCHUNK; threadgroup float *sg_tmp = scale_tmp + sid * 64; threadgroup float *sg_qtile = qtile + sid * 64; threadgroup float *sg_alpha = row_alpha + sid * 8; @@ -117,16 +118,18 @@ } threadgroup_barrier(mem_flags::mem_threadgroup); - thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 kmat; - simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kk * 8), D, ulong2(0, 0), true); - simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); + for (int kc = 0; kc < KCHUNK; kc += 8) { + thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 kmat; + simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D, ulong2(0, 0), true); + simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); + } + simdgroup_store(smat, sg_scores + kc, KCHUNK); } - simdgroup_store(smat, sg_scores, KCHUNK); simdgroup_barrier(mem_flags::mem_threadgroup); - float scores[8]; + float scores[KCHUNK]; float cmax = -3.402823466e+38f; if (lane < (uint)ROWS) { int grow = qrow0 + (int)lane; @@ -178,12 +181,14 @@ } simdgroup_barrier(mem_flags::mem_threadgroup); - thread simdgroup_float8x8 pmat; - simdgroup_load(pmat, sg_scores, KCHUNK); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 vmat; - simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kk * 8), D); - simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); + for (int kc = 0; kc < KCHUNK; kc += 8) { + thread simdgroup_float8x8 pmat; + simdgroup_load(pmat, sg_scores + kc, KCHUNK); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 vmat; + simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D); + simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); + } } threadgroup_barrier(mem_flags::mem_threadgroup); } diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index ca3fe3d11f..f89cc584bd 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -328,3 +328,23 @@ def test_converter_continues_past_mismatched_existing_format(tmp_path, monkeypat converter.main() assert saved == ["int6"] assert (existing / h3.H3_WEIGHTS_FILENAME).read_bytes() == b"existing" + + +@pytest.mark.parametrize("dtype", [mx.float16, mx.bfloat16]) +@pytest.mark.parametrize("exempt", [False, True]) +def test_simd_partial_key_chunks_match_reference(dtype, exempt): + """Nonuniform scores and partial tiles exercise all four 8-key fragments.""" + _require_metal() + geometry = vsa.build_h3_tile_geometry((7, 5), (5, 3, 7), 64) + mx.random.seed(3026) + q, k, value = [mx.random.normal((geometry.total_seq_length, 2, 128)).astype(dtype) + for _ in range(3)] + expected = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="reference") + stats = vsa.MiniMaxH3VSAStats() + actual = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="simd", stats=stats) + mx.eval(actual, expected) + assert stats.impl == "simd" and stats.dense_fallback_reason is None + assert mx.all(mx.isfinite(actual)).item() + # The Metal kernel uses FP32 accumulation with a different reduction order. + np.testing.assert_allclose(np.asarray(actual.astype(mx.float32)), + np.asarray(expected.astype(mx.float32)), atol=.01, rtol=.01) From ef1f277dc7aa668eaad0a8505e21c8e83c8bc78c Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 09:23:13 -0700 Subject: [PATCH 084/107] [bugfix]: apply H3 wired residency through the MLX wired API --- fastvideo/mlx_runtime/minimax_h3_pipeline.py | 38 +++++++++++++------ .../mlx/test_mlx_minimax_h3_fast_mode.py | 24 ++++++++++++ 2 files changed, 50 insertions(+), 12 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index d55131e49c..be898f3bb8 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -232,7 +232,7 @@ def _cleanup_mlx() -> None: def _default_metal_wired_limit_gib(mx) -> float: - """Keep the default below both physical memory and the tested 30 GiB cap.""" + """Legacy helper for allocator capacity, not the wired-residency setting.""" metal = getattr(mx, "metal", None) if metal is None: return 30.0 @@ -245,6 +245,30 @@ def _default_metal_wired_limit_gib(mx) -> float: return min(30.0, 0.84 * total_bytes / 2**30) +def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> None: + """Keep allocator capacity separate from explicitly requested wired residency.""" + set_memory = getattr(mx, "set_memory_limit", None) + if set_memory is None and hasattr(mx, "metal"): + set_memory = getattr(mx.metal, "set_memory_limit", None) + if set_memory is not None: + try: + set_memory(int(_default_metal_wired_limit_gib(mx) * 2**30)) + except Exception as error: # noqa: BLE001 - older MLX best effort + logger.info("Could not set the Metal allocation limit: %s", error) + if wired_limit_gib is None: + return + if not math.isfinite(wired_limit_gib) or wired_limit_gib <= 0: + raise ValueError("metal_wired_limit_gib must be finite and positive") + set_wired = getattr(mx, "set_wired_limit", None) + if set_wired is None and hasattr(mx, "metal"): + set_wired = getattr(mx.metal, "set_wired_limit", None) + if set_wired is None: + raise RuntimeError("This MLX build cannot set the requested wired-memory limit") + # Explicit requests must succeed; do not silently benchmark an unwired model. + previous = set_wired(int(wired_limit_gib * 2**30)) + logger.info("MLX wired limit %.2f GiB (previous %.2f GiB)", wired_limit_gib, previous / 2**30) + + MINIMAX_H3_PROMPT_CACHE_VERSION = "v2-attention-layout" @@ -334,17 +358,7 @@ def __init__( ) -> None: import mlx.core as mx - set_limit = getattr(mx, "set_memory_limit", None) - if set_limit is None and hasattr(mx, "metal"): - set_limit = getattr(mx.metal, "set_memory_limit", None) - if set_limit is not None: - # Keep large resident models inside a predictable wired budget. - try: - if metal_wired_limit_gib is None: - metal_wired_limit_gib = _default_metal_wired_limit_gib(mx) - set_limit(int(metal_wired_limit_gib * 2**30)) - except Exception as error: # noqa: BLE001 - best effort on older MLX - logger.info("Could not raise the Metal wired limit: %s", error) + _configure_metal_memory_limits(mx, metal_wired_limit_gib) self.model_root = Path(model_root) if conditioner_mode not in ("auto", "streamed", "nvfp4"): raise ValueError(f"Unknown H3 conditioner mode: {conditioner_mode}") diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py index e009505547..6e59682545 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py @@ -21,6 +21,7 @@ MiniMaxH3MLXPipeline, _adaln_schedule_union, _center_crop_frames, + _configure_metal_memory_limits, _default_metal_wired_limit_gib, _preflight_media_dependencies, _validate_checkpoint_step_ladder, @@ -143,3 +144,26 @@ def fail(*_args, **_kwargs): assert not output.with_suffix(".tmp.mp4").exists() assert not output.with_suffix(".tmp.wav").exists() + + +def test_explicit_wired_limit_uses_wired_api_separately_from_allocator(): + calls = [] + fake = SimpleNamespace( + metal=SimpleNamespace(device_info=lambda: {"memory_size": 36 * 2**30}), + set_memory_limit=lambda size: calls.append(("allocator", size)), + set_wired_limit=lambda size: calls.append(("wired", size)) or 0, + ) + _configure_metal_memory_limits(fake, 27.0) + assert calls == [("allocator", 30 * 2**30), ("wired", 27 * 2**30)] + + +def test_explicit_wired_limit_failure_is_not_silently_ignored(): + def reject(size): + raise ValueError("exceeds system wired limit") + fake = SimpleNamespace(set_wired_limit=reject) + with pytest.raises(ValueError, match="system wired limit"): + _configure_metal_memory_limits(fake, 31.0) + with pytest.raises(ValueError, match="finite and positive"): + _configure_metal_memory_limits(fake, float("nan")) + with pytest.raises(RuntimeError, match="cannot set"): + _configure_metal_memory_limits(SimpleNamespace(), 27.0) From 6457d71791e1942127aaa1b3ee052c86062105a4 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 10:10:43 -0700 Subject: [PATCH 085/107] [docs]: explain H3 Metal wired-memory limits --- docs/getting_started/installation/mlx.md | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index 1f4974b98e..a2fb02206a 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -128,6 +128,27 @@ components. On a 36 GiB Mac, try INT6 first and measure peak allocation. If loading or inference runs out of memory, use phased loading by leaving `resident=False`. Changing placement does not change frames or resolution. +### Metal wired memory + +MLX's allocation limit and wired-memory limit are separate. H3's optional +`metal_wired_limit_gib` calls `mx.set_wired_limit` so selected Metal allocations +stay in physical memory. It does not increase available RAM. Explicit requests +fail visibly if the installed MLX build cannot apply them. + +Inspect the device's recommended working set before choosing a limit: + +```python +import mlx.core as mx + +print(mx.device_info()) +``` + +For the tested 36 GiB M4 Max, phased generation can request +`metal_wired_limit_gib=27` in `MiniMaxH3MLXPipeline`. Leave room for macOS and +other applications. All-resident generation also needs room for the encoder, +DiT, both decoders and peak activations; wiring cannot make an oversized stack +fit. An omitted wired limit preserves MLX's existing wiring setting. + ## Hardware - FastMetal 1.3B and 5B: 16 GB unified memory and up From ecc91b2cc68b5811f3208ded53e5ac16806e5c11 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 10:48:31 -0700 Subject: [PATCH 086/107] [feat]: cache packed H3 encoder in MLX layout --- .../mlx_runtime/minimax_h3_conditioner.py | 101 +++++++++++++++--- .../tests/test_minimax_h3_resident_nvfp4.py | 36 ++++++- .../convert_minimax_h3_mlx.py | 13 +++ 3 files changed, 133 insertions(+), 17 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index f9b1fe5fca..ccc1bfe7cc 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -403,6 +403,64 @@ def unswizzle_nvfp4_scales(scale: np.ndarray, rows: int, cols: int) -> np.ndarra return np.ascontiguousarray(tiles.transpose(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols]) +MLX_NVFP4_ENCODER_MANIFEST = "mlx_h3_nvfp4_encoder.json" + + +def _read_nvfp4_encoder_config(component_dir: str | Path) -> dict[str, Any]: + raw = json.loads((Path(component_dir) / "config.json").read_text()) + expected = { + "quant_method": "nvfp4", + "fmt": "e2m1", + "group_size": 16, + "scale_fmt": "e4m3", + "scale_layout": "128x4", + "activation_scheme": "dynamic", + } + quant = raw.get("quantization_config", {}) + if any(quant.get(key) != value for key, value in expected.items()): + raise ValueError("MLX NVFP4 conditioning requires the FastVideo group-16, 128x4 encoder export.") + return raw + + +def export_mlx_h3_nvfp4_encoder(component_dir: str | Path, output_dir: str | Path) -> Path: + """Cache the released packed encoder in MLX layout, without requantization. + + Keep original packed nibbles, row-major scale bytes, global scales and + embedding/norm values. Later loads skip CPU scale unswizzling and staging. + """ + raw = _read_nvfp4_encoder_config(component_dir) + output_dir = Path(output_dir) + if output_dir.exists() and any(output_dir.iterdir()): + raise FileExistsError(f"Encoder cache output must be empty: {output_dir}") + output_dir.mkdir(parents=True, exist_ok=True) + index = _ResidentNVFP4Index(_ShardIndex(Path(component_dir))) + try: + arrays = {} + matrices = {} + dense_keys = [] + for key, value in index.weights.items(): + if isinstance(value, NVFP4Matrix): + arrays[key] = value.weight + arrays[key + ".scales"] = value.scales + matrices[key] = {"global_scale": value.global_scale} + else: + arrays[key] = value + dense_keys.append(key) + mx.save_safetensors(str(output_dir / "model.safetensors"), arrays) + (output_dir / "config.json").write_text(json.dumps(raw, indent=2) + "\n") + manifest = { + "format_version": 1, + "language_layers": TEXT_ENCODER_LAYER, + "matrices": matrices, + "dense_keys": dense_keys, + "source_dir": str(Path(component_dir).resolve()) + } + (output_dir / MLX_NVFP4_ENCODER_MANIFEST).write_text(json.dumps(manifest, indent=2) + "\n") + finally: + index.close() + return output_dir + + class _ResidentNVFP4Index: """Load the released 50-layer encoder without expanding packed matrices.""" @@ -441,6 +499,30 @@ def __init__(self, source: _ShardIndex): self.weights[key] = value source.close() + @classmethod + def from_mlx_checkpoint(cls, component_dir: str | Path): + component_dir = Path(component_dir) + manifest = json.loads((component_dir / MLX_NVFP4_ENCODER_MANIFEST).read_text()) + if manifest.get("format_version") != 1 or manifest.get("language_layers") != TEXT_ENCODER_LAYER: + raise ValueError("Unsupported native MLX NVFP4 encoder cache") + arrays = mx.load(str(component_dir / "model.safetensors")) + matrices = manifest["matrices"] + expected = set(manifest["dense_keys"]) | set(matrices) | {key + ".scales" for key in matrices} + if set(arrays) != expected: + raise ValueError("Native MLX encoder arrays do not match the manifest") + index = cls.__new__(cls) + index.weights = {key: arrays[key] for key in manifest["dense_keys"]} + for key, info in matrices.items(): + weight, scales = arrays[key], arrays[key + ".scales"] + factor = float(info["global_scale"]) + if (weight.ndim != 2 or weight.dtype != mx.uint32 or scales.dtype != mx.uint8 + or scales.shape != (weight.shape[0], weight.shape[1] // 2) or weight.shape[1] % 2 + or not np.isfinite(factor) or factor <= 0): + raise ValueError(f"Invalid native MLX NVFP4 matrix: {key}") + index.weights[key] = NVFP4Matrix(weight, scales, factor) + mx.eval(list(arrays.values())) + return index + def get_mlx(self, key: str): return self.weights[key] @@ -458,18 +540,7 @@ class ResidentNVFP4MiniMaxH3TextConditioner(StreamedMiniMaxH3TextConditioner): """ def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = None): - raw = json.loads((Path(component_dir) / "config.json").read_text()) - quant = raw.get("quantization_config", {}) - expected = { - "quant_method": "nvfp4", - "fmt": "e2m1", - "group_size": 16, - "scale_fmt": "e4m3", - "scale_layout": "128x4", - "activation_scheme": "dynamic" - } - if any(quant.get(key) != value for key, value in expected.items()): - raise ValueError("MLX NVFP4 conditioning requires the FastVideo group-16, 128x4 encoder export.") + _read_nvfp4_encoder_config(component_dir) # Fail on an older MLX before reading the encoder's large shards. try: packed, scales = mx.quantize(mx.ones((1, 64)), mode="nvfp4") @@ -477,7 +548,11 @@ def __init__(self, component_dir: str | Path, tokenizer_dir: str | Path | None = except (ValueError, RuntimeError) as error: raise RuntimeError("Native NVFP4 conditioning requires an MLX build with nvfp4 matmul support.") from error super().__init__(component_dir, tokenizer_dir) - self.index = _ResidentNVFP4Index(self.index) + if (Path(component_dir) / MLX_NVFP4_ENCODER_MANIFEST).exists(): + self.index.close() + self.index = _ResidentNVFP4Index.from_mlx_checkpoint(component_dir) + else: + self.index = _ResidentNVFP4Index(self.index) def _embed_tokens(self, token_ids: list[int]): table = self.index.get_mlx("model.language_model.embed_tokens.weight") diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py index 6436ea8141..288b3bf465 100644 --- a/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_resident_nvfp4.py @@ -1,5 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 """Native MLX NVFP4 encoder storage and residency regression checks.""" +import json from types import SimpleNamespace import numpy as np @@ -8,7 +9,7 @@ mx = pytest.importorskip("mlx.core") from fastvideo.mlx_runtime.minimax_h3_conditioner import ( NVFP4Matrix, ResidentNVFP4MiniMaxH3TextConditioner, _ResidentNVFP4Index, - _ShardIndex, unswizzle_nvfp4_scales, + _ShardIndex, export_mlx_h3_nvfp4_encoder, unswizzle_nvfp4_scales, ) from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline @@ -44,7 +45,8 @@ def test_padded_scale_layout_round_trip(): @pytest.mark.parametrize("global_scale", [0.5, 4.0]) -def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, global_scale): +@pytest.mark.parametrize("native_cache", [False, True]) +def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, global_scale, native_cache): from safetensors.numpy import save_file rng = np.random.default_rng(3) @@ -56,6 +58,17 @@ def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, g prefix + ".weight_global_scale": np.array([global_scale], np.float32)}, tmp_path / "model.safetensors") index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + if native_cache: + _write_encoder_config(tmp_path) + cache_dir = export_mlx_h3_nvfp4_encoder(tmp_path, tmp_path / "cache") + cached = _ResidentNVFP4Index.from_mlx_checkpoint(cache_dir) + for key, original in index.weights.items(): + value = cached.get_mlx(key) + np.testing.assert_array_equal(np.array(value.weight), np.array(original.weight)) + np.testing.assert_array_equal(np.array(value.scales), np.array(original.scales)) + assert value.global_scale == original.global_scale + index.close() + index = cached weight = index.get_mlx(prefix + ".weight") assert isinstance(weight, NVFP4Matrix) assert weight.weight.dtype == mx.uint32 @@ -69,14 +82,29 @@ def test_serialized_encoder_linear_matches_independent_fp4_reference(tmp_path, g assert not index.weights -def test_resident_embedding_keeps_bf16_storage(tmp_path): +def _write_encoder_config(path): + (path / "config.json").write_text(json.dumps({"quantization_config": { + "quant_method": "nvfp4", "fmt": "e2m1", "group_size": 16, + "scale_fmt": "e4m3", "scale_layout": "128x4", "activation_scheme": "dynamic", + }})) + + +@pytest.mark.parametrize("native_cache", [False, True]) +def test_resident_embedding_keeps_bf16_storage(tmp_path, native_cache): torch = pytest.importorskip("torch") from safetensors.torch import save_file key = "model.language_model.embed_tokens.weight" table = torch.arange(60).reshape(10, 6).to(torch.bfloat16) save_file({key: table}, tmp_path / "model.safetensors") - index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) + if native_cache: + _write_encoder_config(tmp_path) + cache_dir = export_mlx_h3_nvfp4_encoder(tmp_path, tmp_path / "cache") + index = _ResidentNVFP4Index.from_mlx_checkpoint(cache_dir) + with pytest.raises(FileExistsError, match="empty"): + export_mlx_h3_nvfp4_encoder(tmp_path, cache_dir) + else: + index = _ResidentNVFP4Index(_ShardIndex(tmp_path)) assert index.get_mlx(key).dtype == mx.bfloat16 conditioner = ResidentNVFP4MiniMaxH3TextConditioner.__new__(ResidentNVFP4MiniMaxH3TextConditioner) conditioner.index = index diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index 4abcac2805..6f192889a2 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -99,6 +99,10 @@ def parse_args() -> argparse.Namespace: help=("retain and quantize transformer_blocks.*.attn.to_gate_compress.weight " "(required for MLX VSA inference; omitted by dense conversion)"), ) + parser.add_argument("--nvfp4-conditioner-root", type=Path, + help="also cache the released packed encoder in native MLX layout, without requantization") + parser.add_argument("--nvfp4-conditioner-out", type=Path, + help="empty encoder-cache output directory; defaults to OUT/nvfp4-encoder") return parser.parse_args() @@ -153,6 +157,15 @@ def main() -> None: if hasattr(mx, "clear_cache"): mx.clear_cache() + if args.nvfp4_conditioner_root is not None: + from fastvideo.mlx_runtime.minimax_h3_conditioner import export_mlx_h3_nvfp4_encoder + + encoder_out = args.nvfp4_conditioner_out or out_base / "nvfp4-encoder" + started = time.perf_counter() + export_mlx_h3_nvfp4_encoder(args.nvfp4_conditioner_root, encoder_out) + print(f"[encoder] cached packed NVFP4 encoder in {time.perf_counter() - started:.1f}s at {encoder_out}", + flush=True) + if __name__ == "__main__": main() From 10697e8b45a2b8533d4b8c49835e15d52923f200 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 18:06:44 -0700 Subject: [PATCH 087/107] [feat]: support scaled native H3 floating quantization --- fastvideo/mlx_runtime/fastwan.py | 21 ++++++++-- fastvideo/mlx_runtime/minimax_h3.py | 40 ++++++++++++++----- .../tests/test_minimax_h3_fp8_checkpoint.py | 18 ++++++--- .../convert_minimax_h3_mlx.py | 10 +++-- 4 files changed, 66 insertions(+), 23 deletions(-) diff --git a/fastvideo/mlx_runtime/fastwan.py b/fastvideo/mlx_runtime/fastwan.py index d6a802b2b1..1cf8844ebc 100644 --- a/fastvideo/mlx_runtime/fastwan.py +++ b/fastvideo/mlx_runtime/fastwan.py @@ -90,6 +90,7 @@ class QuantizedMatrix: biases: mx.array | None spec: MLXQuantizationSpec dequantized_dtype: mx.Dtype + global_scale: float = 1.0 def fastwan_shape( @@ -281,14 +282,24 @@ def ensure_quantization_supported(spec: MLXQuantizationSpec | None) -> None: f"(int8 is currently the most reliable quality/memory target).") -def quantize_matrix(weight, spec: MLXQuantizationSpec | None): +def quantize_matrix(weight, spec: MLXQuantizationSpec | None, *, use_nvfp4_global_scale: bool = False): if spec is None: return weight import mlx.core as mx if len(weight.shape) < 2: return weight - q = mx.quantize(weight, group_size=spec.group_size, bits=spec.bits, mode=spec.mode) + global_scale = 1.0 + quantization_input = weight + if spec.mode == "nvfp4" and use_nvfp4_global_scale: + # E4M3 block scales cannot represent typical small model weights + # directly. Normalize into E2M1's max 6 times E4M3's max 448. + maximum = float(mx.max(mx.abs(weight)).item()) + global_scale = maximum / (6.0 * 448.0) if maximum > 0 else 1.0 + # Normalize explicitly so this also works with older MLX operators + # without the global_scale keyword. Matmul restores this multiplier. + quantization_input = weight.astype(mx.float32) / global_scale + q = mx.quantize(quantization_input, group_size=spec.group_size, bits=spec.bits, mode=spec.mode) biases = q[2] if len(q) == 3 else None eval_args = [q[0], q[1]] if biases is not None: @@ -300,6 +311,7 @@ def quantize_matrix(weight, spec: MLXQuantizationSpec | None): biases=biases, spec=spec, dequantized_dtype=weight.dtype, + global_scale=global_scale, ) @@ -371,7 +383,7 @@ def _quantized_linear(x, weight: QuantizedMatrix, *, use_affine_dq_gemm: bool = _dq_gemm_logged = True logger.info("affine dequant+GEMM engaged (rows=%d, floor=%d, bits=%s)", rows, min_m, spec.bits) return y - return mx.quantized_matmul( + result = mx.quantized_matmul( x, weight.weight, weight.scales, @@ -381,6 +393,9 @@ def _quantized_linear(x, weight: QuantizedMatrix, *, use_affine_dq_gemm: bool = bits=spec.bits, mode=spec.mode, ).astype(x.dtype) + if weight.global_scale != 1.0: + result = (result.astype(mx.float32) * weight.global_scale).astype(x.dtype) + return result def linear(x, weight, bias=None, *, use_affine_dq_gemm: bool = False): diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 7cdf1c3dbf..37ff66b215 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -62,7 +62,7 @@ QuantizedMatrix, ensure_quantization_supported, linear as _shared_linear, - quantize_matrix, + quantize_matrix as _shared_quantize_matrix, silu, timestep_embedding, weight_dtype, @@ -87,6 +87,11 @@ def linear(x, weight, bias=None): return _shared_linear(x, weight, bias, use_affine_dq_gemm=True) +def quantize_matrix(weight, spec: MLXQuantizationSpec | None): + """Use a global NVFP4 scale for H3's small transformer weights.""" + return _shared_quantize_matrix(weight, spec, use_nvfp4_global_scale=True) + + # --------------------------------------------------------------------------- # Constants (mirrors fastvideo/pipelines/basic/minimax_h3/packing.py) # --------------------------------------------------------------------------- @@ -1493,6 +1498,7 @@ def assign(key: str, value) -> None: H3_FORMAT_VERSION = 1 +H3_SCALED_NVFP4_FORMAT_VERSION = 2 H3_WEIGHTS_FILENAME = "mlx_h3_dit.safetensors" H3_MANIFEST_FILENAME = "mlx_h3_dit.json" @@ -1562,6 +1568,8 @@ def save_mlx_h3_checkpoint(dit: MLXMiniMaxH3DiT, checkpoint_dir: str | Path) -> "dequantized_dtype": _dtype_name(value.dequantized_dtype), "has_biases": value.biases is not None, } + if value.global_scale != 1.0: + quantized[key]["global_scale"] = value.global_scale else: arrays[key] = value @@ -1580,17 +1588,25 @@ def save_mlx_h3_checkpoint(dit: MLXMiniMaxH3DiT, checkpoint_dir: str | Path) -> arrays["__adaln_cache.norm_out_scale"] = cache.norm_out_scale manifest = { - "format_version": H3_FORMAT_VERSION, - "config": dit.config, - "num_blocks": len(dit.blocks), - "num_refiner_blocks": len(dit.refiner), - "quantization": None if spec is None else { + "format_version": + (H3_SCALED_NVFP4_FORMAT_VERSION if any("global_scale" in info + for info in quantized.values()) else H3_FORMAT_VERSION), + "config": + dit.config, + "num_blocks": + len(dit.blocks), + "num_refiner_blocks": + len(dit.refiner), + "quantization": + None if spec is None else { "mode": spec.mode, "bits": spec.bits, "group_size": spec.group_size, }, - "quantized_keys": quantized, - "adaln_cache": cache_manifest, + "quantized_keys": + quantized, + "adaln_cache": + cache_manifest, "vsa": { "capable": bool(dit.vsa_capable), @@ -1629,9 +1645,9 @@ def load_mlx_h3_checkpoint(checkpoint_dir: str | Path) -> MLXMiniMaxH3DiT: manifest = json.loads(manifest_path.read_text()) version = manifest.get("format_version") - if version != H3_FORMAT_VERSION: + if version not in (H3_FORMAT_VERSION, H3_SCALED_NVFP4_FORMAT_VERSION): raise ValueError(f"MLX H3 checkpoint {checkpoint_dir} has format_version={version}; " - f"this build reads version {H3_FORMAT_VERSION}. Re-export the checkpoint.") + f"this build reads versions 1 and 2. Re-export the checkpoint.") spec = None if manifest["quantization"] is not None: @@ -1646,12 +1662,16 @@ def rebuild(key: str): return arrays[key] info = quantized_keys[key] assert spec is not None, f"Quantized key '{key}' in a checkpoint without a quantization spec" + global_scale = float(info.get("global_scale", 1.0)) + if not math.isfinite(global_scale) or global_scale <= 0: + raise ValueError(f"Invalid global scale for quantized H3 matrix {key}: {global_scale}") return QuantizedMatrix( weight=arrays[key], scales=arrays[f"{key}.scales"], biases=arrays[f"{key}.biases"] if info["has_biases"] else None, spec=spec, dequantized_dtype=_name_to_dtype(info["dequantized_dtype"]), + global_scale=global_scale, ) weights: dict[str, Any] = {} diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py index a0dc7b3216..caac1541e8 100644 --- a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py @@ -1,17 +1,22 @@ # SPDX-License-Identifier: Apache-2.0 -"""Exercise native MXFP8 H3 checkpoint storage on supported MLX builds.""" +"""Exercise native floating-point quantized H3 storage on supported MLX builds.""" import numpy as np import pytest mx = pytest.importorskip('mlx.core') -from fastvideo.mlx_runtime.fastwan import MLXQuantizationSpec, ensure_quantization_supported, linear, quantize_matrix -from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, save_mlx_h3_checkpoint +from fastvideo.mlx_runtime.fastwan import MLXQuantizationSpec, ensure_quantization_supported, linear +from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, quantize_matrix, save_mlx_h3_checkpoint -def test_mxfp8_checkpoint_preserves_quantized_matrix(tmp_path): - spec = MLXQuantizationSpec.from_name('mxfp8') +@pytest.mark.parametrize('mode', ['mxfp8', 'mxfp4', 'nvfp4']) +def test_float_quantized_checkpoint_preserves_matrix(tmp_path, mode): + spec = MLXQuantizationSpec.from_name(mode) ensure_quantization_supported(spec) - weight = quantize_matrix(mx.random.normal((64, 64)).astype(mx.bfloat16), spec) + dense = (mx.random.normal((64, 64)) * 0.001).astype(mx.bfloat16) + weight = quantize_matrix(dense, spec) + restored = mx.dequantize(weight.weight, weight.scales, mode=mode).astype(mx.float32) * weight.global_scale + relative_error = mx.sqrt(mx.sum((restored - dense.astype(mx.float32))**2) / mx.sum(dense.astype(mx.float32)**2)) + assert float(relative_error.item()) < 0.15 x = mx.random.normal((3, 64)).astype(mx.bfloat16) config = dict(hidden_size=64, num_attention_heads=1, attention_head_dim=64, ffn_dim=128, in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=64, @@ -25,3 +30,4 @@ def test_mxfp8_checkpoint_preserves_quantized_matrix(tmp_path): np.testing.assert_array_equal(np.array(actual), np.array(expected)) assert loaded.weights['test.weight'].biases is None assert loaded.weights['test.weight'].spec == spec + assert loaded.weights['test.weight'].global_scale == weight.global_scale diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index 6f192889a2..4a43adaea2 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -23,9 +23,11 @@ quantization. This saves download bytes but quantizes twice; compare its clips with the BF16-sourced export before using it for release. -``--formats mxfp8`` tries native MLX FP8 weight storage and matrix -multiplication. It is experimental and requires operator support from the -installed MLX build; the default formats remain affine INT8/INT6/INT4. +``--formats "mxfp8 mxfp4 nvfp4"`` tries native MLX floating-point quantized +storage and matrix multiplication. These formats are experimental and require +operator support from the installed MLX build. This converts BF16 or FP8 source +weights; it does not import CUDA-packed NVFP4 DiT exports. The default formats +remain affine INT8/INT6/INT4. python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ @@ -62,7 +64,7 @@ logger = init_logger(__name__) -SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8") +SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8", "mxfp4", "nvfp4") DEFAULT_FORMATS = "int8 int6 int4" From 32c8428a9c8d640589ac9a6ed344987b14b6f61a Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sun, 4 Oct 2026 18:06:53 -0700 Subject: [PATCH 088/107] [perf]: parallelize H3 SIMD softmax across four lanes --- fastvideo/mlx_runtime/minimax_h3_vsa_simd.py | 59 ++++++++++---------- 1 file changed, 30 insertions(+), 29 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index 972c1b118f..14068f14de 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -39,6 +39,7 @@ # One threadgroup = one (head, video query tile). 8 SIMD-groups x 32 = 256 # threads cover 64 query rows. K/V stage 32 keys, with 28.25 KiB total smem. # Updating online softmax once per 32 keys reduces accumulator rescaling barriers. +# Four SIMD lanes cooperate per query row in the softmax reduction. _SIMD_SOURCE = """ const int TILE = 64; const int D = 128; @@ -129,25 +130,27 @@ } simdgroup_barrier(mem_flags::mem_threadgroup); - float scores[KCHUNK]; + float scores[KCHUNK / 4]; float cmax = -3.402823466e+38f; - if (lane < (uint)ROWS) { - int grow = qrow0 + (int)lane; - for (int t = 0; t < KCHUNK; t++) { - int gtok = j0 + t; - float sc = -3.402823466e+38f; - if (grow < q_valid && gtok < k_valid && gtok < TILE) { - sc = sg_scores[(int)lane * KCHUNK + t] * scale; - } - scores[t] = sc; - cmax = metal::max(cmax, sc); + int row = (int)lane / 4; + int col_lane = (int)lane % 4; + int grow = qrow0 + row; + for (int t = col_lane; t < KCHUNK; t += 4) { + int gtok = j0 + t; + float sc = -3.402823466e+38f; + if (grow < q_valid && gtok < k_valid && gtok < TILE) { + sc = sg_scores[row * KCHUNK + t] * scale; } - float m_new = metal::max(row_m, cmax); - float alpha = metal::exp(row_m - m_new); - row_lse *= alpha; - row_m = m_new; - sg_alpha[(int)lane] = alpha; + scores[t / 4] = sc; + cmax = metal::max(cmax, sc); } + cmax = metal::max(cmax, simd_shuffle_xor(cmax, 1)); + cmax = metal::max(cmax, simd_shuffle_xor(cmax, 2)); + float m_new = metal::max(row_m, cmax); + float alpha = metal::exp(row_m - m_new); + row_lse *= alpha; + row_m = m_new; + if (col_lane == 0) sg_alpha[row] = alpha; simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { scale_rows_simd8x8(acc[kk], sg_alpha, sg_tmp, lane); @@ -167,18 +170,15 @@ threadgroup_barrier(mem_flags::mem_threadgroup); float local = 0.0f; - if (lane < (uint)ROWS) { - int grow = qrow0 + (int)lane; - for (int t = 0; t < KCHUNK; t++) { - float w = 0.0f; - if (grow < q_valid) { - w = metal::exp(scores[t] - row_m); - } - sg_scores[(int)lane * KCHUNK + t] = w; - local += w; - } - row_lse += local; + for (int t = col_lane; t < KCHUNK; t += 4) { + float w = 0.0f; + if (grow < q_valid) w = metal::exp(scores[t / 4] - row_m); + sg_scores[row * KCHUNK + t] = w; + local += w; } + local += simd_shuffle_xor(local, 1); + local += simd_shuffle_xor(local, 2); + row_lse += local; simdgroup_barrier(mem_flags::mem_threadgroup); for (int kc = 0; kc < KCHUNK; kc += 8) { @@ -194,8 +194,8 @@ } } - if (lane < (uint)ROWS) { - sg_alpha[(int)lane] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; + if (lane % 4 == 0) { + sg_alpha[(int)lane / 4] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; } simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { @@ -216,6 +216,7 @@ } simdgroup_barrier(mem_flags::mem_threadgroup); } + """ From ff8bbcc4c2536e0c34fea6a2c1d0724710dc6419 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:29:19 -0700 Subject: [PATCH 089/107] [bugfix]: publish the packed MLX encoder cache with one rename The export wrote straight into the output directory, so an out-of-memory or out-of-disk failure left partial files and the retry hit FileExistsError. It now builds in a sibling staging directory, renames it into place on success and removes it on failure. --- .../mlx_runtime/minimax_h3_conditioner.py | 22 ++++++++++++++----- 1 file changed, 17 insertions(+), 5 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index ccc1bfe7cc..3acf989e29 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -20,6 +20,8 @@ import gc import json +import shutil +import tempfile from dataclasses import dataclass from pathlib import Path from typing import Any @@ -429,11 +431,15 @@ def export_mlx_h3_nvfp4_encoder(component_dir: str | Path, output_dir: str | Pat embedding/norm values. Later loads skip CPU scale unswizzling and staging. """ raw = _read_nvfp4_encoder_config(component_dir) - output_dir = Path(output_dir) - if output_dir.exists() and any(output_dir.iterdir()): - raise FileExistsError(f"Encoder cache output must be empty: {output_dir}") - output_dir.mkdir(parents=True, exist_ok=True) + final_dir = Path(output_dir) + if final_dir.exists() and any(final_dir.iterdir()): + raise FileExistsError(f"Encoder cache output must be empty: {final_dir}") + # Build in a sibling staging directory and publish it with one rename, so an + # out-of-memory or out-of-disk failure never leaves a partial cache behind. + final_dir.parent.mkdir(parents=True, exist_ok=True) + output_dir = Path(tempfile.mkdtemp(prefix=f".{final_dir.name}.partial-", dir=final_dir.parent)) index = _ResidentNVFP4Index(_ShardIndex(Path(component_dir))) + published = False try: arrays = {} matrices = {} @@ -456,9 +462,15 @@ def export_mlx_h3_nvfp4_encoder(component_dir: str | Path, output_dir: str | Pat "source_dir": str(Path(component_dir).resolve()) } (output_dir / MLX_NVFP4_ENCODER_MANIFEST).write_text(json.dumps(manifest, indent=2) + "\n") + if final_dir.exists(): + final_dir.rmdir() # empty, checked above + output_dir.rename(final_dir) + published = True finally: index.close() - return output_dir + if not published: + shutil.rmtree(output_dir, ignore_errors=True) + return final_dir class _ResidentNVFP4Index: From e2f89dabf4fac18692238c8d318d1b5c38c836ec Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:29:19 -0700 Subject: [PATCH 090/107] [bugfix]: MLX H3 resident preload skips dropped AdaLN weights; restore the wired limit on close With the converter's AdaLN cache, dropped projection weights are None and mx.eval rejected them during resident preparation. metal_wired_limit_gib is process-wide, so close() now restores the previous limit instead of leaking it into later pipelines. --- fastvideo/mlx_runtime/minimax_h3_pipeline.py | 33 +++++++++++++++++--- 1 file changed, 28 insertions(+), 5 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index be898f3bb8..4985461589 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -245,8 +245,12 @@ def _default_metal_wired_limit_gib(mx) -> float: return min(30.0, 0.84 * total_bytes / 2**30) -def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> None: - """Keep allocator capacity separate from explicitly requested wired residency.""" +def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> int | None: + """Keep allocator capacity separate from explicitly requested wired residency. + + Returns the previous wired limit in bytes when one was set, so the caller can + restore it; the limit is process-wide and would otherwise outlive the pipeline. + """ set_memory = getattr(mx, "set_memory_limit", None) if set_memory is None and hasattr(mx, "metal"): set_memory = getattr(mx.metal, "set_memory_limit", None) @@ -256,7 +260,7 @@ def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> None: except Exception as error: # noqa: BLE001 - older MLX best effort logger.info("Could not set the Metal allocation limit: %s", error) if wired_limit_gib is None: - return + return None if not math.isfinite(wired_limit_gib) or wired_limit_gib <= 0: raise ValueError("metal_wired_limit_gib must be finite and positive") set_wired = getattr(mx, "set_wired_limit", None) @@ -267,6 +271,17 @@ def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> None: # Explicit requests must succeed; do not silently benchmark an unwired model. previous = set_wired(int(wired_limit_gib * 2**30)) logger.info("MLX wired limit %.2f GiB (previous %.2f GiB)", wired_limit_gib, previous / 2**30) + return int(previous) + + +def _restore_metal_wired_limit(mx, previous_bytes: int | None) -> None: + if previous_bytes is None: + return + set_wired = getattr(mx, "set_wired_limit", None) + if set_wired is None and hasattr(mx, "metal"): + set_wired = getattr(mx.metal, "set_wired_limit", None) + if set_wired is not None: + set_wired(previous_bytes) MINIMAX_H3_PROMPT_CACHE_VERSION = "v2-attention-layout" @@ -358,7 +373,7 @@ def __init__( ) -> None: import mlx.core as mx - _configure_metal_memory_limits(mx, metal_wired_limit_gib) + self._previous_wired_limit = _configure_metal_memory_limits(mx, metal_wired_limit_gib) self.model_root = Path(model_root) if conditioner_mode not in ("auto", "streamed", "nvfp4"): raise ValueError(f"Unknown H3 conditioner mode: {conditioner_mode}") @@ -458,7 +473,9 @@ def prepare_resident(self) -> None: self._resident_components["dit"] = dit for group in [dit.weights, *dit.blocks, *dit.refiner]: for value in group.values(): - _eval_value(value) + # With the converter's AdaLN cache, dropped AdaLN projection weights are None. + if value is not None: + _eval_value(value) cache = dit._adaln_cache if cache is not None: mx.eval(cache.block_tables, cache.norm_out_shift, cache.norm_out_scale) @@ -479,6 +496,12 @@ def close(self) -> None: conditioner.close() self._resident_components.clear() _cleanup_mlx() + previous = getattr(self, "_previous_wired_limit", None) + if previous is not None: + import mlx.core as mx + + _restore_metal_wired_limit(mx, previous) + self._previous_wired_limit = None def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: """Returns (hidden states (S, hidden), token tags). Uses the cache or From a1c024b1c30fdf37b0fa1e423ff7c978cec9cce9 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Sun, 4 Oct 2026 22:29:19 -0700 Subject: [PATCH 091/107] [test]: update MLX H3 regression stubs for model_root and NVFP4 conditioner args --- fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index f89cc584bd..5be7d672c8 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -166,7 +166,8 @@ def test_dense_generate_ignores_unused_vsa_parameters(tmp_path, monkeypatch): pipeline = object.__new__(pipeline_mod.MiniMaxH3MLXPipeline) pipeline.video_decode_backend = "h3-vae" pipeline.dit_checkpoint = tmp_path - monkeypatch.setattr(pipeline_mod, "_validate_checkpoint_step_ladder", lambda *a: None) + pipeline.model_root = tmp_path + monkeypatch.setattr(pipeline_mod, "_validate_checkpoint_step_ladder", lambda *a, **k: None) def stop_at_media_preflight(**kwargs): raise RuntimeError("reached media preflight") monkeypatch.setattr(pipeline_mod, "_preflight_media_dependencies", stop_at_media_preflight) @@ -321,7 +322,8 @@ def test_converter_continues_past_mismatched_existing_format(tmp_path, monkeypat (existing / h3.H3_MANIFEST_FILENAME).write_text(json.dumps({"vsa": {"capable": False}})) (existing / h3.H3_WEIGHTS_FILENAME).write_bytes(b"existing") monkeypatch.setattr(converter, "parse_args", lambda: argparse.Namespace( - formats="int8 int6", out=tmp_path, model_root="unused", include_vsa=True)) + formats="int8 int6", out=tmp_path, model_root="unused", include_vsa=True, + nvfp4_conditioner_root=None, nvfp4_conditioner_out=None)) monkeypatch.setattr(converter, "mlx_h3_dit_from_diffusers_safetensors", lambda *a, **k: _dit()) saved = [] monkeypatch.setattr(converter, "save_mlx_h3_checkpoint", lambda model, path: saved.append(path.name)) From 25393572dfd76949c87aba59601329cdc0229c0b Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 08:35:32 -0700 Subject: [PATCH 092/107] [test]: update MLX fast-spatial stubs for model_root and decode tile arguments --- fastvideo/tests/mlx/test_mlx_minimax_h3_fast_spatial.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_spatial.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_spatial.py index 370aa1c54d..b8a46a7f31 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_spatial.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_spatial.py @@ -93,9 +93,10 @@ def _generate_with_mocked_phases(monkeypatch, tmp_path, **generate_kwargs): pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) pipeline.video_decode_backend = "h3-vae" pipeline.dit_checkpoint = tmp_path + pipeline.model_root = tmp_path monkeypatch.setattr("fastvideo.mlx_runtime.minimax_h3_pipeline._validate_checkpoint_step_ladder", - lambda _checkpoint, _steps: None) + lambda _checkpoint, _steps, **_kwargs: None) monkeypatch.setattr("fastvideo.mlx_runtime.minimax_h3_pipeline._preflight_media_dependencies", lambda **_kwargs: None) monkeypatch.setattr("fastvideo.mlx_runtime.minimax_h3_pipeline.mlx_h3_checkpoint_vsa_capable", @@ -109,7 +110,7 @@ def fake_denoise(_text_rows, _token_tags, **kwargs): calls["denoise"] = kwargs return np.zeros((4, 4), dtype=np.float32), np.zeros((4, 4), dtype=np.float32) - def fake_decode_video(_rows, *, height, width, num_frames, tiled): + def fake_decode_video(_rows, *, height, width, num_frames, tiled, **_kwargs): events.append("decode_video") calls["decode_video"] = {"height": height, "width": width, "num_frames": num_frames} return np.zeros((num_frames, height, width, 3), dtype=np.uint8) From 6ccdbc761e6b854003f472e39b826c06cb54de60 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 08:37:39 -0700 Subject: [PATCH 093/107] [misc]: one registry entry per FastH3 switch after stacking on the RTX env migration Track C already declared FASTVIDEO_NVFP4_MM_BACKEND and FASTVIDEO_H3_VAE_TILE_BATCH (and FASTVIDEO_VSA_TRITON for Ray workers); drop the duplicate declarations from the RTX migration, override the registered FASTVIDEO_VSA_TRITON in the tile-first test, and regenerate the env-var table. --- docs/contributing/env_vars.md | 2 -- fastvideo/envs.py | 7 ------- fastvideo/tests/transformers/test_minimax_h3_tile_first.py | 2 +- 3 files changed, 1 insertion(+), 10 deletions(-) diff --git a/docs/contributing/env_vars.md b/docs/contributing/env_vars.md index 535ad581a9..ebe98a121f 100644 --- a/docs/contributing/env_vars.md +++ b/docs/contributing/env_vars.md @@ -250,7 +250,6 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_H3_FFN_CHUNK_TOKENS` | int | `0` | performance | Inference-only MiniMax-H3 FFN token chunk size; 0 runs the FFN unchunked. | | `FASTVIDEO_H3_FP8_ATTENTION` | bool | `0` | performance | With NVFP4 layer_profile h3_dit_ffn, run MiniMax-H3 attention projections in FP8. | | `FASTVIDEO_H3_FP8_GRANULARITY` | one of tensor, channel | `tensor` | performance | FP8 scaling granularity for FASTVIDEO_H3_FP8_ATTENTION. | -| `FASTVIDEO_NVFP4_MM_BACKEND` | str | `auto` | performance | FlashInfer mm_fp4 backend for NVFP4 linears, e.g. auto or cutlass. | | `FASTVIDEO_NVFP4_ACT_AMAX` | path | unset | performance | JSON of calibrated NVFP4 input amax per linear, keyed b<block>.<sub> or full prefix; sets a static activation scale. | | `FASTVIDEO_NVFP4_DYNAMIC_ACT` | str | `""` | performance | NVFP4 linears that derive the activation scale per call: all, or comma-separated layer-name suffixes such as ff.fc_out. | | `FASTVIDEO_H3_ADALN_CACHE` | bool | `0` | performance | Cache MiniMax-H3 AdaLN modulation per timestep instead of keeping the projection weights resident. | @@ -260,7 +259,6 @@ longer exists also fails the test, so the fixing pull request deletes its entry. | `FASTVIDEO_H3_SPLICE_FROM_STEP` | int | `4` | eval | First denoising step run by FASTVIDEO_H3_SPLICE_TRANSFORMER. | | `FASTVIDEO_H3_ENCODER_LAYERWISE` | bool | `0` | performance | Stream MiniMax-H3 text-encoder language layers through exact-size pinned host memory (text-only prompts). | | `FASTVIDEO_H3_ENCODER_FUSED_DEQUANT` | bool | `0` | performance | Expand the serialized NVFP4 MiniMax-H3 text encoder with one fused Triton pass on GPUs without FP4 GEMM. | -| `FASTVIDEO_H3_VAE_TILE_BATCH` | int | `1` | performance | Spatial tiles per MiniMax-H3 video VAE decoder call; 1 decodes per tile. | | `FASTVIDEO_H3_VAE_INT8_SHARED_QKV` | bool | `0` | performance | Share the INT8 activation rotation and quantization across the MiniMax-H3 VAE Q/K/V projections. | | `FASTVIDEO_H3_VAE_INT8_TRANSPOSE_VIEW` | bool | `0` | performance | Use transposed weight views in the MiniMax-H3 VAE INT8 projections. | | `FASTVIDEO_H3_VAE_INT8_FUSED_DEQUANT` | bool | `0` | performance | Fused dequantization epilogue for the MiniMax-H3 VAE INT8 projections. | diff --git a/fastvideo/envs.py b/fastvideo/envs.py index c6ea13306c..22e183a95a 100644 --- a/fastvideo/envs.py +++ b/fastvideo/envs.py @@ -557,9 +557,6 @@ def override_external(name: str, value: str | None) -> Iterator[None]: choices=("tensor", "channel"), category="performance", doc="FP8 scaling granularity for FASTVIDEO_H3_FP8_ATTENTION.") -FASTVIDEO_NVFP4_MM_BACKEND = EnvStr("auto", - category="performance", - doc="FlashInfer mm_fp4 backend for NVFP4 linears, e.g. auto or cutlass.") FASTVIDEO_NVFP4_ACT_AMAX = EnvPath(None, category="performance", doc="JSON of calibrated NVFP4 input amax per linear, keyed b. or " @@ -594,10 +591,6 @@ def override_external(name: str, value: str | None) -> Iterator[None]: category="performance", doc="Expand the serialized NVFP4 MiniMax-H3 text encoder with one " "fused Triton pass on GPUs without FP4 GEMM.") -FASTVIDEO_H3_VAE_TILE_BATCH = EnvInt(1, - category="performance", - doc="Spatial tiles per MiniMax-H3 video VAE decoder call; 1 decodes per " - "tile.") FASTVIDEO_H3_VAE_INT8_SHARED_QKV = EnvBool(False, category="performance", doc="Share the INT8 activation rotation and quantization across the " diff --git a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py index fa2ac25fcd..e6d6aa3bc4 100644 --- a/fastvideo/tests/transformers/test_minimax_h3_tile_first.py +++ b/fastvideo/tests/transformers/test_minimax_h3_tile_first.py @@ -65,7 +65,7 @@ def test_tile_first_matches_generic_vsa_with_partial_tiles(env_overrides, distri from fastvideo.models.dits.minimax_h3 import MiniMaxH3Attention env_overrides.enter_context(envs.FASTVIDEO_ATTENTION_BACKEND.override("VIDEO_SPARSE_ATTN_H3")) - env_overrides.enter_context(envs.override_external("FASTVIDEO_VSA_TRITON", "1")) + env_overrides.enter_context(envs.FASTVIDEO_VSA_TRITON.override(True)) env_overrides.enter_context(envs.FASTVIDEO_VSA_SM100A.override(False)) env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_FP4.override(False)) env_overrides.enter_context(envs.FASTVIDEO_H3_VSA_TILE_FIRST.override(False)) From e331f011a9c9487f468549f29ffbbbc97d4e9f7d Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Mon, 5 Oct 2026 13:17:09 -0700 Subject: [PATCH 094/107] [docs]: Align Spark recipes with released FastH3 stacks --- .../installation/spark_pair.md | 23 ++++ .../installation/spark_performance.md | 130 ++++++++---------- .../basic_fasth3_spark_pair_pruned_nvfp4.yaml | 15 +- .../basic_fasth3_spark_pair_v2_nvfp4.yaml | 13 +- .../basic_fasth3_spark_pruned_nvfp4.yaml | 15 +- .../basic/basic_fasth3_spark_v2_nvfp4.yaml | 13 +- .../basic/benchmark_fasth3_spark_nvfp4.py | 13 +- 7 files changed, 120 insertions(+), 102 deletions(-) diff --git a/docs/getting_started/installation/spark_pair.md b/docs/getting_started/installation/spark_pair.md index fe8aee1aa6..928bfe91b2 100644 --- a/docs/getting_started/installation/spark_pair.md +++ b/docs/getting_started/installation/spark_pair.md @@ -149,6 +149,29 @@ fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair.yam Stop the cluster when you are done: `ray stop` on both nodes. +## Released V2 and Trim NVFP4 stacks + +The public eight-forward V2 and Trim stacks include the NVFP4 encoder and +lightweight video VAE. Use the environment above plus the release kernel +settings, then run one of these configs from the head: + +```bash +export FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass +export FASTVIDEO_H3_VAE_TILE_BATCH=1 FASTVIDEO_VSA_TRITON=1 +export FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 +export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 +fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml +# Or basic_fasth3_spark_pair_pruned_nvfp4.yaml for FastH3 Trim. +``` + +Both configs use their released Hugging Face model paths, `h3_dit_vsa`, +832x480, 124 frames, seed 1234, VSA 0.8 and 64-token tiles. All components +stay resident; lazy/sequential loading and compilation are disabled for +these compact stacks. The older BF16/Preview memory guidance below applies +to those larger stacks. Set 1344x768 for native 768p with the same 124 frames. +See [the resident release recipe](spark_performance.md#fasth3-v2-nvfp4-on-one-spark) +for checkpoint contents and the timing protocol. + ## FastH3 frame counts H3 is 24 fps. Legal `num_frames` values are `17n+5`. The pipeline rejects diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index 581f10f9d8..b36bce5f18 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -217,48 +217,30 @@ three resident on one GB10. This stack fits in the Spark's unified memory; benchmark your installed runtime and review the clips before publishing a speed claim. The earlier bf16 H3 memory guidance above concerns a larger checkpoint. -Install FastVideo from a checkout that includes the ModelOpt converter and -FlashInfer FP4 support, following [the Spark install guide](spark.md). Sign in -to Hugging Face with access to the FastVideo model repositories. Download the -V2 scheduler and audio components, the compact encoder and VAE from the pruned -repo, and the ModelOpt V2 transformer. The pruned model's encoder and VAE are -the same components used by V2. +Install FastVideo following [the Spark install guide](spark.md). The released +repositories are complete inference stacks: + +| Model | Repository | Packed DiT profile | +|---|---|---| +| FastH3 V2 | `FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4-Consumer` | `h3_dit_vsa` | +| FastH3 Trim | `FastVideo/FastVideo-FastH3-Trim-8-Step-NVFP4` | `h3_dit_vsa` | + +Each ships its own trained schedule, NVFP4 transformer and 50-layer NVFP4 +text encoder, lightweight 26-layer video VAE with the INT8-weight overlay, +and audio VAE. Download the complete repository; the runtime selects these +components from its model index. The released Trim transformer also packs +attention and VSA gates, unlike the earlier FFN-only pruned export. Neither +release requires local checkpoint conversion or a separate encoder download. ```bash -SPARK_STACK=./FastH3-V2-Spark-NVFP4 -V2_FP4_SRC=./FastH3-V2-ModelOpt-NVFP4 - -hf download FastVideo/FastVideo-FastH3-8-Step-V2 \ - --local-dir "$SPARK_STACK" \ - --exclude 'transformer/*' --exclude 'text_encoder/*' --exclude 'vae/*' -hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ - --local-dir "$SPARK_STACK" \ - --include 'text_encoder/*' --include 'vae/*' -hf download FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4 \ - --local-dir "$V2_FP4_SRC" --include 'transformer/*' - -nice -n 19 python scripts/checkpoint_conversion/convert_minimax_h3_modelopt_nvfp4_dit.py \ - --src "$V2_FP4_SRC/transformer" --dst "$SPARK_STACK/transformer" \ - --quantize-attention --quantize-gate - -test -f "$SPARK_STACK/transformer/nvfp4_weights.safetensors" -test -f "$SPARK_STACK/text_encoder/config.json" -test -f "$SPARK_STACK/vae/config.json" -test -f "$SPARK_STACK/fastvideo_inference.json" -python -m json.tool "$SPARK_STACK/fastvideo_inference.json" >/dev/null +hf download FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4-Consumer +hf download FastVideo/FastVideo-FastH3-Trim-8-Step-NVFP4 ``` -The converter probes each packed linear through FlashInfer `mm_fp4`. If that -probe fails on `sm_121`, convert the transformer on another Blackwell GPU and -copy the resulting `transformer/` directory to the Spark. Do not omit -`fastvideo_inference.json`: it supplies V2's trained denoising ladder. The -recipe's `num_inference_steps: 9` means nine sigma points and eight DiT -forwards. - -The converter preserves ModelOpt's calibrated `input_scale` as the reciprocal -`_nvfp4_input_global_sf`. Reconvert older exports that discarded this scale: -unit activation scaling clips inputs above 2688. An explicit `--act-amax` -table overrides the source calibration. +Keep `fastvideo_inference.json` with the transformer if you stage the stack +in a local directory. It declares the trained eight-forward ladder and +video/audio shifts of 10/3. The recipes use `num_inference_steps: 9` for +nine sigma points and eight DiT forwards. On GB10 with FlashInfer 0.6.18, FastVideo fences activation quantization before releasing its padded input. Without this completion fence, identical H3 @@ -266,7 +248,7 @@ requests produced different DiT latents and occasionally corrupt video. The fence applies to `sm_121`; other architectures retain asynchronous execution. Run `examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml` from the -repository root. It uses 832x480, 243 frames, VSA sparsity 0.8 with +repository root. It uses 832x480, 124 frames and seed 1234, VSA sparsity 0.8 with 64-token tiles, and the light H3 VAE through the `h3-vae` decode backend. It does not use frame dropping or spatial upscaling. @@ -281,40 +263,47 @@ nice -n 19 fastvideo generate \ --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml ``` -For a roughly five-second clip, set `--request.sampling.num_frames 124` and -write to a separate output path. H3 permits frame counts of `17n+5`; 124 is -the closest legal count above five seconds at 24 fps. For the secondary -10-second setting, set `--request.sampling.width 1344` and -`--request.sampling.height 768`, keeping 243 frames. Use the two prompts in -`handoff_spark_mac/benchmark_prompts.json` from the local release handoff. -The benchmark script runs one warmup and at least two timed generations for -each prompt in one process. It saves the MP4s and prints the wall time, stage -times, peak memory, and median. Set the Spark environment before running it: +The defaults produce a roughly five-second clip at 24 fps. For native 768p, +set `--request.sampling.width 1344` and `--request.sampling.height 768`, +keeping 124 frames. For the separate ten-second setting, use 243 frames. +H3 permits frame counts of `17n+5`; 124 is the nearest legal count above +five seconds and 243 is the nearest above ten seconds. + +For Trim, use `basic_fasth3_spark_pruned_nvfp4.yaml` with the same environment. +Both recipes keep the encoder, DiT and VAEs resident, disable compilation, +and use `h3_dit_vsa`. On two Sparks, use the corresponding +`basic_fasth3_spark_pair_{pruned,v2}_nvfp4.yaml` after following the +[pair setup guide](spark_pair.md). The pair uses SP2/TP1 and parallel VAE +gathering, with tile batch 1 on each worker. + +For release timing, create one generator per model/resolution. Run one +untimed ceramics warmup, then two timed ceramics calls and two timed harbor +calls in that same process, using the exact release prompt strings, seed +1234 and the settings above. Measure each `generate()` call through finished +MP4 output and report the median of the two timed calls per prompt. Keep the +warmup excluded. Record the model revision, code commit, command, environment +and peak-memory scope with the results. Review every clip's video and audio +before publishing a quality or speed claim. + +The benchmark helper requires a local stack so it can validate the schedule +before loading. For example: ```bash -export FASTVIDEO_MINIMAX_H3_FUSIONS=all -export FASTVIDEO_NVFP4_MM_BACKEND=cutlass FASTVIDEO_H3_VAE_TILE_BATCH=1 -export FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 -export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 -nice -n 19 python examples/inference/basic/benchmark_fasth3_spark_nvfp4.py \ +hf download FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4-Consumer \ + --local-dir ./FastH3-V2-Consumer +python examples/inference/basic/benchmark_fasth3_spark_nvfp4.py \ --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml \ - --prompts /path/to/fasth3-local-release/handoff_spark_mac/benchmark_prompts.json \ - --output-dir outputs/fasth3_spark_v2_nvfp4/benchmark-243 --frames 243 - -# Repeat with --frames 124 and a different output directory for the five-second check. + --model-path ./FastH3-V2-Consumer --frames 124 \ + --prompts /path/to/benchmark_prompts.json \ + --output-dir outputs/fasth3_spark_v2_nvfp4/benchmark-124 ``` -Record the exact command and commit with the measurements. Review every clip's -video and audio before publishing a quality or speed claim. - -After the V2 baseline works, sweep `FASTVIDEO_H3_VAE_TILE_BATCH` and -`FASTVIDEO_NVFP4_MM_BACKEND` on the same prompts. Compare the optional AdaLN -table and VAE compile only with the same frame count, schedule, and VSA -sparsity. The V2 converter packs VSA gates, so its `h3_dit_vsa` profile must -match the recipe. A later pruned NVFP4 transformer uses the separate -`h3_dit_ffn` profile, with attention and VSA gates left dense. +Set the environment from the generation command above before benchmarking. +The prompt JSON must contain `latency-ceramics-005` and +`latency-harbor-005`. Pass `--width 1344 --height 768` for the native 768p +protocol. Use the Trim repository and config for its corresponding run. -### Native FastH3 release measurements +### Historical eight-forward controls Measured on October 4, 2026, with the trained eight-forward ladder, VSA 0.8, 832x480 native video, the 50-layer NVFP4 encoder and light video/audio VAEs @@ -335,9 +324,10 @@ The final pair uses `715d4a5f`, with matching actual-worker code fingerprints, CUTLASS FP4 GEMMs and Triton VSA. Light-VAE tile batch is 8 on one Spark and 1 on the pair; pair batch 8 was slower at 124 frames (79.993 / 80.022 s pruned). All offload/deferred-loading and compile options are disabled. The pair uses -QSFP RoCE, SP2/TP1 and parallel VAE gathering. See the pair configs -`basic_fasth3_spark_pair_pruned_nvfp4.yaml` and -`basic_fasth3_spark_pair_v2_nvfp4.yaml` beside the benchmark script. +QSFP RoCE, SP2/TP1 and parallel VAE gathering. These are historical controls, not measurements of the new recipe defaults. +The pruned rows used the earlier FFN-only export, and all rows used seed +2026 with the indicated frame counts. Public Trim packs attention and gates +too; its recipe now defaults to seed 1234 and 124 frames. Every final pair warmup and repeat has correct dimensions/frame count, coherent sampled frames and identical full decoded-video hashes within its prompt. diff --git a/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml index 32852ad514..5d59fcfb12 100644 --- a/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml +++ b/examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml @@ -1,24 +1,25 @@ # Start Ray on both Sparks and source spark_pair_env.sh first; see spark_pair.md. -# FastH3 pruned ckpt300 eight-forward video+audio on two DGX Sparks over QSFP RoCE. +# FastH3 Trim eight-forward video+audio on two DGX Sparks over QSFP RoCE. # Download the complete checkpoint, including its trained schedule, as described in # docs/getting_started/installation/spark_performance.md. # -# GB10 uses Triton VSA. The packed transformer uses NVFP4 FFN weights with bf16 attention -# and VSA gates (layer_profile: h3_dit_ffn). No temporal or spatial fast mode. +# GB10 uses Triton VSA. The release packs attention, FFN and VSA gate weights. +# The NVFP4 encoder and lightweight VAE ship in the same repository. # # FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_H3_VAE_TILE_BATCH=1 \ # FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ # fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_pruned_nvfp4.yaml generator: - model_path: FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 + model_path: FastVideo/FastVideo-FastH3-Trim-8-Step-NVFP4 engine: num_gpus: 2 execution_backend: ray use_fsdp_inference: false quantization: transformer_quant: NVFP4 - layer_profile: h3_dit_ffn + layer_profile: h3_dit_vsa parallelism: tp_size: 1 sp_size: 2 @@ -49,10 +50,10 @@ request: prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. negative_prompt: "" sampling: - seed: 2026 + seed: 1234 height: 480 width: 832 - num_frames: 243 + num_frames: 124 fps: 24 num_inference_steps: 9 # nine sigma points, eight DiT forwards guidance_scale: 1.0 diff --git a/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml index 89bef8ed8a..4d1a0d55e3 100644 --- a/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml +++ b/examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml @@ -1,17 +1,18 @@ # Start Ray on both Sparks and source spark_pair_env.sh first; see spark_pair.md. # FastH3 V2 eight-forward video+audio on two DGX Sparks over QSFP RoCE. -# Assemble ./FastH3-V2-Spark-NVFP4 as described in +# Download the complete release checkpoint as described in # docs/getting_started/installation/spark_performance.md. # -# GB10 uses Triton VSA. The packed transformer must include NVFP4 attention -# and VSA gates (layer_profile: h3_dit_vsa). No temporal or spatial fast mode. +# GB10 uses Triton VSA. The release packs attention, FFN and VSA gate weights. +# The NVFP4 encoder and lightweight VAE ship in the same repository. # # FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_H3_VAE_TILE_BATCH=1 \ # FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ # fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair_v2_nvfp4.yaml generator: - model_path: ./FastH3-V2-Spark-NVFP4 + model_path: FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4-Consumer engine: num_gpus: 2 execution_backend: ray @@ -49,10 +50,10 @@ request: prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. negative_prompt: "" sampling: - seed: 2026 + seed: 1234 height: 480 width: 832 - num_frames: 243 + num_frames: 124 fps: 24 num_inference_steps: 9 # nine sigma points, eight DiT forwards guidance_scale: 1.0 diff --git a/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml index a5d1af4c47..5ad4495124 100644 --- a/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml +++ b/examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml @@ -1,22 +1,23 @@ -# FastH3 pruned ckpt300 eight-forward video+audio on one DGX Spark. +# FastH3 Trim eight-forward video+audio on one DGX Spark. # Download the complete checkpoint, including its trained schedule, as described in # docs/getting_started/installation/spark_performance.md. # -# GB10 uses Triton VSA. The packed transformer uses NVFP4 FFN weights with bf16 attention -# and VSA gates (layer_profile: h3_dit_ffn). No temporal or spatial fast mode. +# GB10 uses Triton VSA. The release packs attention, FFN and VSA gate weights. +# The NVFP4 encoder and lightweight VAE ship in the same repository. # # FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_H3_VAE_TILE_BATCH=1 \ # FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ # fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pruned_nvfp4.yaml generator: - model_path: FastVideo/FastH3-Pruned-8Step-NVFP4-ckpt300 + model_path: FastVideo/FastVideo-FastH3-Trim-8-Step-NVFP4 engine: num_gpus: 1 use_fsdp_inference: false quantization: transformer_quant: NVFP4 - layer_profile: h3_dit_ffn + layer_profile: h3_dit_vsa parallelism: tp_size: 1 sp_size: 1 @@ -46,10 +47,10 @@ request: prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. negative_prompt: "" sampling: - seed: 2026 + seed: 1234 height: 480 width: 832 - num_frames: 243 + num_frames: 124 fps: 24 num_inference_steps: 9 # nine sigma points, eight DiT forwards guidance_scale: 1.0 diff --git a/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml b/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml index c3986fabd0..27531db0e7 100644 --- a/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml +++ b/examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml @@ -1,16 +1,17 @@ # FastH3 V2 eight-forward video+audio on one DGX Spark. -# Assemble ./FastH3-V2-Spark-NVFP4 as described in +# Download the complete release checkpoint as described in # docs/getting_started/installation/spark_performance.md. # -# GB10 uses Triton VSA. The packed transformer must include NVFP4 attention -# and VSA gates (layer_profile: h3_dit_vsa). No temporal or spatial fast mode. +# GB10 uses Triton VSA. The release packs attention, FFN and VSA gate weights. +# The NVFP4 encoder and lightweight VAE ship in the same repository. # # FASTVIDEO_MINIMAX_H3_FUSIONS=all FASTVIDEO_NVFP4_MM_BACKEND=cutlass \ +# FASTVIDEO_H3_VAE_TILE_BATCH=1 \ # FASTVIDEO_VSA_TRITON=1 FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ # FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_STAGE_LOGGING=1 \ # fastvideo generate --config examples/inference/basic/basic_fasth3_spark_v2_nvfp4.yaml generator: - model_path: ./FastH3-V2-Spark-NVFP4 + model_path: FastVideo/FastVideo-FastH3-8-Step-V2-NVFP4-Consumer engine: num_gpus: 1 use_fsdp_inference: false @@ -46,10 +47,10 @@ request: prompt: A quiet pottery studio with a potter finishing a bowl at the wheel. negative_prompt: "" sampling: - seed: 2026 + seed: 1234 height: 480 width: 832 - num_frames: 243 + num_frames: 124 fps: 24 num_inference_steps: 9 # nine sigma points, eight DiT forwards guidance_scale: 1.0 diff --git a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py index f8e3ee7fcd..20d31043a7 100644 --- a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -1,8 +1,8 @@ # SPDX-License-Identifier: Apache-2.0 """Time a resident FastH3 eight-forward Spark recipe with the release prompts. -One process loads the model, then each prompt gets one excluded warmup and at -least two timed calls. Every call writes a video. This script does not alter +One process loads the model, runs one excluded ceramics warmup, then times +each prompt at least twice. Every call writes a video. This script does not alter the V2 schedule, VSA sparsity, or video resolution. """ @@ -67,14 +67,14 @@ def main() -> None: parser.add_argument("--prompts", type=Path, required=True) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--model-path", type=Path) - parser.add_argument("--frames", type=int, default=243) + parser.add_argument("--frames", type=int, default=124) parser.add_argument("--width", type=int, default=832) parser.add_argument("--height", type=int, default=480) parser.add_argument("--repeats", type=int, default=2) args = parser.parse_args() if args.frames not in (124, 243): - parser.error("use 124 frames for roughly five seconds or 243 for the ten-second headline") + parser.error("use 124 frames for roughly five seconds or 243 for roughly ten seconds") if args.repeats < 2: parser.error("the release protocol requires at least two timed calls") @@ -99,9 +99,10 @@ def main() -> None: generator = VideoGenerator.from_config(config.generator) try: - for prompt_id in PROMPT_IDS: + for prompt_index, prompt_id in enumerate(PROMPT_IDS): times = [] - for index in range(args.repeats + 1): + first_index = 0 if prompt_index == 0 else 1 + for index in range(first_index, args.repeats + 1): warmup = index == 0 label = "warmup" if warmup else f"run-{index:02d}" requested_path = args.output_dir / f"{prompt_id}-{args.width}x{args.height}-{args.frames}-{label}.mp4" From 31687469e0e31830f867fae6f7a74c7bbd63671a Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Mon, 5 Oct 2026 13:27:34 -0700 Subject: [PATCH 095/107] [docs]: Keep Spark launch table on released models --- .../installation/spark_performance.md | 56 +++++++------------ 1 file changed, 20 insertions(+), 36 deletions(-) diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index b36bce5f18..c16dc9e9a7 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -303,42 +303,26 @@ The prompt JSON must contain `latency-ceramics-005` and `latency-harbor-005`. Pass `--width 1344 --height 768` for the native 768p protocol. Use the Trim repository and config for its corresponding run. -### Historical eight-forward controls - -Measured on October 4, 2026, with the trained eight-forward ladder, VSA 0.8, -832x480 native video, the 50-layer NVFP4 encoder and light video/audio VAEs -resident. Each cell is the median of two timed calls after one warmup, in -seconds, for `latency-ceramics-005` / `latency-harbor-005`, seed 2026. - -| Model | Frames | One Spark | Two Sparks, SP2 | -|---|---:|---:|---:| -| Pruned ckpt300, FFN NVFP4 | 124 | 134.361 / 134.675 | **78.245 / 78.272** | -| Pruned ckpt300, FFN NVFP4 | 243 | 284.188 / 280.593 | **165.987 / 163.001** | -| Full V2, NVFP4 FFN/attention/gates | 124 | 142.460 / 140.425 | **88.356 / 86.101** | -| Full V2, NVFP4 FFN/attention/gates | 243 | 308.204 / 306.141 | **179.876 / 180.213** | - -The base integrates upstream main `0cc41a22` with experimental NVFP4 support. -One-Spark tested commits are `6d6b57fe` (pruned 124), `3f24557a` (pruned 243, -with `CUDA_LAUNCH_BLOCKING=1`), `f5126f78` (V2 124) and `6e9d7a0b` (V2 243). -The final pair uses `715d4a5f`, with matching actual-worker code fingerprints, -CUTLASS FP4 GEMMs and Triton VSA. Light-VAE tile batch is 8 on one Spark and 1 -on the pair; pair batch 8 was slower at 124 frames (79.993 / 80.022 s pruned). -All offload/deferred-loading and compile options are disabled. The pair uses -QSFP RoCE, SP2/TP1 and parallel VAE gathering. These are historical controls, not measurements of the new recipe defaults. -The pruned rows used the earlier FFN-only export, and all rows used seed -2026 with the indicated frame counts. Public Trim packs attention and gates -too; its recipe now defaults to seed 1234 and 124 frames. - -Every final pair warmup and repeat has correct dimensions/frame count, coherent -sampled frames and identical full decoded-video hashes within its prompt. -The V2 one-Spark 124-frame harbor warmup differs from the timed clips but remains -coherent. These checks establish repeat reliability for the tested recipes; -BF16 reference parity, speech accuracy and lip sync need separate review. - -The [older H3 local blog](https://haoailab.com/blogs/fasth3-local/) reports -243 s on one Spark and 209 s on two Sparks at 124 frames. It uses the four-step -Preview checkpoint and full VAE, so the old and new values are context, not a -matched optimization comparison. It has no matching 243-frame baseline. +### Released model measurements + +The released Trim stack at revision `cae9ceb6feefe77d34a56640782cda3909363f19` +completed the native 832x480, 124-frame protocol on one Spark, seed 1234. +The tested code is `6ccdbc761e6b854003f472e39b826c06cb54de60`, using the +resident recipe above. Medians exclude warmups and cover finished MP4 output. + +| Released model | Resolution | One Spark, ceramics / harbor | Two Sparks | +|---|---|---:|---| +| FastH3 Trim NVFP4 | 832x480 | 124.090 / 127.519 s | Pending | +| FastH3 Trim NVFP4 | 1344x768 | Review pending | Pending | +| FastH3 V2 NVFP4-Consumer | 832x480 | Repeatability review pending | Pending | +| FastH3 V2 NVFP4-Consumer | 1344x768 | Review pending | Pending | + +The completed Trim 480p batch used one extra untimed harbor warmup in the +same process, six calls total. Both warmups are excluded from the medians. +Each prompt's three decoded videos and audio streams match exactly, and +sampled frames are coherent. These checks do not establish BF16 parity, +speech accuracy or lip sync. Pending cells are not measured substitutes +from older checkpoints. ## Reproduce these numbers From ecbdbe59cab835d7747add8b0e0dd6e2d9c3c029 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Mon, 5 Oct 2026 13:33:14 -0700 Subject: [PATCH 096/107] [misc]: ship validated Track C MLX release paths --- fastvideo/mlx_runtime/fastwan.py | 21 +--- fastvideo/mlx_runtime/minimax_h3.py | 54 +++------- fastvideo/mlx_runtime/minimax_h3_vsa_simd.py | 98 +++++++++---------- .../tests/test_minimax_h3_fp8_checkpoint.py | 33 ------- .../test_mlx_minimax_h3_vsa_regressions.py | 20 ---- .../convert_minimax_h3_mlx.py | 15 +-- 6 files changed, 62 insertions(+), 179 deletions(-) delete mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py diff --git a/fastvideo/mlx_runtime/fastwan.py b/fastvideo/mlx_runtime/fastwan.py index 1cf8844ebc..d6a802b2b1 100644 --- a/fastvideo/mlx_runtime/fastwan.py +++ b/fastvideo/mlx_runtime/fastwan.py @@ -90,7 +90,6 @@ class QuantizedMatrix: biases: mx.array | None spec: MLXQuantizationSpec dequantized_dtype: mx.Dtype - global_scale: float = 1.0 def fastwan_shape( @@ -282,24 +281,14 @@ def ensure_quantization_supported(spec: MLXQuantizationSpec | None) -> None: f"(int8 is currently the most reliable quality/memory target).") -def quantize_matrix(weight, spec: MLXQuantizationSpec | None, *, use_nvfp4_global_scale: bool = False): +def quantize_matrix(weight, spec: MLXQuantizationSpec | None): if spec is None: return weight import mlx.core as mx if len(weight.shape) < 2: return weight - global_scale = 1.0 - quantization_input = weight - if spec.mode == "nvfp4" and use_nvfp4_global_scale: - # E4M3 block scales cannot represent typical small model weights - # directly. Normalize into E2M1's max 6 times E4M3's max 448. - maximum = float(mx.max(mx.abs(weight)).item()) - global_scale = maximum / (6.0 * 448.0) if maximum > 0 else 1.0 - # Normalize explicitly so this also works with older MLX operators - # without the global_scale keyword. Matmul restores this multiplier. - quantization_input = weight.astype(mx.float32) / global_scale - q = mx.quantize(quantization_input, group_size=spec.group_size, bits=spec.bits, mode=spec.mode) + q = mx.quantize(weight, group_size=spec.group_size, bits=spec.bits, mode=spec.mode) biases = q[2] if len(q) == 3 else None eval_args = [q[0], q[1]] if biases is not None: @@ -311,7 +300,6 @@ def quantize_matrix(weight, spec: MLXQuantizationSpec | None, *, use_nvfp4_globa biases=biases, spec=spec, dequantized_dtype=weight.dtype, - global_scale=global_scale, ) @@ -383,7 +371,7 @@ def _quantized_linear(x, weight: QuantizedMatrix, *, use_affine_dq_gemm: bool = _dq_gemm_logged = True logger.info("affine dequant+GEMM engaged (rows=%d, floor=%d, bits=%s)", rows, min_m, spec.bits) return y - result = mx.quantized_matmul( + return mx.quantized_matmul( x, weight.weight, weight.scales, @@ -393,9 +381,6 @@ def _quantized_linear(x, weight: QuantizedMatrix, *, use_affine_dq_gemm: bool = bits=spec.bits, mode=spec.mode, ).astype(x.dtype) - if weight.global_scale != 1.0: - result = (result.astype(mx.float32) * weight.global_scale).astype(x.dtype) - return result def linear(x, weight, bias=None, *, use_affine_dq_gemm: bool = False): diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 37ff66b215..5b5e546cc6 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -62,7 +62,7 @@ QuantizedMatrix, ensure_quantization_supported, linear as _shared_linear, - quantize_matrix as _shared_quantize_matrix, + quantize_matrix, silu, timestep_embedding, weight_dtype, @@ -87,11 +87,6 @@ def linear(x, weight, bias=None): return _shared_linear(x, weight, bias, use_affine_dq_gemm=True) -def quantize_matrix(weight, spec: MLXQuantizationSpec | None): - """Use a global NVFP4 scale for H3's small transformer weights.""" - return _shared_quantize_matrix(weight, spec, use_nvfp4_global_scale=True) - - # --------------------------------------------------------------------------- # Constants (mirrors fastvideo/pipelines/basic/minimax_h3/packing.py) # --------------------------------------------------------------------------- @@ -1395,8 +1390,6 @@ def assign(key: str, value) -> None: for shard in _safetensors_shards(transformer_path): shard_arrays = mx.load(str(shard)) for key, source in shard_arrays.items(): - if key.endswith(".weight_scale"): - continue # paired with its FP8 weight if _is_ignored_dense_key(key, include_vsa=include_vsa): continue if temb is not None and (key.startswith("time_embedder.") or key == "adaln_basis.weight"): @@ -1411,17 +1404,7 @@ def assign(key: str, value) -> None: factorized_adaln = config.get("adaln_rank") is not None and (".adaln_proj." in key or key.startswith( ("norm_out.linear.", "adaln_basis."))) target_dtype = mx.float32 if keep_fp32 else (mx.float16 if factorized_adaln else cast_dtype) - if source.dtype == mx.uint8 and key.endswith(".weight"): - scale_key = key + "_scale" - if scale_key not in shard_arrays: - raise KeyError(f"FP8 weight {key} needs {scale_key} in the same safetensors shard") - scale = shard_arrays[scale_key].astype(mx.float32) - if scale.size != source.shape[0]: - raise ValueError(f"FP8 scale for {key} has {scale.size} entries, expected {source.shape[0]}") - array = (mx.from_fp8(source, dtype=mx.float16) * scale.reshape(-1, 1)).astype(target_dtype) - mx.eval(array) - else: - array = _load_array(source, target_dtype) + array = _load_array(source, target_dtype) if temb is not None and ".adaln_proj.linear." in key: _, index_str, sub = key.split(".", 2) index = int(index_str) @@ -1498,7 +1481,6 @@ def assign(key: str, value) -> None: H3_FORMAT_VERSION = 1 -H3_SCALED_NVFP4_FORMAT_VERSION = 2 H3_WEIGHTS_FILENAME = "mlx_h3_dit.safetensors" H3_MANIFEST_FILENAME = "mlx_h3_dit.json" @@ -1568,8 +1550,6 @@ def save_mlx_h3_checkpoint(dit: MLXMiniMaxH3DiT, checkpoint_dir: str | Path) -> "dequantized_dtype": _dtype_name(value.dequantized_dtype), "has_biases": value.biases is not None, } - if value.global_scale != 1.0: - quantized[key]["global_scale"] = value.global_scale else: arrays[key] = value @@ -1588,25 +1568,17 @@ def save_mlx_h3_checkpoint(dit: MLXMiniMaxH3DiT, checkpoint_dir: str | Path) -> arrays["__adaln_cache.norm_out_scale"] = cache.norm_out_scale manifest = { - "format_version": - (H3_SCALED_NVFP4_FORMAT_VERSION if any("global_scale" in info - for info in quantized.values()) else H3_FORMAT_VERSION), - "config": - dit.config, - "num_blocks": - len(dit.blocks), - "num_refiner_blocks": - len(dit.refiner), - "quantization": - None if spec is None else { + "format_version": H3_FORMAT_VERSION, + "config": dit.config, + "num_blocks": len(dit.blocks), + "num_refiner_blocks": len(dit.refiner), + "quantization": None if spec is None else { "mode": spec.mode, "bits": spec.bits, "group_size": spec.group_size, }, - "quantized_keys": - quantized, - "adaln_cache": - cache_manifest, + "quantized_keys": quantized, + "adaln_cache": cache_manifest, "vsa": { "capable": bool(dit.vsa_capable), @@ -1645,9 +1617,9 @@ def load_mlx_h3_checkpoint(checkpoint_dir: str | Path) -> MLXMiniMaxH3DiT: manifest = json.loads(manifest_path.read_text()) version = manifest.get("format_version") - if version not in (H3_FORMAT_VERSION, H3_SCALED_NVFP4_FORMAT_VERSION): + if version != H3_FORMAT_VERSION: raise ValueError(f"MLX H3 checkpoint {checkpoint_dir} has format_version={version}; " - f"this build reads versions 1 and 2. Re-export the checkpoint.") + f"this build reads version {H3_FORMAT_VERSION}. Re-export the checkpoint.") spec = None if manifest["quantization"] is not None: @@ -1662,16 +1634,12 @@ def rebuild(key: str): return arrays[key] info = quantized_keys[key] assert spec is not None, f"Quantized key '{key}' in a checkpoint without a quantization spec" - global_scale = float(info.get("global_scale", 1.0)) - if not math.isfinite(global_scale) or global_scale <= 0: - raise ValueError(f"Invalid global scale for quantized H3 matrix {key}: {global_scale}") return QuantizedMatrix( weight=arrays[key], scales=arrays[f"{key}.scales"], biases=arrays[f"{key}.biases"] if info["has_biases"] else None, spec=spec, dequantized_dtype=_name_to_dtype(info["dequantized_dtype"]), - global_scale=global_scale, ) weights: dict[str, Any] = {} diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index 14068f14de..bb21b1912a 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -37,18 +37,16 @@ """ # One threadgroup = one (head, video query tile). 8 SIMD-groups x 32 = 256 -# threads cover 64 query rows. K/V stage 32 keys, with 28.25 KiB total smem. -# Updating online softmax once per 32 keys reduces accumulator rescaling barriers. -# Four SIMD lanes cooperate per query row in the softmax reduction. +# threads cover 64 query rows. Q is half smem (16 KiB). K/V stage 8 keys. _SIMD_SOURCE = """ const int TILE = 64; const int D = 128; const int SG = 32; const int N_SG = 8; const int ROWS = 8; - const int KCHUNK = 32; - threadgroup float kvsmem[KCHUNK * D]; - threadgroup float score_smem[N_SG * ROWS * KCHUNK]; + const int KCHUNK = 8; + threadgroup float kvsmem[8 * 128]; + threadgroup float score_smem[8 * 64]; threadgroup float scale_tmp[8 * 64]; threadgroup float qtile[8 * 64]; threadgroup float row_alpha[8 * 8]; @@ -67,7 +65,7 @@ int qt = n_prefix + (int)q_tile; int q_valid = active ? vbs[qt] : 0; int q_base_tile = (((int)head * S) + qt * TILE) * D; - threadgroup float *sg_scores = score_smem + sid * ROWS * KCHUNK; + threadgroup float *sg_scores = score_smem + sid * 64; threadgroup float *sg_tmp = scale_tmp + sid * 64; threadgroup float *sg_qtile = qtile + sid * 64; threadgroup float *sg_alpha = row_alpha + sid * 8; @@ -119,38 +117,34 @@ } threadgroup_barrier(mem_flags::mem_threadgroup); - for (int kc = 0; kc < KCHUNK; kc += 8) { - thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 kmat; - simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D, ulong2(0, 0), true); - simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); - } - simdgroup_store(smat, sg_scores + kc, KCHUNK); + thread simdgroup_float8x8 smat = make_filled_simdgroup_matrix(0.0f); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 kmat; + simdgroup_load(kmat, (const threadgroup float*)(kvsmem + kk * 8), D, ulong2(0, 0), true); + simdgroup_multiply_accumulate(smat, qfrag[kk], kmat, smat); } + simdgroup_store(smat, sg_scores, KCHUNK); simdgroup_barrier(mem_flags::mem_threadgroup); - float scores[KCHUNK / 4]; + float scores[8]; float cmax = -3.402823466e+38f; - int row = (int)lane / 4; - int col_lane = (int)lane % 4; - int grow = qrow0 + row; - for (int t = col_lane; t < KCHUNK; t += 4) { - int gtok = j0 + t; - float sc = -3.402823466e+38f; - if (grow < q_valid && gtok < k_valid && gtok < TILE) { - sc = sg_scores[row * KCHUNK + t] * scale; + if (lane < (uint)ROWS) { + int grow = qrow0 + (int)lane; + for (int t = 0; t < KCHUNK; t++) { + int gtok = j0 + t; + float sc = -3.402823466e+38f; + if (grow < q_valid && gtok < k_valid && gtok < TILE) { + sc = sg_scores[(int)lane * KCHUNK + t] * scale; + } + scores[t] = sc; + cmax = metal::max(cmax, sc); } - scores[t / 4] = sc; - cmax = metal::max(cmax, sc); + float m_new = metal::max(row_m, cmax); + float alpha = metal::exp(row_m - m_new); + row_lse *= alpha; + row_m = m_new; + sg_alpha[(int)lane] = alpha; } - cmax = metal::max(cmax, simd_shuffle_xor(cmax, 1)); - cmax = metal::max(cmax, simd_shuffle_xor(cmax, 2)); - float m_new = metal::max(row_m, cmax); - float alpha = metal::exp(row_m - m_new); - row_lse *= alpha; - row_m = m_new; - if (col_lane == 0) sg_alpha[row] = alpha; simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { scale_rows_simd8x8(acc[kk], sg_alpha, sg_tmp, lane); @@ -170,32 +164,33 @@ threadgroup_barrier(mem_flags::mem_threadgroup); float local = 0.0f; - for (int t = col_lane; t < KCHUNK; t += 4) { - float w = 0.0f; - if (grow < q_valid) w = metal::exp(scores[t / 4] - row_m); - sg_scores[row * KCHUNK + t] = w; - local += w; + if (lane < (uint)ROWS) { + int grow = qrow0 + (int)lane; + for (int t = 0; t < KCHUNK; t++) { + float w = 0.0f; + if (grow < q_valid) { + w = metal::exp(scores[t] - row_m); + } + sg_scores[(int)lane * KCHUNK + t] = w; + local += w; + } + row_lse += local; } - local += simd_shuffle_xor(local, 1); - local += simd_shuffle_xor(local, 2); - row_lse += local; simdgroup_barrier(mem_flags::mem_threadgroup); - for (int kc = 0; kc < KCHUNK; kc += 8) { - thread simdgroup_float8x8 pmat; - simdgroup_load(pmat, sg_scores + kc, KCHUNK); - for (int kk = 0; kk < 16; kk++) { - simdgroup_float8x8 vmat; - simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kc * D + kk * 8), D); - simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); - } + thread simdgroup_float8x8 pmat; + simdgroup_load(pmat, sg_scores, KCHUNK); + for (int kk = 0; kk < 16; kk++) { + simdgroup_float8x8 vmat; + simdgroup_load(vmat, (const threadgroup float*)(kvsmem + kk * 8), D); + simdgroup_multiply_accumulate(acc[kk], pmat, vmat, acc[kk]); } threadgroup_barrier(mem_flags::mem_threadgroup); } } - if (lane % 4 == 0) { - sg_alpha[(int)lane / 4] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; + if (lane < (uint)ROWS) { + sg_alpha[(int)lane] = row_lse > 0.0f ? 1.0f / row_lse : 0.0f; } simdgroup_barrier(mem_flags::mem_threadgroup); for (int kk = 0; kk < 16; kk++) { @@ -216,7 +211,6 @@ } simdgroup_barrier(mem_flags::mem_threadgroup); } - """ diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py deleted file mode 100644 index caac1541e8..0000000000 --- a/fastvideo/mlx_runtime/tests/test_minimax_h3_fp8_checkpoint.py +++ /dev/null @@ -1,33 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -"""Exercise native floating-point quantized H3 storage on supported MLX builds.""" -import numpy as np -import pytest - -mx = pytest.importorskip('mlx.core') -from fastvideo.mlx_runtime.fastwan import MLXQuantizationSpec, ensure_quantization_supported, linear -from fastvideo.mlx_runtime.minimax_h3 import MLXMiniMaxH3DiT, load_mlx_h3_checkpoint, quantize_matrix, save_mlx_h3_checkpoint - - -@pytest.mark.parametrize('mode', ['mxfp8', 'mxfp4', 'nvfp4']) -def test_float_quantized_checkpoint_preserves_matrix(tmp_path, mode): - spec = MLXQuantizationSpec.from_name(mode) - ensure_quantization_supported(spec) - dense = (mx.random.normal((64, 64)) * 0.001).astype(mx.bfloat16) - weight = quantize_matrix(dense, spec) - restored = mx.dequantize(weight.weight, weight.scales, mode=mode).astype(mx.float32) * weight.global_scale - relative_error = mx.sqrt(mx.sum((restored - dense.astype(mx.float32))**2) / mx.sum(dense.astype(mx.float32)**2)) - assert float(relative_error.item()) < 0.15 - x = mx.random.normal((3, 64)).astype(mx.bfloat16) - config = dict(hidden_size=64, num_attention_heads=1, attention_head_dim=64, ffn_dim=128, - in_channels=24, audio_in_channels=24, patch_size=[1, 1, 1], text_dim=64, - freq_dim=64, time_embed_dim=64, rope_freq_dim=4, rope_theta=10000., - norm_eps=1e-5, qk_norm_eps=1e-5, final_norm_eps=1e-5) - dit = MLXMiniMaxH3DiT({'test.weight': weight}, [], [], config) - save_mlx_h3_checkpoint(dit, tmp_path) - loaded = load_mlx_h3_checkpoint(tmp_path) - actual = linear(x, loaded.weights['test.weight']).astype(mx.float32) - expected = linear(x, weight).astype(mx.float32) - np.testing.assert_array_equal(np.array(actual), np.array(expected)) - assert loaded.weights['test.weight'].biases is None - assert loaded.weights['test.weight'].spec == spec - assert loaded.weights['test.weight'].global_scale == weight.global_scale diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index 5be7d672c8..293a8e66de 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -330,23 +330,3 @@ def test_converter_continues_past_mismatched_existing_format(tmp_path, monkeypat converter.main() assert saved == ["int6"] assert (existing / h3.H3_WEIGHTS_FILENAME).read_bytes() == b"existing" - - -@pytest.mark.parametrize("dtype", [mx.float16, mx.bfloat16]) -@pytest.mark.parametrize("exempt", [False, True]) -def test_simd_partial_key_chunks_match_reference(dtype, exempt): - """Nonuniform scores and partial tiles exercise all four 8-key fragments.""" - _require_metal() - geometry = vsa.build_h3_tile_geometry((7, 5), (5, 3, 7), 64) - mx.random.seed(3026) - q, k, value = [mx.random.normal((geometry.total_seq_length, 2, 128)).astype(dtype) - for _ in range(3)] - expected = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="reference") - stats = vsa.MiniMaxH3VSAStats() - actual = vsa.h3_vsa_attention(q, k, value, geometry, sparsity=.5, exempt=exempt, impl="simd", stats=stats) - mx.eval(actual, expected) - assert stats.impl == "simd" and stats.dense_fallback_reason is None - assert mx.all(mx.isfinite(actual)).item() - # The Metal kernel uses FP32 accumulation with a different reduction order. - np.testing.assert_allclose(np.asarray(actual.astype(mx.float32)), - np.asarray(expected.astype(mx.float32)), atol=.01, rtol=.01) diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index 4a43adaea2..931da4aa23 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -18,17 +18,6 @@ INT8/INT6/INT4 grid, and record ``vsa.capable`` in the manifest. Write VSA checkpoints to a new directory — do not overwrite an existing dense export. -An FP8 transformer with per-channel ``weight_scale`` can also be a source. -The loader dequantizes each FP8 matrix before applying the requested MLX -quantization. This saves download bytes but quantizes twice; compare its clips -with the BF16-sourced export before using it for release. - -``--formats "mxfp8 mxfp4 nvfp4"`` tries native MLX floating-point quantized -storage and matrix multiplication. These formats are experimental and require -operator support from the installed MLX build. This converts BF16 or FP8 source -weights; it does not import CUDA-packed NVFP4 DiT exports. The default formats -remain affine INT8/INT6/INT4. - python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \\ --model-root ~/models/FastH3-Preview-v0.2/transformer \\ --out ~/models/FastH3-MLX-vsa \\ @@ -64,8 +53,8 @@ logger = init_logger(__name__) -SUPPORTED_FORMATS = ("int8", "int6", "int4", "mxfp8", "mxfp4", "nvfp4") -DEFAULT_FORMATS = "int8 int6 int4" +SUPPORTED_FORMATS = ("int8", "int6", "int4") +DEFAULT_FORMATS = " ".join(SUPPORTED_FORMATS) def _adaln_cache_timesteps(model_root: str | Path | None = None) -> np.ndarray: From 8f33ee9e3d2dfdb3f28850e59d5f4d7f4b744fa8 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Mon, 5 Oct 2026 13:33:35 -0700 Subject: [PATCH 097/107] [bugfix]: reject packed H3 transformer sources --- fastvideo/mlx_runtime/minimax_h3.py | 3 +++ .../tests/test_minimax_h3_bf16_source.py | 17 +++++++++++++++++ 2 files changed, 20 insertions(+) create mode 100644 fastvideo/mlx_runtime/tests/test_minimax_h3_bf16_source.py diff --git a/fastvideo/mlx_runtime/minimax_h3.py b/fastvideo/mlx_runtime/minimax_h3.py index 5b5e546cc6..c31180de31 100644 --- a/fastvideo/mlx_runtime/minimax_h3.py +++ b/fastvideo/mlx_runtime/minimax_h3.py @@ -1392,6 +1392,9 @@ def assign(key: str, value) -> None: for key, source in shard_arrays.items(): if _is_ignored_dense_key(key, include_vsa=include_vsa): continue + if source.dtype == mx.uint8 and key.endswith(".weight"): + raise ValueError(f"Packed transformer weight {key} is not a floating-point source. " + "Convert the released BF16 transformer to MLX INT6 instead.") if temb is not None and (key.startswith("time_embedder.") or key == "adaln_basis.weight"): continue if key.startswith("transformer_blocks."): diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_bf16_source.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_bf16_source.py new file mode 100644 index 0000000000..9c9646b7ca --- /dev/null +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_bf16_source.py @@ -0,0 +1,17 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Reject packed source bytes before they can become invalid affine weights.""" +import pytest + +mx = pytest.importorskip("mlx.core") +from fastvideo.mlx_runtime.minimax_h3 import mlx_h3_dit_from_diffusers_safetensors + + +def test_packed_transformer_source_is_rejected(tmp_path): + mx.save_safetensors(str(tmp_path / "diffusion_pytorch_model.safetensors"), { + "transformer_blocks.0.attn.to_q.weight": mx.zeros((128, 128), dtype=mx.uint8), + }) + with pytest.raises(ValueError, match="released BF16 transformer"): + mlx_h3_dit_from_diffusers_safetensors( + tmp_path, config={"num_layers": 1, "num_refiner_layers": 0}, + quantization="int6", include_vsa=True, + ) From ec7c971ec15700451d66dd9647deb30a746ef4d4 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Mon, 5 Oct 2026 13:33:35 -0700 Subject: [PATCH 098/107] [docs]: use released Track C INT6 Mac recipe --- docs/getting_started/installation/mlx.md | 128 ++++++++++------------- 1 file changed, 54 insertions(+), 74 deletions(-) diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index a2fb02206a..4e1de0be09 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -46,108 +46,88 @@ is the higher-quality FastH3. Recorded shapes and evidence live in the [support matrix](../../inference/support_matrix.md#apple-silicon-native-runtime). -## Pruned eight-forward checkpoint +## FastH3 V2 and Trim with INT6 -The pruned FastH3 checkpoint has 42 transformer blocks and rank-16 AdaLN. -Its `fastvideo_inference.json` fixes eight denoising forwards, video/audio -shifts of 10/3, and VSA sparsity 0.8. Keep that file beside the transformer -when converting. The converter reads its schedule to build the AdaLN cache. +The released sources are `FastVideo/FastVideo-FastH3-8-Step-V2` and +`FastVideo/FastVideo-FastH3-Trim-8-Step`. Trim has 42 transformer blocks and +rank-16 AdaLN. Both use eight denoising forwards, video/audio shifts of 10/3, +VSA sparsity 0.8, a native NVFP4 text encoder, and the 26-layer light video VAE. +Keep `fastvideo_inference.json` beside `transformer/`; conversion reads its +schedule to build the AdaLN cache. + +Convert the BF16 transformer to affine INT6 with its VSA gates: ```bash -hf download FastVideo/FastH3-Pruned-8Step-BF16-ckpt300 \ - --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ - --exclude 'text_encoder/*' - -# Optional BF16 encoder fallback: stream the first 50 language layers. -# The last three shards are unused. The packed NVFP4 option is described below. -hf download MiniMaxAI/MiniMax-H3 \ - --local-dir ./FastH3-Pruned-8Step-BF16-ckpt300 \ - --include 'text_encoder/model-0000[1-9]-of-00014.safetensors' \ - --include 'text_encoder/model-0001[0-1]-of-00014.safetensors' \ - --include 'text_encoder/model.safetensors.index.json' \ - --include 'text_encoder/config.json' +hf download FastVideo/FastVideo-FastH3-Trim-8-Step \ + --local-dir ./FastH3-Trim python scripts/checkpoint_conversion/convert_minimax_h3_mlx.py \ - --model-root ./FastH3-Pruned-8Step-BF16-ckpt300/transformer \ - --out ./FastH3-Pruned-MLX-vsa \ - --formats "int8 int6" --include-vsa - -python examples/inference/basic/mlx_fasth3.py \ - --model-root ./FastH3-Pruned-8Step-BF16-ckpt300 \ - --mlx-checkpoint ./FastH3-Pruned-MLX-vsa/int8 \ - --prompt "(S1) A potter asks [English] Is the rim ready?" \ - --height 480 --width 832 --num-frames 243 --steps 8 \ - --vsa --vsa-sparsity 0.8 --vsa-tile-size 64 \ - --output-path ./outputs/fasth3_pruned_int8_480p.mp4 + --model-root ./FastH3-Trim/transformer \ + --out ./FastH3-Trim-MLX \ + --formats "int6" --include-vsa ``` -At 24 fps, 124 frames is the legal H3 count for a roughly five-second clip. -Use `--num-frames 124` and a separate output path for that run. The `--fast` -and `--fast-spatial` options change the workload and are not part of the -native-resolution benchmark. A 36 GB Mac may need INT6 and phased loading; -measure memory before claiming all-resident operation. - -### Packed encoder and resident loading +For V2, use `FastVideo/FastVideo-FastH3-8-Step-V2` and separate source/output +directories. Preconverted release snapshots use the same names with the +`-MLX-INT6` suffix. Each snapshot includes the encoder in MLX layout, both VAEs, +and the trained schedule, so it does not require a second encoder download. -The experimental MLX conditioner can read the released FastVideo NVFP4 -text encoder directly, using native `nvfp4` matrix multiplication. It keeps -the packed weights and BF16 embedding table in memory, with FP32 -activations. CUDA uses quantized activations, so the two encoders are not -bit-exact. Validate generated video and audio before publishing a timing. -MLX 0.32.2 supports the required operator on Apple Silicon. - -Pass the packed encoder directory as `conditioner_dir`; `conditioner_mode="auto"` -selects it from `config.json`. The BF16 fallback continues to stream layers. -To request all-resident generation through the Python API: +The 36 GiB M4 Max release recipe uses phased placement. It loads the encoder, +DiT, and decoders in turn. Use reference attention and native output geometry: ```python +from pathlib import Path from fastvideo.mlx_runtime.minimax_h3_pipeline import MiniMaxH3MLXPipeline +root = Path("./FastH3-Trim") pipeline = MiniMaxH3MLXPipeline( - model_root="./FastH3-Pruned-8Step-BF16-ckpt300", - mlx_dit_checkpoint="./FastH3-Pruned-MLX-vsa/int6", - conditioner_dir="./FastH3-NVFP4-encoder", + model_root=root, + mlx_dit_checkpoint="./FastH3-Trim-MLX/int6", conditioner_mode="nvfp4", - resident=True, + resident=False, vae_dtype="fp16", + metal_wired_limit_gib=27, ) try: - pipeline.prepare_resident() # Load encoder, DiT, video VAE and audio VAE. - result = pipeline.generate( - "(S1) A potter asks [English] Is the rim ready?", - output_path="./outputs/fasth3_pruned_resident.mp4", - height=480, width=832, num_frames=243, num_steps=8, - vsa=True, vsa_sparsity=0.8, vsa_tile_size=64, + pipeline.generate( + "A corgi news anchor sits behind a desk and gives a cheerful bark.", + output_path="./outputs/trim-int6-corgi.mp4", + width=832, height=480, num_frames=124, seed=1234, + num_steps=8, vsa=True, vsa_sparsity=0.8, vsa_tile_size=64, + vsa_impl="reference", vae_tile_height=256, vae_tile_width=256, ) finally: pipeline.close() ``` -Resident placement requires space for activations as well as all four -components. On a 36 GiB Mac, try INT6 first and measure peak allocation. -If loading or inference runs out of memory, use phased loading by leaving -`resident=False`. Changing placement does not change frames or resolution. +Set `FASTVIDEO_MLX_DQ_GEMM=1` before running this Python command. It selects the +validated affine dequantization followed by dense matrix multiplication. +124 frames at 24 fps is roughly five seconds. The recipe preserves all frames +and the requested resolution. -### Metal wired memory +### Native NVFP4 encoder and cache -MLX's allocation limit and wired-memory limit are separate. H3's optional -`metal_wired_limit_gib` calls `mx.set_wired_limit` so selected Metal allocations -stay in physical memory. It does not increase available RAM. Explicit requests -fail visibly if the installed MLX build cannot apply them. +The MLX conditioner reads the released packed NVFP4 weights without +requantization. It retains the layers H3 reads and can cache the packed weights +in MLX layout. The cache is written in a staging directory and published by a +single rename. A cache hit changes storage layout, not encoder arithmetic. -Inspect the device's recommended working set before choosing a limit: +MLX uses BF16 embeddings and FP32 activations; CUDA uses quantized activations. +Generated video and audio must be reviewed before claiming cross-runtime +quality parity. MLX 0.32.2 supports the required operator on Apple Silicon. -```python -import mlx.core as mx +### Metal wired memory -print(mx.device_info()) -``` +MLX's allocation limit and wired-memory limit are separate. The optional +`metal_wired_limit_gib` calls `mx.set_wired_limit` to keep selected Metal +allocations in physical memory. It does not add RAM. An explicit request fails +if the installed MLX build cannot apply it. `close()` restores the previous +wired limit. -For the tested 36 GiB M4 Max, phased generation can request -`metal_wired_limit_gib=27` in `MiniMaxH3MLXPipeline`. Leave room for macOS and -other applications. All-resident generation also needs room for the encoder, -DiT, both decoders and peak activations; wiring cannot make an oversized stack -fit. An omitted wired limit preserves MLX's existing wiring setting. +The tested 36 GiB M4 Max recipe uses 27 GiB and phased placement. Resident +placement also requires capacity for all components and peak activations; +wiring cannot make an oversized stack fit. Inspect `mx.device_info()` before +choosing a limit on another Mac. ## Hardware From 17987f6f0e48b84e2d72a003cbeb5c15c9b86bf3 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 15:13:39 -0700 Subject: [PATCH 099/107] [test]: MLX dequant-GEMM test allows CPU BF16 accumulation error --- fastvideo/tests/mlx/test_mlx_affine_dq_gemm.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/fastvideo/tests/mlx/test_mlx_affine_dq_gemm.py b/fastvideo/tests/mlx/test_mlx_affine_dq_gemm.py index 9c8d3813ab..81fade2f2d 100644 --- a/fastvideo/tests/mlx/test_mlx_affine_dq_gemm.py +++ b/fastvideo/tests/mlx/test_mlx_affine_dq_gemm.py @@ -73,7 +73,11 @@ def test_dq_gemm_matches_qmm_for_supported_bit_widths(bits: int, group_size: int ref_np = np.asarray(ref.astype(mx.float32)) got_np = np.asarray(got.astype(mx.float32)) scale = max(float(np.max(np.abs(ref_np))), 1e-3) - assert rel < 2e-2, rel + # MLX's CPU backend accumulates the BF16 dense GEMM less precisely than Metal; at group size 128 + # (512 inputs) its error against qmm reaches about 2.6% for every bit width, including 8, so the + # gap is accumulation rather than quantization. + rel_limit = 3e-2 if mx.default_device() == mx.cpu else 2e-2 + assert rel < rel_limit, rel assert float(np.max(np.abs(got_np - ref_np))) / scale < 0.08 From 0881b38c3f951992538dff6f81268f468024ec40 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Tue, 6 Oct 2026 11:37:23 -0700 Subject: [PATCH 100/107] [bugfix]: make FastH3 MLX limits reversible and restore CI coverage --- .github/workflows/ci-macos-mlx.yml | 8 +-- .gitignore | 1 + docs/getting_started/installation/mlx.md | 4 +- fastvideo/mlx_runtime/minimax_h3_pipeline.py | 61 +++++++++++++------ .../mlx/test_mlx_minimax_h3_fast_mode.py | 49 ++++++++++++++- .../ssim/.reference_videos_download.lock | 0 6 files changed, 98 insertions(+), 25 deletions(-) delete mode 100644 fastvideo/tests/ssim/.reference_videos_download.lock diff --git a/.github/workflows/ci-macos-mlx.yml b/.github/workflows/ci-macos-mlx.yml index 3ef19a3118..7d4da9ded7 100644 --- a/.github/workflows/ci-macos-mlx.yml +++ b/.github/workflows/ci-macos-mlx.yml @@ -41,7 +41,6 @@ jobs: - uses: actions/setup-python@v5 with: python-version: "3.12" - cache: pip - uses: astral-sh/setup-uv@v3 @@ -49,7 +48,7 @@ jobs: run: | uv pip install --system \ --index-url https://download.pytorch.org/whl/cpu \ - torch==2.11.0 torchvision torchaudio + torch==2.12.0 torchvision torchaudio uv pip install --system \ pytest numpy scipy pillow imageio einops cloudpickle filelock \ PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \ @@ -74,6 +73,7 @@ jobs: - name: Run MLX smoke tests run: | python -m pytest \ + fastvideo/mlx_runtime/tests/ \ fastvideo/tests/mlx/test_dmd_sampling.py \ fastvideo/tests/mlx/test_memory_limits.py \ fastvideo/tests/mlx/test_quant_capability.py \ @@ -122,7 +122,6 @@ jobs: - uses: actions/setup-python@v5 with: python-version: "3.12" - cache: pip - uses: astral-sh/setup-uv@v3 @@ -130,7 +129,7 @@ jobs: run: | uv pip install --system \ --index-url https://download.pytorch.org/whl/cpu \ - torch==2.11.0 torchvision torchaudio + torch==2.12.0 torchvision torchaudio uv pip install --system \ pytest numpy scipy pillow imageio einops cloudpickle filelock \ PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \ @@ -139,6 +138,7 @@ jobs: - name: Run MLX smoke tests (CPU backend) run: | python -m pytest \ + fastvideo/mlx_runtime/tests/ \ fastvideo/tests/mlx/test_dmd_sampling.py \ fastvideo/tests/mlx/test_memory_limits.py \ fastvideo/tests/mlx/test_quant_capability.py \ diff --git a/.gitignore b/.gitignore index 88035e8c2b..098369a7e2 100644 --- a/.gitignore +++ b/.gitignore @@ -134,6 +134,7 @@ openspec/ fastvideo/tests/ssim/reference_videos/** !fastvideo/tests/ssim/reference_videos/**/*.mp4 !fastvideo/tests/ssim/reference_videos/**/*.png +fastvideo/tests/ssim/.reference_videos_download.lock # Local H3 MLX kernel / exactness benches (JSON, logs, frames, videos) .kernel_bench/ diff --git a/docs/getting_started/installation/mlx.md b/docs/getting_started/installation/mlx.md index 4e1de0be09..fb94d3e9f1 100644 --- a/docs/getting_started/installation/mlx.md +++ b/docs/getting_started/installation/mlx.md @@ -122,7 +122,9 @@ MLX's allocation limit and wired-memory limit are separate. The optional `metal_wired_limit_gib` calls `mx.set_wired_limit` to keep selected Metal allocations in physical memory. It does not add RAM. An explicit request fails if the installed MLX build cannot apply it. `close()` restores the previous -wired limit. +wired limit. The phased pipeline also sets a 30 GiB maximum allocator guideline +and restores its previous value on close. Resident placement keeps the existing +allocator limit so larger Macs can hold all components. The tested 36 GiB M4 Max recipe uses 27 GiB and phased placement. Resident placement also requires capacity for all components and peak activations; diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index 4985461589..02708e2d1e 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -245,33 +245,40 @@ def _default_metal_wired_limit_gib(mx) -> float: return min(30.0, 0.84 * total_bytes / 2**30) -def _configure_metal_memory_limits(mx, wired_limit_gib: float | None) -> int | None: - """Keep allocator capacity separate from explicitly requested wired residency. +def _configure_metal_memory_limits(mx, + wired_limit_gib: float | None, + *, + resident: bool = False) -> tuple[int | None, int | None]: + """Set Metal limits and return their previous process-wide values.""" + if wired_limit_gib is not None and (not math.isfinite(wired_limit_gib) or wired_limit_gib <= 0): + raise ValueError("metal_wired_limit_gib must be finite and positive") + set_wired = getattr(mx, "set_wired_limit", None) + if set_wired is None and hasattr(mx, "metal"): + set_wired = getattr(mx.metal, "set_wired_limit", None) + if wired_limit_gib is not None and set_wired is None: + raise RuntimeError("This MLX build cannot set the requested wired-memory limit") - Returns the previous wired limit in bytes when one was set, so the caller can - restore it; the limit is process-wide and would otherwise outlive the pipeline. - """ + previous_memory = None set_memory = getattr(mx, "set_memory_limit", None) if set_memory is None and hasattr(mx, "metal"): set_memory = getattr(mx.metal, "set_memory_limit", None) - if set_memory is not None: + if set_memory is not None and not resident: try: - set_memory(int(_default_metal_wired_limit_gib(mx) * 2**30)) + previous_memory = int(set_memory(int(_default_metal_wired_limit_gib(mx) * 2**30))) except Exception as error: # noqa: BLE001 - older MLX best effort logger.info("Could not set the Metal allocation limit: %s", error) if wired_limit_gib is None: - return None - if not math.isfinite(wired_limit_gib) or wired_limit_gib <= 0: - raise ValueError("metal_wired_limit_gib must be finite and positive") - set_wired = getattr(mx, "set_wired_limit", None) - if set_wired is None and hasattr(mx, "metal"): - set_wired = getattr(mx.metal, "set_wired_limit", None) - if set_wired is None: - raise RuntimeError("This MLX build cannot set the requested wired-memory limit") + return previous_memory, None # Explicit requests must succeed; do not silently benchmark an unwired model. - previous = set_wired(int(wired_limit_gib * 2**30)) + assert set_wired is not None + try: + previous = set_wired(int(wired_limit_gib * 2**30)) + except Exception: + if previous_memory is not None and set_memory is not None: + set_memory(previous_memory) + raise logger.info("MLX wired limit %.2f GiB (previous %.2f GiB)", wired_limit_gib, previous / 2**30) - return int(previous) + return previous_memory, int(previous) def _restore_metal_wired_limit(mx, previous_bytes: int | None) -> None: @@ -284,6 +291,16 @@ def _restore_metal_wired_limit(mx, previous_bytes: int | None) -> None: set_wired(previous_bytes) +def _restore_metal_memory_limit(mx, previous_bytes: int | None) -> None: + if previous_bytes is None: + return + set_memory = getattr(mx, "set_memory_limit", None) + if set_memory is None and hasattr(mx, "metal"): + set_memory = getattr(mx.metal, "set_memory_limit", None) + if set_memory is not None: + set_memory(previous_bytes) + + MINIMAX_H3_PROMPT_CACHE_VERSION = "v2-attention-layout" @@ -373,7 +390,6 @@ def __init__( ) -> None: import mlx.core as mx - self._previous_wired_limit = _configure_metal_memory_limits(mx, metal_wired_limit_gib) self.model_root = Path(model_root) if conditioner_mode not in ("auto", "streamed", "nvfp4"): raise ValueError(f"Unknown H3 conditioner mode: {conditioner_mode}") @@ -406,6 +422,9 @@ def __init__( self._dit_in_channels = int(dit_config["in_channels"]) self.last_dit_forward_s = 0.0 self.last_vsa_stats: dict[str, Any] | None = None + self._previous_memory_limit, self._previous_wired_limit = _configure_metal_memory_limits(mx, + metal_wired_limit_gib, + resident=resident) # -- input validation (before anything heavy loads) ------------------- @@ -502,6 +521,12 @@ def close(self) -> None: _restore_metal_wired_limit(mx, previous) self._previous_wired_limit = None + previous_memory = getattr(self, "_previous_memory_limit", None) + if previous_memory is not None: + import mlx.core as mx + + _restore_metal_memory_limit(mx, previous_memory) + self._previous_memory_limit = None def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: """Returns (hidden states (S, hidden), token tags). Uses the cache or diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py index 6e59682545..1682372092 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py @@ -24,6 +24,8 @@ _configure_metal_memory_limits, _default_metal_wired_limit_gib, _preflight_media_dependencies, + _restore_metal_memory_limit, + _restore_metal_wired_limit, _validate_checkpoint_step_ladder, plan_fast_temporal, ) @@ -150,11 +152,25 @@ def test_explicit_wired_limit_uses_wired_api_separately_from_allocator(): calls = [] fake = SimpleNamespace( metal=SimpleNamespace(device_info=lambda: {"memory_size": 36 * 2**30}), - set_memory_limit=lambda size: calls.append(("allocator", size)), + set_memory_limit=lambda size: calls.append(("allocator", size)) or 40 * 2**30, set_wired_limit=lambda size: calls.append(("wired", size)) or 0, ) - _configure_metal_memory_limits(fake, 27.0) + previous_memory, previous_wired = _configure_metal_memory_limits(fake, 27.0) assert calls == [("allocator", 30 * 2**30), ("wired", 27 * 2**30)] + assert (previous_memory, previous_wired) == (40 * 2**30, 0) + _restore_metal_wired_limit(fake, previous_wired) + _restore_metal_memory_limit(fake, previous_memory) + assert calls[-2:] == [("wired", 0), ("allocator", 40 * 2**30)] + + +def test_resident_placement_keeps_the_existing_allocator_limit(): + calls = [] + fake = SimpleNamespace( + set_memory_limit=lambda size: calls.append(("allocator", size)) or 40 * 2**30, + set_wired_limit=lambda size: calls.append(("wired", size)) or 0, + ) + assert _configure_metal_memory_limits(fake, 36.0, resident=True) == (None, 0) + assert calls == [("wired", 36 * 2**30)] def test_explicit_wired_limit_failure_is_not_silently_ignored(): @@ -167,3 +183,32 @@ def reject(size): _configure_metal_memory_limits(fake, float("nan")) with pytest.raises(RuntimeError, match="cannot set"): _configure_metal_memory_limits(SimpleNamespace(), 27.0) + + +def test_wired_limit_failure_restores_allocator_limit(): + calls = [] + + def set_memory(size): + calls.append(("allocator", size)) + return 40 * 2**30 + + def reject(size): + calls.append(("wired", size)) + raise ValueError("exceeds system wired limit") + + fake = SimpleNamespace(metal=SimpleNamespace(device_info=lambda: {"memory_size": 36 * 2**30}), + set_memory_limit=set_memory, + set_wired_limit=reject) + with pytest.raises(ValueError, match="system wired limit"): + _configure_metal_memory_limits(fake, 27.0) + assert calls == [("allocator", 30 * 2**30), ("wired", 27 * 2**30), ("allocator", 40 * 2**30)] + + +def test_invalid_pipeline_options_do_not_change_process_memory_limits(monkeypatch): + calls = [] + monkeypatch.setattr("fastvideo.mlx_runtime.minimax_h3_pipeline._configure_metal_memory_limits", + lambda *args, **kwargs: calls.append((args, kwargs))) + with pytest.raises(ValueError, match="Unknown H3 conditioner mode"): + MiniMaxH3MLXPipeline(model_root="missing", mlx_dit_checkpoint="missing", conditioner_mode="invalid", + metal_wired_limit_gib=27.0) + assert calls == [] diff --git a/fastvideo/tests/ssim/.reference_videos_download.lock b/fastvideo/tests/ssim/.reference_videos_download.lock deleted file mode 100644 index e69de29bb2..0000000000 From 40cffbee2746c35da249810a6a6da001531f7dd4 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Tue, 6 Oct 2026 11:42:46 -0700 Subject: [PATCH 101/107] [misc]: keep RTX showcase prompts in the merged base --- .../minimax_h3_pro6000/showcase_prompts.json | 12 ------------ 1 file changed, 12 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json b/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json index be61105719..048f7d9d2c 100644 --- a/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json +++ b/scripts/benchmarks/minimax_h3_pro6000/showcase_prompts.json @@ -1,16 +1,4 @@ { - "violinist-archway": "A street violinist plays a short, bright melody under a stone archway at dusk. Medium shot from a fixed camera at chest height. Warm string lights hang overhead and a few pedestrians pass softly out of focus in the background. The bow moves in long, even strokes, and the violinist's fingers press precisely on the strings. The sound is a clear solo violin with light echo from the stone, distant footsteps and a faint murmur of the street. No other instruments. The violinist finishes the phrase, lowers the bow slightly and smiles.", - "chef-tasting": "In a bright home kitchen, a chef in a white jacket stands behind a wooden counter and looks straight at the camera. She lifts a wooden spoon from a simmering pot, tastes the sauce and says, \"A little more lemon, and it's perfect.\" Steam rises from the pot. The camera is fixed at eye level in a medium close-up. Soft bubbling, the light tap of the spoon on the pot rim and her clear, warm voice; no music.", - "fox-snow": "A red fox stands in a quiet snowy meadow at sunrise, listening with its head tilted. It crouches, then leaps high and dives nose-first into the deep snow, sending up a spray of powder, and pops back up with snow on its face. Wide shot from a low fixed camera, soft golden light, breath visible in the cold air. Sounds: gentle wind, the crunch of snow, a soft thump as the fox lands, a distant bird call. No music.", - "surfer-wave": "A single surfer rides a glassy turquoise wave at midday, carving a smooth turn along the face as the lip curls behind them. The camera follows from the water at board level, keeping the surfer centered. Spray catches the sunlight. Sounds: the roar of the breaking wave, rushing water and wind, a short exhilarated shout from the surfer. No music.", - "robot-windowsill": "3D animated short. A small round robot with big glowing eyes waters a single potted sunflower on a sunny windowsill. It tilts a tiny watering can carefully, the flower perks up, and the robot does a happy little wiggle. Fixed medium shot, soft cinematic lighting, gentle depth of field. Sounds: water trickling, soft mechanical whirs and a cheerful two-note beep. Light playful music box melody.", - "jazz-drummer": "A jazz drummer in a dim club plays a crisp brush pattern on the snare, then a short fill across the toms ending on a soft cymbal crash. Medium shot from a fixed camera slightly to the side, warm amber stage light, a hint of haze. Only drums are heard: the swish of brushes, tom hits and the shimmering cymbal, with a faint murmur from the audience.", - "rain-cafe": "A young woman sits by a rain-streaked café window, writing in a leather notebook with a fountain pen. She pauses, looks out at the street, then back at the camera and says softly, \"Some days are made for staying in.\" Close-up from a fixed camera, soft gray daylight with warm interior lamps. Sounds: steady rain on the glass, the scratch of the pen, a cup set down gently in the background. No music.", - "horse-beach": "A dark brown horse gallops along the edge of a wide beach at sunrise, splashing through shallow water. The camera tracks alongside at a steady distance, keeping the horse in profile. Low golden light, long shadows, spray from the hooves. Sounds: rhythmic hoofbeats on wet sand, splashing water, waves breaking, the horse's breathing. No music.", - "anime-rooftop": "2D anime style. A girl with short black hair and large headphones around her neck stands on a city rooftop at sunset, her jacket moving in the wind. She looks over the skyline, then turns to the camera and says, \"Tomorrow, we try again.\" Fixed medium shot, vivid orange and purple sky, clean cel shading. Sounds: wind, distant city traffic, a soft piano note under her line.", - "potter-hands": "Close-up of a potter's hands shaping a spinning clay bowl on a wheel. Wet fingers press gently, and the rim rises and smooths with each rotation. The camera is fixed just above wheel height, warm window light from the left. Sounds: the steady hum of the wheel, wet clay sliding under fingers, a soft splash. No music, no speech.", - "astronaut-radio": "Inside a spacecraft cockpit, an astronaut in a white suit looks out the window at Earth rising over the horizon. She presses a button on the console and says into her headset, \"Houston, the view is incredible today.\" Medium close-up from a fixed camera, soft blue light from the window, small floating particles. Sounds: a quiet cabin hum, a short radio click before and after her line.", - "grandma-bread": "In a cozy kitchen, a grandmother kneads bread dough on a floured wooden table, folding and pressing it with practiced hands. She looks up toward someone off camera, smiles and says, \"Patience makes the best bread.\" Medium shot from a fixed camera, warm afternoon light, flour dust drifting in the air. Sounds: dough thumping on the table, a ticking wall clock, her gentle voice. No music.", "corgi-weather": "Live-action TV news parody, 16:9, bright studio lighting, 35mm lens, locked-off medium shot. A fluffy tan-and-white corgi with a tiny red bow tie sits upright behind a glossy news desk, a weather map of puffy clouds glowing on the screen behind it. [0-2 s] The corgi looks straight into the camera, ears perked, and says in a cheerful deep announcer voice, \"Tomorrow: one hundred percent chance of zoomies.\" [2-5 s] It tilts its head, its tongue flops out, and it gives one proud little bark as a short brassy news jingle plays. Sound: crisp studio voice, the jingle, the single bark. Only one corgi. No subtitles, no on-screen text, no watermark, no extra animals.", "street-food-jingle": "Documentary-style phone video, 16:9, sunny afternoon, phone held still at chest height. A cheerful street-food vendor in a green apron and white bandana stands behind a stall of golden fried snacks. [0-2 s] She smiles at the camera and claps her hands twice to set a rhythm, swaying gently. [2-5 s] She sings a bouncy jingle, \"Crispy, crunchy, ready to munch!\", and finishes with a wink and a thumbs-up. Sound: her clear singing voice in a playful melody, two claps, sizzling oil and a soft market murmur. One vendor in focus, hands empty. No subtitles, no text, no watermark.", "parrot-pirate": "Cinematic live-action, 16:9, golden hour on the deck of a wooden pirate ship, 50mm lens, locked-off medium shot. A weathered pirate captain with a braided beard, red bandana and long coat holds a rolled treasure map; a bright green parrot sits still on his shoulder. [0-2 s] He taps the map and declares in a gravelly voice, \"X marks the spot!\" [2-5 s] The parrot bobs its head and squawks back, \"Spot! Spot!\", and the captain slowly turns to give it a tired look. Sound: his gravelly voice, the parrot's squawk, creaking wood and soft waves. One captain and one parrot. No subtitles, no text, no watermark.", From f10bf3e30a9104d0d1cf6c00e0b66390a75f28ac Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Tue, 6 Oct 2026 11:50:49 -0700 Subject: [PATCH 102/107] [ci]: bind MLX Gloo smoke tests to loopback and bound hangs --- .github/workflows/ci-macos-mlx.yml | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/.github/workflows/ci-macos-mlx.yml b/.github/workflows/ci-macos-mlx.yml index 7d4da9ded7..546b61a599 100644 --- a/.github/workflows/ci-macos-mlx.yml +++ b/.github/workflows/ci-macos-mlx.yml @@ -33,8 +33,9 @@ jobs: env: FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA TOKENIZERS_PARALLELISM: "false" - MASTER_ADDR: localhost + MASTER_ADDR: "127.0.0.1" MASTER_PORT: "29513" + GLOO_SOCKET_IFNAME: lo0 steps: - uses: actions/checkout@v4 @@ -50,7 +51,7 @@ jobs: --index-url https://download.pytorch.org/whl/cpu \ torch==2.12.0 torchvision torchaudio uv pip install --system \ - pytest numpy scipy pillow imageio einops cloudpickle filelock \ + pytest pytest-timeout numpy scipy pillow imageio einops cloudpickle filelock \ PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \ "ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0" @@ -101,7 +102,7 @@ jobs: fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \ fastvideo/tests/platforms/test_mps_vsa_error.py \ fastvideo/tests/platforms/test_cpu_sdpa.py \ - -v -s -o faulthandler_timeout=120 + -v -s --timeout=120 -o faulthandler_timeout=120 # Same tests on MLX's CPU backend. Hosted macOS runners are scarce and # slower to schedule; this Linux job gives fast PR signal on the identical @@ -131,7 +132,7 @@ jobs: --index-url https://download.pytorch.org/whl/cpu \ torch==2.12.0 torchvision torchaudio uv pip install --system \ - pytest numpy scipy pillow imageio einops cloudpickle filelock \ + pytest pytest-timeout numpy scipy pillow imageio einops cloudpickle filelock \ PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \ "ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0" @@ -166,4 +167,4 @@ jobs: fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \ fastvideo/tests/platforms/test_mps_vsa_error.py \ fastvideo/tests/platforms/test_cpu_sdpa.py \ - -v -s -o faulthandler_timeout=120 + -v -s --timeout=120 -o faulthandler_timeout=120 From 4e84449c9a1b5014d3304d7d927e6b0459399bc1 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Tue, 6 Oct 2026 11:59:33 -0700 Subject: [PATCH 103/107] [bugfix]: fall back from inaccurate M1 SIMD VSA --- fastvideo/mlx_runtime/minimax_h3_vsa_simd.py | 7 +++++++ .../tests/mlx/test_mlx_minimax_h3_vsa_regressions.py | 10 ++++++++++ 2 files changed, 17 insertions(+) diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index bb21b1912a..556d9a0a6f 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -242,6 +242,13 @@ def _simd_kernel() -> Any | None: if not mx.metal.is_available(): _SIMD_KERNEL_ERROR = "Metal is not available in this MLX build" return None + # The SIMD kernel compiles on M1 but fails the BF16 reference parity + # check on the macOS 15 hosted M1 runner. Keep this opt-in path on + # validated newer GPUs and use the reference VSA on M1. + device_name = str(mx.metal.device_info().get("device_name", "")) + if device_name.startswith("Apple M1"): + _SIMD_KERNEL_ERROR = "SIMD VSA does not match the BF16 reference on Apple M1" + return None if not hasattr(mx.fast, "metal_kernel"): _SIMD_KERNEL_ERROR = "mx.fast.metal_kernel is not available in this MLX build" return None diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index 293a8e66de..c25f777353 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -262,6 +262,16 @@ def test_no_metal_probe_is_cached(monkeypatch): assert not simd.simd_kernel_available() +def test_m1_simd_falls_back_before_compilation(monkeypatch): + monkeypatch.setattr(simd, "_SIMD_KERNEL", None) + monkeypatch.setattr(simd, "_SIMD_KERNEL_ERROR", None) + monkeypatch.setattr(mx.metal, "is_available", lambda: True) + monkeypatch.setattr(mx.metal, "device_info", lambda: {"device_name": "Apple M1"}) + monkeypatch.setattr(mx.fast, "metal_kernel", lambda **kwargs: pytest.fail("M1 must not compile SIMD VSA")) + assert not simd.simd_kernel_available() + assert "BF16 reference on Apple M1" in simd.simd_kernel_error() + + def _require_metal(): if not mx.metal.is_available() or mx.default_device() != mx.gpu: pytest.skip("requires actual Metal execution") From d10f4e42103a9d365ed823d36683ba6c6af8d0a4 Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Tue, 6 Oct 2026 12:03:44 -0700 Subject: [PATCH 104/107] [bugfix]: use reference VSA on small virtual Metal runners --- .github/workflows/ci-macos-mlx.yml | 4 ++-- fastvideo/mlx_runtime/minimax_h3_vsa_simd.py | 14 ++++++++------ .../mlx/test_mlx_minimax_h3_vsa_regressions.py | 12 +++++++++++- 3 files changed, 21 insertions(+), 9 deletions(-) diff --git a/.github/workflows/ci-macos-mlx.yml b/.github/workflows/ci-macos-mlx.yml index 546b61a599..29a3261ebc 100644 --- a/.github/workflows/ci-macos-mlx.yml +++ b/.github/workflows/ci-macos-mlx.yml @@ -65,8 +65,8 @@ jobs: print("machine:", platform.machine()) print("processor:", platform.processor()) print("mlx default device:", mx.default_device()) - memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable" - print("mlx memory_size:", memory_size) + device_info = mx.metal.device_info() if mx.metal.is_available() else "metal unavailable" + print("mlx device_info:", device_info) print("torch:", torch.__version__) print("torch mps available:", torch.backends.mps.is_available()) PY diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py index 556d9a0a6f..910384ac9d 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa_simd.py @@ -242,12 +242,14 @@ def _simd_kernel() -> Any | None: if not mx.metal.is_available(): _SIMD_KERNEL_ERROR = "Metal is not available in this MLX build" return None - # The SIMD kernel compiles on M1 but fails the BF16 reference parity - # check on the macOS 15 hosted M1 runner. Keep this opt-in path on - # validated newer GPUs and use the reference VSA on M1. - device_name = str(mx.metal.device_info().get("device_name", "")) - if device_name.startswith("Apple M1"): - _SIMD_KERNEL_ERROR = "SIMD VSA does not match the BF16 reference on Apple M1" + # The SIMD kernel compiles but fails BF16 parity on the 7 GiB virtual + # Metal device in GitHub's M1 runner. Physical M1 devices also use the + # reference path until this kernel is validated on that GPU family. + device_info = mx.metal.device_info() + device_name = str(device_info.get("device_name", "")) + memory_size = int(device_info.get("memory_size", 0)) + if device_name.startswith("Apple M1") or 0 < memory_size <= 8 * 2**30: + _SIMD_KERNEL_ERROR = "SIMD VSA BF16 parity is unvalidated on Apple M1 or 8 GiB Metal devices" return None if not hasattr(mx.fast, "metal_kernel"): _SIMD_KERNEL_ERROR = "mx.fast.metal_kernel is not available in this MLX build" diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index c25f777353..717ecb95c4 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -269,7 +269,17 @@ def test_m1_simd_falls_back_before_compilation(monkeypatch): monkeypatch.setattr(mx.metal, "device_info", lambda: {"device_name": "Apple M1"}) monkeypatch.setattr(mx.fast, "metal_kernel", lambda **kwargs: pytest.fail("M1 must not compile SIMD VSA")) assert not simd.simd_kernel_available() - assert "BF16 reference on Apple M1" in simd.simd_kernel_error() + assert "BF16 parity is unvalidated" in simd.simd_kernel_error() + + +def test_small_virtual_metal_simd_falls_back_before_compilation(monkeypatch): + monkeypatch.setattr(simd, "_SIMD_KERNEL", None) + monkeypatch.setattr(simd, "_SIMD_KERNEL_ERROR", None) + monkeypatch.setattr(mx.metal, "is_available", lambda: True) + monkeypatch.setattr(mx.metal, "device_info", lambda: {"device_name": "Virtual GPU", "memory_size": 7 * 2**30}) + monkeypatch.setattr(mx.fast, "metal_kernel", lambda **kwargs: pytest.fail("small GPU must not compile SIMD VSA")) + assert not simd.simd_kernel_available() + assert "8 GiB Metal devices" in simd.simd_kernel_error() def _require_metal(): From ddd67fa3143a8a78c1abced21611452f70100084 Mon Sep 17 00:00:00 2001 From: SolitaryThinker Date: Tue, 6 Oct 2026 18:55:47 +0000 Subject: [PATCH 105/107] [bugfix]: address review findings on FastH3 Spark and MLX paths - MLX resident DiT: reload when a new step ladder needs dropped AdaLN weights; reset VSA to the default when a request passes none; reject a caller DiT whose AdaLN weights were already dropped for another ladder. - An explicit Metal wired limit also raises the allocator cap. - Prompt cache keyed by the effective encoder kind and encoder/tokenizer files. - NVFP4 encoder config and scale unswizzle accept what the CUDA loader accepts. - Encoder cache published with umask permissions; converter re-runs skip it. - mlx_fasth3 CLI restores Metal limits; VSA gather budget scales with memory. - One fenced NVFP4 quantize helper (cached GB10 check) for inference and QAT. - Spark benchmark: decode time excludes post-decode frame processing. Co-Authored-By: Claude Opus 5.5 --- .../basic/benchmark_fasth3_spark_nvfp4.py | 10 +- examples/inference/basic/mlx_fasth3.py | 59 ++++---- fastvideo/layers/fp4linear.py | 17 +-- fastvideo/layers/quantization/nvfp4_config.py | 61 +++++--- .../mlx_runtime/minimax_h3_conditioner.py | 25 ++-- fastvideo/mlx_runtime/minimax_h3_pipeline.py | 139 ++++++++++++------ fastvideo/mlx_runtime/minimax_h3_vsa.py | 21 ++- .../test_benchmark_fasth3_spark_nvfp4.py | 29 ++++ .../mlx/test_mlx_minimax_h3_conditioner.py | 82 ++++++++++- .../mlx/test_mlx_minimax_h3_fast_mode.py | 86 +++++++++++ .../test_mlx_minimax_h3_vsa_regressions.py | 35 +++++ .../quantization/test_nvfp4_spark_fence.py | 96 ++++++++++++ .../convert_minimax_h3_mlx.py | 16 +- 13 files changed, 543 insertions(+), 133 deletions(-) create mode 100644 fastvideo/tests/inference/test_benchmark_fasth3_spark_nvfp4.py create mode 100644 fastvideo/tests/ops/quantization/test_nvfp4_spark_fence.py diff --git a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py index 20d31043a7..96778cf6a2 100644 --- a/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py +++ b/examples/inference/basic/benchmark_fasth3_spark_nvfp4.py @@ -42,8 +42,11 @@ def _peak_mb(metrics: dict[str, dict], key: str) -> float | None: return max(values) if values else None -def _stage_total(stages: dict[str, float], fragment: str) -> float | None: - matches = [seconds for name, seconds in stages.items() if fragment in name.lower()] +def _stage_total(stages: dict[str, float], fragment: str, exclude: str | None = None) -> float | None: + matches = [ + seconds for name, seconds in stages.items() + if fragment in name.lower() and (exclude is None or exclude not in name.lower()) + ] return sum(matches) if matches else None @@ -124,7 +127,8 @@ def main() -> None: "height": args.height, "e2e_seconds": round(wall, 3), "denoise_seconds": _stage_total(stages, "denois"), - "decode_seconds": _stage_total(stages, "decod"), + "decode_seconds": _stage_total(stages, "decod", exclude="postdecode"), + "postprocess_seconds": _stage_total(stages, "postdecode"), "peak_memory_mb": _peak_mb(metrics, "peak_allocated_mb"), "peak_reserved_mb": _peak_mb(metrics, "peak_reserved_mb"), "result_peak_memory_mb": result.peak_memory_mb, diff --git a/examples/inference/basic/mlx_fasth3.py b/examples/inference/basic/mlx_fasth3.py index 62ec602ea3..11277a2008 100644 --- a/examples/inference/basic/mlx_fasth3.py +++ b/examples/inference/basic/mlx_fasth3.py @@ -190,34 +190,37 @@ def run(args: argparse.Namespace) -> None: taeh3_chunk_size=args.taeh3_chunk_size, prompt_cache_dir=args.prompt_cache_dir, ) - result = pipeline.generate( - args.prompt, - output_path=args.output_path, - height=args.height, - width=args.width, - num_frames=args.num_frames, - seed=args.seed, - num_steps=args.steps, - tiled_video_decode=args.tiled_video_decode, - vae_tile_height=args.vae_tile_height, - vae_tile_width=args.vae_tile_width, - inter_step_cooldown_s=args.inter_step_cooldown_s, - fast=args.fast, - fast_factor=args.fast_factor, - fast_sharpen=args.fast_sharpen, - rife_weights_dir=args.rife_weights_dir, - fast_spatial=args.fast_spatial, - fast_spatial_scale=args.fast_spatial_scale, - fast_spatial_upsample_mode=args.fast_spatial_upsample_mode, - fast_spatial_sharpen=args.fast_spatial_sharpen, - vsa=args.vsa, - vsa_sparsity=args.vsa_sparsity, - vsa_tile_size=args.vsa_tile_size, - vsa_prefix_mode=args.vsa_prefix_mode, - vsa_dense_first_n_steps=args.vsa_dense_first_n_steps, - vsa_dense_layers=args.vsa_dense_layers, - vsa_impl=args.vsa_impl, - ) + try: + result = pipeline.generate( + args.prompt, + output_path=args.output_path, + height=args.height, + width=args.width, + num_frames=args.num_frames, + seed=args.seed, + num_steps=args.steps, + tiled_video_decode=args.tiled_video_decode, + vae_tile_height=args.vae_tile_height, + vae_tile_width=args.vae_tile_width, + inter_step_cooldown_s=args.inter_step_cooldown_s, + fast=args.fast, + fast_factor=args.fast_factor, + fast_sharpen=args.fast_sharpen, + rife_weights_dir=args.rife_weights_dir, + fast_spatial=args.fast_spatial, + fast_spatial_scale=args.fast_spatial_scale, + fast_spatial_upsample_mode=args.fast_spatial_upsample_mode, + fast_spatial_sharpen=args.fast_spatial_sharpen, + vsa=args.vsa, + vsa_sparsity=args.vsa_sparsity, + vsa_tile_size=args.vsa_tile_size, + vsa_prefix_mode=args.vsa_prefix_mode, + vsa_dense_first_n_steps=args.vsa_dense_first_n_steps, + vsa_dense_layers=args.vsa_dense_layers, + vsa_impl=args.vsa_impl, + ) + finally: + pipeline.close() print(json.dumps({ "video_path": result.video_path, "timings_s": {k: round(v, 2) for k, v in result.timings.items()}, diff --git a/fastvideo/layers/fp4linear.py b/fastvideo/layers/fp4linear.py index 6b2b613681..82a2c75739 100644 --- a/fastvideo/layers/fp4linear.py +++ b/fastvideo/layers/fp4linear.py @@ -2,6 +2,8 @@ import torch +from fastvideo.layers.quantization.nvfp4_config import nvfp4_quantize_fenced + try: import flashinfer except ImportError: @@ -50,18 +52,9 @@ def forward(ctx, x, weight, bias, backend="cutlass", block_size=16, use_128x4_sf global_sf_a = _global_sf(x2d) global_sf_b = _global_sf(weight_cast) - a_fp4, a_inv_s = flashinfer_mod.nvfp4_quantize( - x2d, - global_sf_a, - sfLayout=a_sf_layout, - do_shuffle=False, - ) - b_fp4, b_inv_s = flashinfer_mod.nvfp4_quantize( - weight_cast, - global_sf_b, - sfLayout=flashinfer_mod.SfLayout.layout_128x4, - do_shuffle=False, - ) + # Shared helper: carries the DGX Spark (GB10) ordering fence. + a_fp4, a_inv_s = nvfp4_quantize_fenced(x2d, global_sf_a, a_sf_layout.value) + b_fp4, b_inv_s = nvfp4_quantize_fenced(weight_cast, global_sf_b, flashinfer_mod.SfLayout.layout_128x4.value) alpha = 1.0 / (global_sf_a * global_sf_b) diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index fcf78ea00f..17460f46f8 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -19,6 +19,7 @@ """ from __future__ import annotations +import functools import logging import os import re @@ -171,6 +172,44 @@ def _get_ltx2_fp4_stage_profile(default: str = "refine") -> str: return default +@functools.cache +def _is_dgx_spark(device_index: int) -> bool: + return torch.cuda.get_device_capability(device_index) == (12, 1) + + +def nvfp4_quantize_fenced( + x: torch.Tensor, + global_sf: torch.Tensor, + sf_layout: int, + do_shuffle: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + """FlashInfer NVFP4 quantization with the DGX Spark (GB10) ordering fence. + + Every FastVideo NVFP4 quantization goes through here, so the GB10 + workaround covers inference and the QAT straight-through linear alike. + """ + device_index = x.device.index if x.device.index is not None else torch.cuda.current_device() + spark = _is_dgx_spark(device_index) + if spark and torch.cuda.is_current_stream_capturing(): + raise RuntimeError("NVFP4 activation quantization on DGX Spark requires a completion fence; " + "disable CUDA graph capture.") + SfLayout, _, nvfp4_quantize = _require_flashinfer() + if not spark: + return nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) + # FlashInfer's PDL kernel reads the global scale before its + # dependency wait. Fresh dynamic scales require normal ordering. + quantized, scales = nvfp4_quantize(x, + global_sf, + sfLayout=SfLayout(sf_layout), + do_shuffle=do_shuffle, + enable_pdl=False) + # With FlashInfer 0.6.18 on GB10, queued activation quantization + # plus GEMM can diverge. Completing quantization while its padded + # input is alive prevents the observed intermittent corruption. + torch.cuda.current_stream(x.device).synchronize() + return quantized, scales + + _OPS_REGISTERED = False @@ -193,27 +232,7 @@ def _nvfp4_quantize_op( sf_layout: int, do_shuffle: bool = False, ) -> tuple[torch.Tensor, torch.Tensor]: - spark = torch.cuda.get_device_capability(x.device) == (12, 1) - if spark and torch.cuda.is_current_stream_capturing(): - raise RuntimeError("NVFP4 activation quantization on DGX Spark requires a completion fence; " - "disable CUDA graph capture.") - SfLayout, _, nvfp4_quantize = _require_flashinfer() - if spark: - # FlashInfer's PDL kernel reads the global scale before its - # dependency wait. Fresh dynamic scales require normal ordering. - quantized, scales = nvfp4_quantize(x, - global_sf, - sfLayout=SfLayout(sf_layout), - do_shuffle=do_shuffle, - enable_pdl=False) - else: - quantized, scales = nvfp4_quantize(x, global_sf, sfLayout=SfLayout(sf_layout), do_shuffle=do_shuffle) - if spark: - # With FlashInfer 0.6.18 on GB10, queued activation quantization - # plus GEMM can diverge. Completing quantization while its padded - # input is alive prevents the observed intermittent corruption. - torch.cuda.current_stream(x.device).synchronize() - return quantized, scales + return nvfp4_quantize_fenced(x, global_sf, sf_layout, do_shuffle) @_nvfp4_quantize_op.register_fake def _nvfp4_quantize_op_fake( diff --git a/fastvideo/mlx_runtime/minimax_h3_conditioner.py b/fastvideo/mlx_runtime/minimax_h3_conditioner.py index 3acf989e29..12dcd8554b 100644 --- a/fastvideo/mlx_runtime/minimax_h3_conditioner.py +++ b/fastvideo/mlx_runtime/minimax_h3_conditioner.py @@ -20,6 +20,7 @@ import gc import json +import os import shutil import tempfile from dataclasses import dataclass @@ -399,9 +400,9 @@ def close(self) -> None: def unswizzle_nvfp4_scales(scale: np.ndarray, rows: int, cols: int) -> np.ndarray: """FlashInfer 128x4 scale bytes -> MLX row-major group-16 scale bytes.""" pad_rows, pad_cols = -(-rows // 128) * 128, -(-cols // 4) * 4 - if scale.size != pad_rows * pad_cols: + if scale.size < pad_rows * pad_cols: raise ValueError(f"NVFP4 scales need {pad_rows * pad_cols} bytes, got {scale.size}.") - tiles = scale.reshape(pad_rows // 128, pad_cols // 4, 32, 4, 4) + tiles = scale.reshape(-1)[:pad_rows * pad_cols].reshape(pad_rows // 128, pad_cols // 4, 32, 4, 4) return np.ascontiguousarray(tiles.transpose(0, 3, 2, 1, 4).reshape(pad_rows, pad_cols)[:rows, :cols]) @@ -410,16 +411,17 @@ def unswizzle_nvfp4_scales(scale: np.ndarray, rows: int, cols: int) -> np.ndarra def _read_nvfp4_encoder_config(component_dir: str | Path) -> dict[str, Any]: raw = json.loads((Path(component_dir) / "config.json").read_text()) - expected = { - "quant_method": "nvfp4", - "fmt": "e2m1", - "group_size": 16, - "scale_fmt": "e4m3", - "scale_layout": "128x4", - "activation_scheme": "dynamic", + # Same accepted spellings as the CUDA loader (MiniMaxH3SerializedNVFP4Config.from_config). + accepted = { + "quant_method": ("nvfp4", ), + "fmt": ("e2m1", "float4_e2m1fn", "nvfp4"), + "scale_fmt": ("e4m3", "float8_e4m3fn"), + "activation_scheme": ("dynamic", ), } quant = raw.get("quantization_config", {}) - if any(quant.get(key) != value for key, value in expected.items()): + group_size = quant.get("group_size") + if (any(str(quant.get(key, "")).lower() not in values for key, values in accepted.items()) + or isinstance(group_size, bool) or group_size != 16 or quant.get("scale_layout") != "128x4"): raise ValueError("MLX NVFP4 conditioning requires the FastVideo group-16, 128x4 encoder export.") return raw @@ -462,6 +464,9 @@ def export_mlx_h3_nvfp4_encoder(component_dir: str | Path, output_dir: str | Pat "source_dir": str(Path(component_dir).resolve()) } (output_dir / MLX_NVFP4_ENCODER_MANIFEST).write_text(json.dumps(manifest, indent=2) + "\n") + umask = os.umask(0) + os.umask(umask) + output_dir.chmod(0o777 & ~umask) # mkdtemp is owner-only; shared model dirs need the umask mode if final_dir.exists(): final_dir.rmdir() # empty, checked above output_dir.rename(final_dir) diff --git a/fastvideo/mlx_runtime/minimax_h3_pipeline.py b/fastvideo/mlx_runtime/minimax_h3_pipeline.py index 02708e2d1e..0fc805d137 100644 --- a/fastvideo/mlx_runtime/minimax_h3_pipeline.py +++ b/fastvideo/mlx_runtime/minimax_h3_pipeline.py @@ -65,6 +65,7 @@ video_latent_num_frames, ) from fastvideo.mlx_runtime.minimax_h3_vsa import MiniMaxH3VSAConfig +from fastvideo.mlx_runtime.prompt_cache import fingerprint_digest, text_encoder_fingerprint logger = init_logger(__name__) @@ -263,8 +264,10 @@ def _configure_metal_memory_limits(mx, if set_memory is None and hasattr(mx, "metal"): set_memory = getattr(mx.metal, "set_memory_limit", None) if set_memory is not None and not resident: + # An explicit wired request must fit under the allocator cap, or it pins memory MLX cannot allocate. + memory_limit_gib = max(_default_metal_wired_limit_gib(mx), wired_limit_gib or 0.0) try: - previous_memory = int(set_memory(int(_default_metal_wired_limit_gib(mx) * 2**30))) + previous_memory = int(set_memory(int(memory_limit_gib * 2**30))) except Exception as error: # noqa: BLE001 - older MLX best effort logger.info("Could not set the Metal allocation limit: %s", error) if wired_limit_gib is None: @@ -331,6 +334,27 @@ def _cached_adaln_timesteps(checkpoint_dir: str | Path) -> np.ndarray | None: return np.asarray(cache_info["timesteps"], dtype=np.float32) +def _load_resident_dit(checkpoint_dir: str | Path): + """Load an H3 DiT and materialize every weight before timed requests.""" + import mlx.core as mx + + dit = load_mlx_h3_checkpoint(checkpoint_dir) + for group in [dit.weights, *dit.blocks, *dit.refiner]: + for value in group.values(): + # With the converter's AdaLN cache, dropped AdaLN projection weights are None. + if value is not None: + _eval_value(value) + cache = dit._adaln_cache + if cache is not None: + mx.eval(cache.block_tables, cache.norm_out_shift, cache.norm_out_scale) + return dit + + +def _adaln_weights_dropped(dit: Any) -> bool: + key = "adaln_proj.linear.weight" + return any(key in block and block[key] is None for block in getattr(dit, "blocks", ())) + + def _validate_checkpoint_step_ladder(checkpoint_dir: str | Path, num_steps: int, *, @@ -372,6 +396,10 @@ def _preflight_media_dependencies(*, class MiniMaxH3MLXPipeline: """Text-to-video-with-audio generation through the native MLX runtime.""" + # Defaults for pipelines built without __init__ (unit tests use __new__). + resident = False + conditioner_mode = "auto" + def __init__( self, *, @@ -488,16 +516,7 @@ def prepare_resident(self) -> None: conditioner.close() raise ValueError("All-resident generation requires the packed NVFP4 text encoder.") self._resident_components["conditioner"] = conditioner - dit = load_mlx_h3_checkpoint(self.dit_checkpoint) - self._resident_components["dit"] = dit - for group in [dit.weights, *dit.blocks, *dit.refiner]: - for value in group.values(): - # With the converter's AdaLN cache, dropped AdaLN projection weights are None. - if value is not None: - _eval_value(value) - cache = dit._adaln_cache - if cache is not None: - mx.eval(cache.block_tables, cache.norm_out_shift, cache.norm_out_scale) + self._resident_components["dit"] = _load_resident_dit(self.dit_checkpoint) self._resident_components["video_vae"] = mlx_h3_video_vae_from_dir(self.model_root / "vae", include_encoder=False, storage_dtype=self.vae_dtype) @@ -531,24 +550,19 @@ def close(self) -> None: def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: """Returns (hidden states (S, hidden), token tags). Uses the cache or the streamed conditioner.""" - cache_key = None - if self.prompt_cache_dir is not None: - identity = (f"{self.model_root}::conditioner=" - f"{getattr(self, 'conditioner_dir', self.model_root / 'text_encoder')}::" - f"{getattr(self, 'conditioner_mode', 'auto')}") - cache_key = prompt_cache_path(self.prompt_cache_dir, identity, prompt) - if cache_key.exists(): - data = np.load(cache_key) - logger.info("Loaded prompt embeddings from cache %s", cache_key) - return data["hidden_states"], data["token_tags"] - - if getattr(self, "resident", False): + cache_key = self._prompt_cache_key(prompt) + if cache_key is not None and cache_key.exists(): + data = np.load(cache_key) + logger.info("Loaded prompt embeddings from cache %s", cache_key) + return data["hidden_states"], data["token_tags"] + + if self.resident: self.prepare_resident() conditioner = self._resident_components["conditioner"] else: conditioner = self._load_conditioner() hidden, tags = conditioner.encode_prompt(prompt) - if not getattr(self, "resident", False): + if not self.resident: conditioner.close() _cleanup_mlx() if cache_key is not None: @@ -562,6 +576,19 @@ def encode_prompt(self, prompt: str) -> tuple[np.ndarray, np.ndarray]: tmp_cache.unlink(missing_ok=True) return hidden, tags + def _prompt_cache_key(self, prompt: str) -> Path | None: + """Cache path bound to the effective encoder and its files; None skips the cache.""" + if self.prompt_cache_dir is None: + return None + fingerprint = { + "conditioner": self._conditioner_kind(), + "text_encoder": text_encoder_fingerprint(self.conditioner_dir), + "tokenizer": text_encoder_fingerprint(self.tokenizer_dir), + } + if not (fingerprint["text_encoder"]["complete"] and fingerprint["tokenizer"]["complete"]): + return None + return prompt_cache_path(self.prompt_cache_dir, fingerprint_digest(fingerprint), prompt) + def load_prompt_cache(self, path: str | Path) -> tuple[np.ndarray, np.ndarray]: data = np.load(path) return data["hidden_states"], data["token_tags"] @@ -577,13 +604,19 @@ def _load_conditioner(self): StreamedMiniMaxH3TextConditioner, ) + if self._conditioner_kind() == "nvfp4": + return ResidentNVFP4MiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + return StreamedMiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + + def _conditioner_kind(self) -> str: + """The encoder that conditioner_mode selects for these weights: 'nvfp4' or 'streamed'.""" config = json.loads((self.conditioner_dir / "config.json").read_text()) - packed = config.get("quantization_config", {}).get("quant_method") == "nvfp4" + packed = str(config.get("quantization_config", {}).get("quant_method", "")).lower() == "nvfp4" if self.conditioner_mode == "nvfp4" or (self.conditioner_mode == "auto" and packed): - return ResidentNVFP4MiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + return "nvfp4" if packed: raise ValueError("The streamed conditioner requires BF16 weights; use conditioner_mode='nvfp4'.") - return StreamedMiniMaxH3TextConditioner(self.conditioner_dir, self.tokenizer_dir) + return "streamed" # -- phase 2: denoise -------------------------------------------------- @@ -609,7 +642,8 @@ def denoise( geometry = self.resolve_geometry(height, width, num_frames, enforce_duration=audio_num_frames is None) audio_frames = geometry["num_frames"] if audio_num_frames is None else align_num_frames(audio_num_frames) - if dit is None and getattr(self, "resident", False): + resident_dit = dit is None and self.resident + if resident_dit: _validate_checkpoint_step_ladder(self.dit_checkpoint, num_steps, model_root=self.model_root) self.prepare_resident() dit = self._resident_components["dit"] @@ -619,6 +653,35 @@ def denoise( t0 = time.perf_counter() dit = load_mlx_h3_checkpoint(self.dit_checkpoint) logger.info("Loaded MLX H3 DiT from %s in %.1fs", self.dit_checkpoint, time.perf_counter() - t0) + + schedule = resolve_h3_denoise_schedule(self.model_root, num_steps) + video_scheduler = schedule.video + audio_scheduler = schedule.audio + union = schedule.adaln_timesteps + num_steps = schedule.num_steps + # The keyframe-noise timestep (0.999) is only exercised by FL2VA/Ref2VA + # conditioning rows; those modes recompute the ladder before denoise. + + cache = getattr(dit, "_adaln_cache", None) + if cache is None or not np.array_equal(cache.timesteps.astype(np.float32), union): + if cache is not None: + extra = np.setdiff1d(union, cache.timesteps) + logger.info("Recomputing AdaLN cache for %d-step ladder (extra timesteps %s).", num_steps, extra) + if _adaln_weights_dropped(dit): + # An earlier request on this DiT already released the AdaLN projections. + if not resident_dit: + raise ValueError("This H3 DiT dropped its AdaLN weights for another step ladder; " + "pass a freshly loaded DiT for a different num_steps.") + logger.info("Reloading the resident H3 DiT for the %d-step ladder.", num_steps) + cache = dit = None + self._resident_components.pop("dit", None) + _cleanup_mlx() + dit = self._resident_components["dit"] = _load_resident_dit(self.dit_checkpoint) + dit.precompute_adaln(union, drop_weights=True) + + if resident_dit: + # The resident DiT keeps the previous request's VSA mode unless reset to the default. + vsa_config = vsa_config or MiniMaxH3VSAConfig() if vsa_config is not None: dit.configure_vsa(vsa_config) if hasattr(dit, "reset_vsa_stats"): @@ -637,22 +700,6 @@ def denoise( if getattr(dit, "vsa_config", None) is not None and dit.vsa_config.enabled: dit.prepare_vsa_geometry(layout) - schedule = resolve_h3_denoise_schedule(self.model_root, num_steps) - video_scheduler = schedule.video - audio_scheduler = schedule.audio - union = schedule.adaln_timesteps - num_steps = schedule.num_steps - # The keyframe-noise timestep (0.999) is only exercised by FL2VA/Ref2VA - # conditioning rows; those modes recompute the ladder before denoise. - - cache = getattr(dit, "_adaln_cache", None) - if cache is None: - dit.precompute_adaln(union, drop_weights=True) - elif not np.array_equal(cache.timesteps.astype(np.float32), union): - extra = np.setdiff1d(union, cache.timesteps) - logger.info("Recomputing AdaLN cache for %d-step ladder (extra timesteps %s).", num_steps, extra) - dit.precompute_adaln(union, drop_weights=True) - video_key, audio_key = mx.random.split(mx.random.key(seed)) target_video_rows = int(layout.video_indices.shape[0] - layout.num_condition_video_rows) target_audio_rows = int(layout.audio_indices.shape[0] - layout.num_condition_audio_rows) @@ -759,7 +806,7 @@ def decode_video(self, raise RuntimeError(f"TAEH3 produced unexpected frame shape: {frames.shape}") _cleanup_mlx() return frames - if getattr(self, "resident", False): + if self.resident: self.prepare_resident() vae = self._resident_components["video_vae"] else: @@ -807,7 +854,7 @@ def decode_audio(self, audio_rows: np.ndarray, *, num_frames: int) -> np.ndarray num_audio_latents = audio_latent_num_frames(align_num_frames(num_frames)) latents = unpack_audio_tokens(audio_rows, num_audio_latents) - if getattr(self, "resident", False): + if self.resident: self.prepare_resident() vae = self._resident_components["audio_vae"] else: diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa.py b/fastvideo/mlx_runtime/minimax_h3_vsa.py index 48ca15fa32..126e6f1ac0 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa.py @@ -29,7 +29,7 @@ import math from dataclasses import dataclass, field -from functools import cached_property +from functools import cached_property, lru_cache from numbers import Integral from typing import Any, Literal @@ -544,13 +544,26 @@ def _key_valid_mask(block_idx, variable_block_sizes, tile_elems: int): return offsets[None, None, None, :] < selected_sizes[:, :, :, None] -_REFERENCE_GATHER_TARGET_BYTES = 256 * 1024**2 +_REFERENCE_GATHER_MIN_BYTES = 256 * 1024**2 +_REFERENCE_GATHER_MAX_BYTES = 2 * 1024**3 + + +@lru_cache(maxsize=1) +def _reference_gather_target_bytes() -> int: + """1/128 of unified memory in 256 MiB-2 GiB: 32 GB Macs keep room for resident weights.""" + import mlx.core as mx + + try: + total = int(mx.metal.device_info().get("memory_size", 0)) + except (AttributeError, RuntimeError, TypeError, ValueError): + total = 0 + return min(_REFERENCE_GATHER_MAX_BYTES, max(_REFERENCE_GATHER_MIN_BYTES, total // 128)) def _reference_gather_query_chunk(heads: int, dim: int, k_sel: int, tile_elems: int, n_q: int) -> int: - """Bound gathered BF16 K/V to 256 MiB, leaving space for resident weights.""" + """Bound gathered BF16 K/V by the device's gather budget.""" bytes_per_query = 4 * heads * max(k_sel, 1) * tile_elems * dim - chunk = min(n_q, max(1, _REFERENCE_GATHER_TARGET_BYTES // max(bytes_per_query, 1))) + chunk = min(n_q, max(1, _reference_gather_target_bytes() // max(bytes_per_query, 1))) return int(chunk) diff --git a/fastvideo/tests/inference/test_benchmark_fasth3_spark_nvfp4.py b/fastvideo/tests/inference/test_benchmark_fasth3_spark_nvfp4.py new file mode 100644 index 0000000000..5fd58c0564 --- /dev/null +++ b/fastvideo/tests/inference/test_benchmark_fasth3_spark_nvfp4.py @@ -0,0 +1,29 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU checks of examples/inference/basic/benchmark_fasth3_spark_nvfp4.py stage accounting.""" +from __future__ import annotations + +import importlib.util +from pathlib import Path + +REPO_ROOT = Path(__file__).resolve().parents[3] +EXAMPLE_PATH = REPO_ROOT / "examples" / "inference" / "basic" / "benchmark_fasth3_spark_nvfp4.py" + + +def _load_benchmark(): + spec = importlib.util.spec_from_file_location("benchmark_fasth3_spark_nvfp4", EXAMPLE_PATH) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_decode_seconds_exclude_post_decode_frame_processing() -> None: + benchmark = _load_benchmark() + stages = { + "MiniMaxH3DenoisingStage": 9.0, + "MiniMaxH3VideoDecodingStage": 2.0, + "MiniMaxH3AudioDecodingStage": 0.5, + "PostDecodeFrameProcessStage": 1.0, + } + assert benchmark._stage_total(stages, "decod", exclude="postdecode") == 2.5 + assert benchmark._stage_total(stages, "postdecode") == 1.0 + assert benchmark._stage_total(stages, "denois") == 9.0 diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_conditioner.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_conditioner.py index 4fb261ce7d..a94d3b7a6a 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_conditioner.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_conditioner.py @@ -4,6 +4,7 @@ from __future__ import annotations import inspect +import json import numpy as np import pytest @@ -133,6 +134,33 @@ def test_prompt_cache_is_versioned_for_conditioner_layout() -> None: assert current.name != f"prompt_embeds_{legacy_digest}.npz" +def _cache_pipeline(tmp_path, quant_method=None): + encoder = tmp_path / "model" / "text_encoder" + tokenizer = tmp_path / "model" / "tokenizer" + encoder.mkdir(parents=True, exist_ok=True) + tokenizer.mkdir(parents=True, exist_ok=True) + config = {"quantization_config": {"quant_method": quant_method}} if quant_method else {} + (encoder / "config.json").write_text(json.dumps(config)) + (encoder / "model.safetensors").write_bytes(b"weights") + (tokenizer / "tokenizer.json").write_text("{}") + pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) + pipeline.prompt_cache_dir = tmp_path / "cache" + pipeline.model_root = tmp_path / "model" + pipeline.conditioner_dir = encoder + pipeline.tokenizer_dir = tokenizer + return pipeline + + +def test_prompt_cache_key_tracks_effective_encoder_and_files(tmp_path) -> None: + pipeline = _cache_pipeline(tmp_path, quant_method="nvfp4") + auto_key = pipeline._prompt_cache_key("prompt") + pipeline.conditioner_mode = "nvfp4" + assert pipeline._prompt_cache_key("prompt") == auto_key + (pipeline.conditioner_dir / "model.safetensors").write_bytes(b"replaced weights") + assert pipeline._prompt_cache_key("prompt") != auto_key + assert pipeline._prompt_cache_key("other prompt") != pipeline._prompt_cache_key("prompt") + + def test_prompt_cache_write_is_atomic(tmp_path, monkeypatch) -> None: class FakeConditioner: def encode_prompt(self, _prompt): @@ -141,9 +169,7 @@ def encode_prompt(self, _prompt): def close(self): return None - pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) - pipeline.prompt_cache_dir = tmp_path - pipeline.model_root = tmp_path / "model" + pipeline = _cache_pipeline(tmp_path) monkeypatch.setattr(MiniMaxH3MLXPipeline, "_load_conditioner", lambda _self: FakeConditioner()) def interrupted_save(handle, **_arrays): @@ -151,7 +177,7 @@ def interrupted_save(handle, **_arrays): raise RuntimeError("interrupted") monkeypatch.setattr(np, "savez", interrupted_save) - cache_path = prompt_cache_path(tmp_path, pipeline.model_root, "prompt") + cache_path = pipeline._prompt_cache_key("prompt") with pytest.raises(RuntimeError, match="interrupted"): pipeline.encode_prompt("prompt") @@ -173,3 +199,51 @@ def test_video_vae_decode_defaults_to_reference_spatial_tiles() -> None: def test_h3_audio_duration_rounds_up_to_cover_last_video_packet() -> None: assert _audio_sample_count(124) == 165334 + + +@pytest.mark.parametrize("fmt,scale_fmt,activation", [("e2m1", "e4m3", "dynamic"), + ("nvfp4", "float8_e4m3fn", "Dynamic"), + ("FLOAT4_E2M1FN", "E4M3", "dynamic")]) +def test_nvfp4_encoder_config_accepts_cuda_loader_spellings(tmp_path, fmt, scale_fmt, activation) -> None: + from fastvideo.mlx_runtime.minimax_h3_conditioner import _read_nvfp4_encoder_config + + quant = {"quant_method": "NVFP4", "fmt": fmt, "group_size": 16, "scale_fmt": scale_fmt, + "scale_layout": "128x4", "activation_scheme": activation} + (tmp_path / "config.json").write_text(json.dumps({"quantization_config": quant})) + assert _read_nvfp4_encoder_config(tmp_path)["quantization_config"] == quant + (tmp_path / "config.json").write_text(json.dumps({"quantization_config": {**quant, "group_size": 32}})) + with pytest.raises(ValueError, match="group-16"): + _read_nvfp4_encoder_config(tmp_path) + + +def test_unswizzle_nvfp4_scales_ignores_trailing_padding() -> None: + from fastvideo.mlx_runtime.minimax_h3_conditioner import unswizzle_nvfp4_scales + + exact = np.arange(128 * 4, dtype=np.uint8) + padded = np.concatenate([exact, np.full(64, 255, dtype=np.uint8)]) + np.testing.assert_array_equal(unswizzle_nvfp4_scales(padded, 100, 3), unswizzle_nvfp4_scales(exact, 100, 3)) + with pytest.raises(ValueError, match="need 512 bytes"): + unswizzle_nvfp4_scales(exact[:-1], 100, 3) + + +def test_encoder_cache_is_published_with_umask_permissions(tmp_path, monkeypatch) -> None: + import os + + from fastvideo.mlx_runtime import minimax_h3_conditioner as conditioner + + class FakeIndex: + def __init__(self, _source): + self.weights = {"model.norm.weight": mx.ones((4, ))} + + def close(self): + pass + + monkeypatch.setattr(conditioner, "_read_nvfp4_encoder_config", lambda _dir: {}) + monkeypatch.setattr(conditioner, "_ShardIndex", lambda path: path) + monkeypatch.setattr(conditioner, "_ResidentNVFP4Index", FakeIndex) + old_umask = os.umask(0o022) + try: + out = conditioner.export_mlx_h3_nvfp4_encoder(tmp_path / "src", tmp_path / "cache") + finally: + os.umask(old_umask) + assert out.stat().st_mode & 0o777 == 0o755 diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py index 1682372092..62e91df61c 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_fast_mode.py @@ -30,6 +30,7 @@ plan_fast_temporal, ) from fastvideo.mlx_runtime import rife_interp # noqa: E402 +from fastvideo.mlx_runtime.minimax_h3_vsa import MiniMaxH3VSAConfig # noqa: E402 def test_fast_plan_keeps_full_audio_and_reduces_only_video() -> None: @@ -212,3 +213,88 @@ def test_invalid_pipeline_options_do_not_change_process_memory_limits(monkeypatc MiniMaxH3MLXPipeline(model_root="missing", mlx_dit_checkpoint="missing", conditioner_mode="invalid", metal_wired_limit_gib=27.0) assert calls == [] + + +def test_explicit_wired_limit_raises_the_allocator_cap(): + calls = [] + fake = SimpleNamespace( + set_memory_limit=lambda size: calls.append(("allocator", size)) or 40 * 2**30, + set_wired_limit=lambda size: calls.append(("wired", size)) or 0, + ) + _configure_metal_memory_limits(fake, 48.0) + assert calls == [("allocator", 48 * 2**30), ("wired", 48 * 2**30)] + + +class _StopAfterPrologue(Exception): + pass + + +class _FakeAdaLNDiT: + patch_size = (1, 2, 2) + + def __init__(self): + self.blocks = [{"adaln_proj.linear.weight": np.ones(1)}] + self._adaln_cache = None + self.vsa_config = MiniMaxH3VSAConfig() + + def precompute_adaln(self, timesteps, *, drop_weights=True): + assert all(block["adaln_proj.linear.weight"] is not None for block in self.blocks) + self._adaln_cache = SimpleNamespace(timesteps=np.asarray(timesteps, dtype=np.float32)) + if drop_weights: + for block in self.blocks: + block["adaln_proj.linear.weight"] = None + + def configure_vsa(self, config): + self.vsa_config = config + + def reset_vsa_stats(self): + pass + + +def _denoise_prologue(pipeline, monkeypatch, **kwargs): + import fastvideo.mlx_runtime.minimax_h3_pipeline as pipeline_mod + + def stop(*_args, **_kwargs): + raise _StopAfterPrologue + + monkeypatch.setattr(pipeline_mod, "build_packed_layout", stop) + with pytest.raises(_StopAfterPrologue): + pipeline.denoise(np.zeros((2, 4), dtype=np.float32), np.zeros(2, dtype=np.int64), height=256, width=256, + num_frames=124, audio_num_frames=124, seed=0, **kwargs) + + +def test_resident_dit_reloads_for_a_new_step_ladder_and_resets_vsa(tmp_path, monkeypatch) -> None: + import fastvideo.mlx_runtime.minimax_h3_pipeline as pipeline_mod + + first, reloaded = _FakeAdaLNDiT(), _FakeAdaLNDiT() + monkeypatch.setattr(pipeline_mod, "_validate_checkpoint_step_ladder", lambda *a, **k: None) + monkeypatch.setattr(pipeline_mod, "_load_resident_dit", lambda _path: reloaded) + pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) + pipeline.resident = True + pipeline.model_root = tmp_path + pipeline.dit_checkpoint = tmp_path + pipeline._resident_components = {"dit": first} + + sparse = MiniMaxH3VSAConfig(enabled=True, sparsity=0.8) + _denoise_prologue(pipeline, monkeypatch, num_steps=4, vsa_config=sparse) + assert first.vsa_config == sparse + _denoise_prologue(pipeline, monkeypatch, num_steps=4) + assert pipeline._resident_components["dit"] is first + assert first.vsa_config == MiniMaxH3VSAConfig() + + _denoise_prologue(pipeline, monkeypatch, num_steps=8) + assert pipeline._resident_components["dit"] is reloaded + np.testing.assert_array_equal(reloaded._adaln_cache.timesteps, _adaln_schedule_union(8)) + + +def test_caller_dit_with_dropped_adaln_weights_rejects_a_new_ladder(tmp_path, monkeypatch) -> None: + import fastvideo.mlx_runtime.minimax_h3_pipeline as pipeline_mod + + monkeypatch.setattr(pipeline_mod, "_validate_checkpoint_step_ladder", lambda *a, **k: None) + pipeline = MiniMaxH3MLXPipeline.__new__(MiniMaxH3MLXPipeline) + pipeline.model_root = tmp_path + dit = _FakeAdaLNDiT() + _denoise_prologue(pipeline, monkeypatch, num_steps=4, dit=dit) + with pytest.raises(ValueError, match="dropped its AdaLN weights"): + pipeline.denoise(np.zeros((2, 4), dtype=np.float32), np.zeros(2, dtype=np.int64), height=256, width=256, + num_frames=124, audio_num_frames=124, seed=0, num_steps=8, dit=dit) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index 717ecb95c4..8f1360cdde 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -350,3 +350,38 @@ def test_converter_continues_past_mismatched_existing_format(tmp_path, monkeypat converter.main() assert saved == ["int6"] assert (existing / h3.H3_WEIGHTS_FILENAME).read_bytes() == b"existing" + + +def test_converter_rerun_skips_published_encoder_cache(tmp_path, monkeypatch): + path = Path(__file__).resolve().parents[3] / "scripts/checkpoint_conversion/convert_minimax_h3_mlx.py" + spec = importlib.util.spec_from_file_location("h3_converter_encoder_rerun", path) + converter = importlib.util.module_from_spec(spec) + spec.loader.exec_module(converter) + from fastvideo.mlx_runtime import minimax_h3_conditioner as conditioner + + existing = tmp_path / "int8" + existing.mkdir() + (existing / h3.H3_MANIFEST_FILENAME).write_text(json.dumps({"vsa": {"capable": False}})) + (existing / h3.H3_WEIGHTS_FILENAME).write_bytes(b"existing") + encoder = tmp_path / "nvfp4-encoder" + encoder.mkdir() + (encoder / conditioner.MLX_NVFP4_ENCODER_MANIFEST).write_text("{}") + (encoder / "model.safetensors").write_bytes(b"cached") + monkeypatch.setattr(converter, "parse_args", lambda: argparse.Namespace( + formats="int8", out=tmp_path, model_root="unused", include_vsa=False, + nvfp4_conditioner_root=tmp_path / "packed", nvfp4_conditioner_out=None)) + exported = [] + monkeypatch.setattr(conditioner, "export_mlx_h3_nvfp4_encoder", lambda *args: exported.append(args)) + converter.main() + assert exported == [] + assert (encoder / "model.safetensors").read_bytes() == b"cached" + + +@pytest.mark.parametrize("memory_gib,budget_mib", [(16, 256), (32, 256), (64, 512), (128, 1024), (512, 2048)]) +def test_reference_gather_budget_scales_with_unified_memory(monkeypatch, memory_gib, budget_mib): + monkeypatch.setattr(mx.metal, "device_info", lambda: {"memory_size": memory_gib * 2**30}) + vsa._reference_gather_target_bytes.cache_clear() + try: + assert vsa._reference_gather_target_bytes() == budget_mib * 2**20 + finally: + vsa._reference_gather_target_bytes.cache_clear() diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_spark_fence.py b/fastvideo/tests/ops/quantization/test_nvfp4_spark_fence.py new file mode 100644 index 0000000000..fda1aedc82 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_spark_fence.py @@ -0,0 +1,96 @@ +# SPDX-License-Identifier: Apache-2.0 +"""The DGX Spark (GB10) NVFP4 quantization fence covers inference and the QAT linear. + +GB10 is emulated by patching ``_is_dgx_spark``; any Blackwell GPU runs the kernels. +""" +from __future__ import annotations + +import pytest +import torch + +flashinfer = pytest.importorskip("flashinfer") + +from fastvideo.layers import fp4linear # noqa: E402 +from fastvideo.layers.quantization import nvfp4_config as nv # noqa: E402 + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 10, + reason="NVFP4 kernels need a Blackwell GPU") + +LAYOUT_128X4 = flashinfer.SfLayout.layout_128x4 + + +def _activation(): + torch.manual_seed(0) + x = torch.randn(200, 256, device="cuda", dtype=torch.bfloat16) + return x, (448.0 * 6.0) / x.float().abs().amax() + + +def _pad_rows(x: torch.Tensor) -> torch.Tensor: + return torch.nn.functional.pad(x, (0, 0, 0, (-x.shape[0]) % 128)) + + +def _spy_quantize(monkeypatch): + real = flashinfer.nvfp4_quantize + pdl = [] + + def spy(*args, **kwargs): + pdl.append(kwargs.get("enable_pdl")) + return real(*args, **kwargs) + + monkeypatch.setattr(flashinfer, "nvfp4_quantize", spy) + return real, pdl + + +def _spy_sync(monkeypatch): + real = torch.cuda.Stream.synchronize + syncs = [] + monkeypatch.setattr(torch.cuda.Stream, "synchronize", lambda self: syncs.append(self) or real(self)) + return syncs + + +def test_fenced_quantize_matches_flashinfer_without_a_fence_off_spark(monkeypatch) -> None: + real, pdl = _spy_quantize(monkeypatch) + syncs = _spy_sync(monkeypatch) + x, global_sf = _activation() + quantized, scales = nv.nvfp4_quantize_fenced(x, global_sf, LAYOUT_128X4.value) + expected = real(x, global_sf, sfLayout=LAYOUT_128X4, do_shuffle=False) + assert pdl == [None] and syncs == [] + assert torch.equal(quantized, expected[0]) and torch.equal(scales, expected[1]) + + +def test_spark_inference_op_disables_pdl_and_fences(monkeypatch) -> None: + real, pdl = _spy_quantize(monkeypatch) + syncs = _spy_sync(monkeypatch) + monkeypatch.setattr(nv, "_is_dgx_spark", lambda _index: True) + x, global_sf = _activation() + quantized, scales = nv._nvfp4_quantize(x, global_sf, sfLayout=LAYOUT_128X4) + expected = real(_pad_rows(x), global_sf, sfLayout=LAYOUT_128X4, do_shuffle=False) + assert pdl == [False] and len(syncs) == 1 + assert torch.equal(quantized, expected[0][:x.shape[0]]) and torch.equal(scales, expected[1]) + + +def test_spark_rejects_cuda_graph_capture(monkeypatch) -> None: + monkeypatch.setattr(nv, "_is_dgx_spark", lambda _index: True) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + x, global_sf = _activation() + with pytest.raises(RuntimeError, match="completion fence"): + nv.nvfp4_quantize_fenced(x, global_sf, LAYOUT_128X4.value) + + +def test_qat_linear_quantizes_through_the_fence_with_unchanged_numerics(monkeypatch) -> None: + torch.manual_seed(1) + x = torch.randn(4, 64, 256, device="cuda", dtype=torch.bfloat16) + weight = torch.randn(512, 256, device="cuda", dtype=torch.bfloat16) / 16 + bias = torch.randn(512, device="cuda", dtype=torch.bfloat16) + off_spark = fp4linear._LinearFWD4BWD16Fn.apply(x, weight, bias, "cutlass", 16, True) + + _, pdl = _spy_quantize(monkeypatch) + syncs = _spy_sync(monkeypatch) + monkeypatch.setattr(nv, "_is_dgx_spark", lambda _index: True) + on_spark = fp4linear._LinearFWD4BWD16Fn.apply(x, weight, bias, "cutlass", 16, True) + assert pdl == [False, False] and len(syncs) == 2 + assert torch.equal(on_spark, off_spark) + + reference = torch.nn.functional.linear(x.float(), weight.float(), bias.float()) + relative = (on_spark.float() - reference).norm() / reference.norm() + assert relative < 0.15 diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py index 931da4aa23..688e2f3267 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_mlx.py @@ -149,13 +149,19 @@ def main() -> None: mx.clear_cache() if args.nvfp4_conditioner_root is not None: - from fastvideo.mlx_runtime.minimax_h3_conditioner import export_mlx_h3_nvfp4_encoder + from fastvideo.mlx_runtime.minimax_h3_conditioner import ( + MLX_NVFP4_ENCODER_MANIFEST, + export_mlx_h3_nvfp4_encoder, + ) encoder_out = args.nvfp4_conditioner_out or out_base / "nvfp4-encoder" - started = time.perf_counter() - export_mlx_h3_nvfp4_encoder(args.nvfp4_conditioner_root, encoder_out) - print(f"[encoder] cached packed NVFP4 encoder in {time.perf_counter() - started:.1f}s at {encoder_out}", - flush=True) + if (encoder_out / MLX_NVFP4_ENCODER_MANIFEST).exists() and (encoder_out / "model.safetensors").exists(): + print(f"[skip] NVFP4 encoder already cached at {encoder_out}", flush=True) + else: + started = time.perf_counter() + export_mlx_h3_nvfp4_encoder(args.nvfp4_conditioner_root, encoder_out) + print(f"[encoder] cached packed NVFP4 encoder in {time.perf_counter() - started:.1f}s at {encoder_out}", + flush=True) if __name__ == "__main__": From f06681a31990c950bd7ee14e9aae4110982fb15d Mon Sep 17 00:00:00 2001 From: SolitaryThinker Date: Tue, 6 Oct 2026 18:56:12 +0000 Subject: [PATCH 106/107] [misc]: read device memory through mx.device_info when available Co-Authored-By: Claude Opus 5.5 --- fastvideo/mlx_runtime/minimax_h3_vsa.py | 3 ++- fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py | 4 +++- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/fastvideo/mlx_runtime/minimax_h3_vsa.py b/fastvideo/mlx_runtime/minimax_h3_vsa.py index 126e6f1ac0..833bd139f0 100644 --- a/fastvideo/mlx_runtime/minimax_h3_vsa.py +++ b/fastvideo/mlx_runtime/minimax_h3_vsa.py @@ -553,8 +553,9 @@ def _reference_gather_target_bytes() -> int: """1/128 of unified memory in 256 MiB-2 GiB: 32 GB Macs keep room for resident weights.""" import mlx.core as mx + device_info = getattr(mx, "device_info", None) or mx.metal.device_info try: - total = int(mx.metal.device_info().get("memory_size", 0)) + total = int(device_info().get("memory_size", 0)) except (AttributeError, RuntimeError, TypeError, ValueError): total = 0 return min(_REFERENCE_GATHER_MAX_BYTES, max(_REFERENCE_GATHER_MIN_BYTES, total // 128)) diff --git a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py index 8f1360cdde..433b007b7e 100644 --- a/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py +++ b/fastvideo/tests/mlx/test_mlx_minimax_h3_vsa_regressions.py @@ -379,7 +379,9 @@ def test_converter_rerun_skips_published_encoder_cache(tmp_path, monkeypatch): @pytest.mark.parametrize("memory_gib,budget_mib", [(16, 256), (32, 256), (64, 512), (128, 1024), (512, 2048)]) def test_reference_gather_budget_scales_with_unified_memory(monkeypatch, memory_gib, budget_mib): - monkeypatch.setattr(mx.metal, "device_info", lambda: {"memory_size": memory_gib * 2**30}) + info = lambda: {"memory_size": memory_gib * 2**30} # noqa: E731 + monkeypatch.setattr(mx, "device_info", info, raising=False) + monkeypatch.setattr(mx.metal, "device_info", info) vsa._reference_gather_target_bytes.cache_clear() try: assert vsa._reference_gather_target_bytes() == budget_mib * 2**20 From 31702474eefba3d0006f1803a7655a6a3dee88b5 Mon Sep 17 00:00:00 2001 From: SolitaryThinker Date: Tue, 6 Oct 2026 19:10:40 +0000 Subject: [PATCH 107/107] [misc]: patch the gather budget function in the gather-budget test Co-Authored-By: Claude Opus 5.5 --- .../mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py index 376478624a..b35bd59426 100644 --- a/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py +++ b/fastvideo/mlx_runtime/tests/test_minimax_h3_vsa_gather_budget.py @@ -15,10 +15,10 @@ def test_gather_budget_preserves_attention_with_partial_tiles(monkeypatch): # Reverse video order also exercises the selected-key order, not a dense mask. selected = np.array([0, 1, geom.num_tiles - 1, geom.num_prefix_tiles], dtype=np.int32) idx = mx.array(np.broadcast_to(selected, (2, geom.num_video_tiles, selected.size)).copy()) - monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 2 * 1024**3) + monkeypatch.setattr(vsa, '_reference_gather_target_bytes', lambda: 2 * 1024**3) expected = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) mx.eval(expected) - monkeypatch.setattr(vsa, '_REFERENCE_GATHER_TARGET_BYTES', 1) + monkeypatch.setattr(vsa, '_reference_gather_target_bytes', lambda: 1) actual = vsa._reference_gather_sdpa(q, k, value, idx, geom, 128**-0.5) mx.eval(actual) np.testing.assert_array_equal(np.array(actual.astype(mx.float32)), np.array(expected.astype(mx.float32)))