From 2d0ebf1eefa2d09fb668aff658382690af75ef1e Mon Sep 17 00:00:00 2001 From: Aryan Kumar Date: Sat, 19 Sep 2026 14:25:03 -0700 Subject: [PATCH 01/68] [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 02/68] [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 03/68] [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 04/68] [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 05/68] [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 06/68] [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 07/68] [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 08/68] [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 09/68] [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 10/68] [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 11/68] [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 12/68] [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 13/68] [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 14/68] [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 15/68] [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 16/68] [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 17/68] [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 18/68] [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 19/68] [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 20/68] [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 21/68] [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 22/68] [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 23/68] [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 24/68] [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 25/68] [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 26/68] [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 27/68] [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 28/68] [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 29/68] [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 30/68] [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 31/68] [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 32/68] [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 33/68] [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 34/68] [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 35/68] [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 36/68] [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 37/68] [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 38/68] [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 39/68] [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 40/68] [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 41/68] [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 42/68] [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 43/68] [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 44/68] [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 45/68] [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 46/68] [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 47/68] [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 48/68] [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 49/68] [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 50/68] [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 51/68] [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 52/68] [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 53/68] [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 54/68] [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 55/68] [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 56/68] [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 57/68] [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 58/68] [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 59/68] [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 0453c083a485c17099e2703f96b10aa0fed04fbf Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 12:08:42 -0700 Subject: [PATCH 60/68] [bench]: allow the launch benchmark seed to be specified --- scripts/benchmarks/minimax_h3_4090/bench_pod.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index b1f4dde625..114e649979 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -71,6 +71,7 @@ def main(): 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("--seed", type=int, default=20260929) ap.add_argument("--height", type=int, default=768) ap.add_argument("--width", type=int, default=1344) ap.add_argument("--warmup", type=int, default=1) @@ -137,7 +138,7 @@ def main(): 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, + sampling = {"seed": a.seed, "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_")) From 1d5d9c9703a7b1a6f8a7e691cc5d46aeaf9ed269 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 12:40:57 -0700 Subject: [PATCH 61/68] [offload]: honor pageable host storage for layerwise models --- fastvideo/hooks/layerwise_offload.py | 28 +++++++++++++---- .../models/encoders/minimax_h3_qwen3_vl.py | 5 ++-- fastvideo/models/loader/component_loader.py | 4 +-- .../test_minimax_h3_encoder_layerwise.py | 12 +++++--- .../tests/hooks/test_layerwise_offload.py | 30 +++++++++++++++++++ scripts/benchmarks/minimax_h3_4090/README.md | 13 ++++++++ .../benchmarks/minimax_h3_4090/bench_pod.py | 3 +- 7 files changed, 80 insertions(+), 15 deletions(-) diff --git a/fastvideo/hooks/layerwise_offload.py b/fastvideo/hooks/layerwise_offload.py index 211ea0f3e1..2d384fd31b 100644 --- a/fastvideo/hooks/layerwise_offload.py +++ b/fastvideo/hooks/layerwise_offload.py @@ -47,6 +47,7 @@ def __init__( async_copy_stream: torch.cuda.Stream, device: torch.device, next_state: "LayerwiseOffloadState | None" = None, + pin_cpu_memory: bool = True, ) -> None: self.async_copy_stream = async_copy_stream self.next_state = next_state @@ -55,6 +56,7 @@ def __init__( self.module_ref: nn.Module = None # type: ignore self.device: torch.device = device self.cpu_arena: PinnedTensorArena | None = None + self.pin_cpu_memory = pin_cpu_memory def _will_offload(self, name: str) -> bool: return True @@ -63,12 +65,18 @@ def _will_offload(self, name: str) -> bool: 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)) + if self.pin_cpu_memory: + 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): - host = self.cpu_arena.empty_like(name, param) - host.copy_(param.data.detach()) + if self.cpu_arena is not None: + host = self.cpu_arena.empty_like(name, param) + host.copy_(param.data.detach()) + else: + # Retain checkpoint-backed CPU storage so the OS can reclaim + # inactive file pages instead of holding an anonymous pinned copy. + host = param.data.detach().to("cpu") self.cpu_named_parameters[name] = host param.data = _tensor_placeholder(param.data, self.device) @@ -176,7 +184,13 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False, *, resident_blocks: int | None = None, - cyclic: bool = True): + cyclic: bool = True, + pin_cpu_memory: bool = True): + """Stream blocks, optionally retaining their existing pageable CPU storage. + + Disabling pinning avoids a private copy of file-backed checkpoint tensors. + Transfers can be slower, but inactive checkpoint pages remain reclaimable. + """ if torch.cuda.is_available(): device = torch.device("cuda", torch.cuda.current_device()) else: @@ -199,7 +213,9 @@ def enable_layerwise_offload(model: nn.Module, for idx, module_entry in enumerate(submodule): if idx < resident: continue - state = LayerwiseOffloadState(async_copy_stream=async_stream, device=device) + state = LayerwiseOffloadState(async_copy_stream=async_stream, + device=device, + pin_cpu_memory=pin_cpu_memory) state_list.append(state) hook_mgr = ModuleHookManager.get_from_or_default(module_entry) hook = LayerwiseOffloadHook(state) diff --git a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py index 4886ae1db8..ee34c47cd8 100644 --- a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py +++ b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py @@ -755,7 +755,7 @@ 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: + def prepare_layerwise_offload(self, device: torch.device, *, pin_cpu_memory: bool = True) -> 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 @@ -770,7 +770,8 @@ def prepare_layerwise_offload(self, device: torch.device) -> None: 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) + enable_layerwise_offload(self.language_model, resident_blocks=0, cyclic=False, + pin_cpu_memory=pin_cpu_memory) self._h3_encoder_layerwise_device = device def forward( diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 66a1d2dceb..106e8d7962 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -473,7 +473,7 @@ def load_model( 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) + prepare_layerwise(runtime_device, pin_cpu_memory=fastvideo_args.pin_cpu_memory) use_cpu_offload = False logger.info("Enabled text-only layerwise H3 encoder with CPU token embeddings") @@ -1211,7 +1211,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # 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) + enable_layerwise_offload(model, pin_cpu_memory=fastvideo_args.pin_cpu_memory) # Blocks now hold placeholders; the remaining (non-block) weights and buffers belong on the GPU. model = model.to(get_local_torch_device()) else: diff --git a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py index b96d0042f5..ecc0c51eee 100644 --- a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py +++ b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py @@ -20,8 +20,9 @@ @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)]) +@pytest.mark.parametrize("pin_cpu_memory", [True, False]) def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, env_overrides, quantized, - fused): + fused, pin_cpu_memory): # DiT residency must not accidentally keep encoder layers resident too. env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.override(6)) env_overrides.enter_context(envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.override(False)) @@ -63,8 +64,8 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup 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 + model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) + model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) # 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): @@ -74,7 +75,10 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup 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 + state = manager.forward_hooks["LayerwiseOffloadHook"].state + assert not state.gpu_named_parameters + assert state.pin_cpu_memory == pin_cpu_memory + assert (state.cpu_arena is not None) == pin_cpu_memory 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)) diff --git a/fastvideo/tests/hooks/test_layerwise_offload.py b/fastvideo/tests/hooks/test_layerwise_offload.py index f1505029e6..ed95b6a0d9 100644 --- a/fastvideo/tests/hooks/test_layerwise_offload.py +++ b/fastvideo/tests/hooks/test_layerwise_offload.py @@ -44,6 +44,36 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return x +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +def test_pageable_offload_retains_mapped_storage_and_matches_repeated_forwards(tmp_path): + """File-backed weights stay reclaimable; GPU prefetch preserves their values.""" + model = SimpleModelWithModuleList(num_blocks=3, hidden_size=32) + pointers = {} + for block_index, block in enumerate(model.blocks): + for name, parameter in block.named_parameters(): + path = tmp_path / f"{block_index}-{name}.bin" + mapped = torch.from_file(str(path), shared=True, size=parameter.numel(), dtype=parameter.dtype) + mapped.copy_(parameter.detach().reshape(-1)) + parameter.data = mapped.view_as(parameter) + pointers[block_index, name] = parameter.data_ptr() + reference = SimpleModelWithModuleList(num_blocks=3, hidden_size=32).cuda() + reference.load_state_dict(model.state_dict()) + x = torch.randn(2, 9, 32, device="cuda") + with torch.inference_mode(): + expected = reference(x) + enable_layerwise_offload(model, pin_cpu_memory=False) + for block_index, block in enumerate(model.blocks): + manager = ModuleHookManager.get_from(block) + state = manager.forward_hooks["LayerwiseOffloadHook"].state + assert state.cpu_arena is None + for name, host in state.cpu_named_parameters.items(): + assert not host.is_pinned() + assert host.data_ptr() == pointers[block_index, name] + with torch.inference_mode(): + for _ in range(3): + torch.testing.assert_close(model(x), expected, rtol=0, atol=0) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") def test_layerwise_offload_basic(): """Test basic functionality of layerwise offloading.""" diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 11c8b75e8a..736b415f7d 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -6,6 +6,19 @@ DMD forwards. Preserve VSA sparsity 0.8 and tile size 64. ## Setup and validation +The consumer launch uses seed 1234 and 124 frames for both resolutions. +Pass `--seed 1234 --frames 124 --timed 4` with the two launch benchmark +prompts to collect one warmup and two timed samples per prompt. + +On hosts with limited RAM, `--pageable-host` sets the existing engine +`offload.pin_cpu_memory` option to false. Layerwise DiT and encoder offload +then retain existing CPU storage instead of copying it into pinned arenas. +Checkpoint-backed pages can be reclaimed by the OS, at the cost of slower +host-to-device transfers. This also allows a cached-component recipe to be +tested before resorting to `--lazy`, which reloads components per request. +Validate actual host and GPU peaks for each recipe; disabling pinning does +not make ordinary heap allocations reclaimable checkpoint pages. + 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. diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index 114e649979..29f2b8f212 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -69,6 +69,7 @@ def main(): 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("--pageable-host", action="store_true", help="retain file-backed CPU offload weights without pinning") ap.add_argument("--tile-batch", default=None) ap.add_argument("--frames", type=int, default=243) ap.add_argument("--seed", type=int, default=20260929) @@ -120,7 +121,7 @@ def main(): 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}, + "vae": layerwise, "pin_cpu_memory": not a.pageable_host, "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"} From 00bd0a1412336bc4a76655ee5587cab3abf8a29b Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 12:55:50 -0700 Subject: [PATCH 62/68] [bench]: measure memory on cgroup v1 pods --- scripts/benchmarks/minimax_h3_4090/README.md | 5 ++ .../benchmarks/minimax_h3_4090/bench_pod.py | 35 ++++++---- .../benchmarks/test_h3_benchmark_memory.py | 65 +++++++++++++++++++ 3 files changed, 93 insertions(+), 12 deletions(-) create mode 100644 tests/local_tests/benchmarks/test_h3_benchmark_memory.py diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 736b415f7d..23619ec065 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -19,6 +19,11 @@ tested before resorting to `--lazy`, which reloads components per request. Validate actual host and GPU peaks for each recipe; disabling pinning does not make ordinary heap allocations reclaimable checkpoint pages. +The memory sampler supports cgroup v2 (`memory.current`, `anon`) and v1 +(`memory.usage_in_bytes`, anonymous RSS). Both report pod-wide usage including +file cache separately from anonymous memory. If host counters cannot be read, +the host peak is unknown and NVML GPU capacity sampling continues. + 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. diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index 29f2b8f212..502c59b494 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -18,7 +18,8 @@ class HostMemoryPeak: """Sample pod-wide cgroup usage; anon excludes cached checkpoint file pages.""" - def __init__(self): + def __init__(self, cgroup_root=pathlib.Path("/sys/fs/cgroup")): + self.cgroup_root = cgroup_root self.stop = threading.Event() self.peak_bytes = 0 self.peak_anon_bytes = 0 @@ -38,17 +39,27 @@ def __init__(self): 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"])) - 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) as exc: - self.host_error = f"{type(exc).__name__}: {exc}" - print(f"host memory sampling stopped: {self.host_error}", flush=True) - return + if self.host_error is None: + try: + root = self.cgroup_root + if (root / "memory.current").is_file(): + used = int((root / "memory.current").read_text()) + stats = dict(line.split() for line in (root / "memory.stat").read_text().splitlines()) + anon = int(stats["anon"]) + else: + root = root / "memory" + used = int((root / "memory.usage_in_bytes").read_text()) + stats = dict(line.split() for line in (root / "memory.stat").read_text().splitlines()) + # cgroup v1 RSS counts anonymous memory, excluding file cache. + anon = int(stats["total_rss"] if "total_rss" in stats else stats["rss"]) + self.peak_bytes = max(self.peak_bytes, used) + self.peak_anon_bytes = max(self.peak_anon_bytes, anon) + 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) + # Capacity verification must continue even if host counters are unavailable. + if self._gpu_used is not None: + self.peak_gpu_bytes = max(self.peak_gpu_bytes or 0, self._gpu_used()) self.stop.wait(0.1) def __enter__(self): diff --git a/tests/local_tests/benchmarks/test_h3_benchmark_memory.py b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py new file mode 100644 index 0000000000..44299b4398 --- /dev/null +++ b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py @@ -0,0 +1,65 @@ +"""Benchmark capacity sampling supports both pod cgroup versions, without CUDA.""" +import importlib.util +from pathlib import Path +import sys +from types import SimpleNamespace + +import pytest + + +SCRIPT = Path(__file__).resolve().parents[3] / "scripts/benchmarks/minimax_h3_4090/bench_pod.py" +SPEC = importlib.util.spec_from_file_location("h3_benchmark_memory", SCRIPT) +BENCH = importlib.util.module_from_spec(SPEC) +SPEC.loader.exec_module(BENCH) + + +class OneSample: + def __init__(self): + self.finished = False + + def is_set(self): + return self.finished + + def wait(self, _seconds): + self.finished = True + + +def sample(root, monkeypatch): + monkeypatch.setitem(sys.modules, "pynvml", SimpleNamespace( + nvmlInit=lambda: None, nvmlShutdown=lambda: None, + nvmlDeviceGetHandleByIndex=lambda _index: None, + nvmlDeviceGetMemoryInfo=lambda _handle: SimpleNamespace(used=8192))) + peak = BENCH.HostMemoryPeak(root) + peak.stop = OneSample() + peak._sample() + return peak + + +@pytest.mark.parametrize("version", [1, 2]) +def test_cgroup_memory_excludes_cached_files_from_anonymous_peak(tmp_path, monkeypatch, version): + if version == 2: + (tmp_path / "memory.current").write_text("4096") + (tmp_path / "memory.stat").write_text("anon 1024\nfile 3072\n") + else: + root = tmp_path / "memory" + root.mkdir() + (root / "memory.usage_in_bytes").write_text("4096") + (root / "memory.stat").write_text("rss 512\ntotal_rss 1024\ntotal_cache 3072\n") + peak = sample(tmp_path, monkeypatch) + assert peak.host_error is None + assert peak.peak_bytes == 4096 + assert peak.peak_anon_bytes == 1024 + assert peak.peak_gpu_bytes == 8192 + + +def test_missing_host_counters_do_not_disable_gpu_capacity_measurement(tmp_path, monkeypatch): + peak = sample(tmp_path, monkeypatch) + assert peak.host_error is not None + assert peak.peak_gpu_bytes == 8192 + + +def test_malformed_host_counters_do_not_report_a_measured_zero(tmp_path, monkeypatch): + (tmp_path / "memory.current").write_text("invalid") + peak = sample(tmp_path, monkeypatch) + assert peak.host_error is not None + assert peak.peak_gpu_bytes == 8192 From 659a4840f381621c388fa1e2d10af84ea242f784 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 13:02:54 -0700 Subject: [PATCH 63/68] [bugfix]: keep CPU-targeted checkpoint loads off the GPU --- fastvideo/models/loader/fsdp_load.py | 4 +- .../loader/test_cpu_target_checkpoint.py | 49 +++++++++++++++++++ 2 files changed, 52 insertions(+), 1 deletion(-) create mode 100644 fastvideo/tests/loader/test_cpu_target_checkpoint.py diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 2e1d5831a3..87b53cf653 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -360,7 +360,9 @@ def maybe_load_fsdp_model( packed_nvfp4_export, ) packed_nvfp4_export = None - load_weights_to_cpu = cpu_offload or packed_nvfp4_export is not None + # CPU-targeted layerwise/table loads must not stage the full checkpoint on + # the GPU before offload hooks or skipped projection weights are applied. + load_weights_to_cpu = device.type == "cpu" or 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) diff --git a/fastvideo/tests/loader/test_cpu_target_checkpoint.py b/fastvideo/tests/loader/test_cpu_target_checkpoint.py new file mode 100644 index 0000000000..093c6bf7b8 --- /dev/null +++ b/fastvideo/tests/loader/test_cpu_target_checkpoint.py @@ -0,0 +1,49 @@ +"""CPU-targeted inference must never stage its checkpoint on the GPU.""" +from types import SimpleNamespace + +import torch +from torch import nn +from safetensors.torch import save_file + +from fastvideo.models.loader import fsdp_load + + +class TinyCPUModel(nn.Module): + param_names_mapping = {} + + def __init__(self, config): + super().__init__() + self.config = config + self.linear = nn.Linear(8, 8) + + +def test_cpu_target_reads_cpu_checkpoint_and_preserves_weights(tmp_path, monkeypatch): + weights = {"linear.weight": torch.arange(64, dtype=torch.float32).view(8, 8), + "linear.bias": torch.arange(8, dtype=torch.float32)} + path = tmp_path / "model.safetensors" + save_file(weights, path) + real_iterator = fsdp_load.safetensors_weights_iterator + placements = [] + + def cpu_iterator(files, *, to_cpu): + assert to_cpu, "CPU target must not stage checkpoint tensors on the GPU" + for name, tensor in real_iterator(files, to_cpu=to_cpu): + placements.append(tensor.device.type) + yield name, tensor + + monkeypatch.setattr(fsdp_load, "safetensors_weights_iterator", cpu_iterator) + model = fsdp_load.maybe_load_fsdp_model( + model_cls=TinyCPUModel, + init_params={"config": SimpleNamespace(quant_config=None)}, + weight_dir_list=[str(path)], device=torch.device("cpu"), + hsdp_replicate_dim=1, hsdp_shard_dim=1, + default_dtype=torch.float32, param_dtype=torch.float32, + reduce_dtype=torch.float32, training_mode=False, cpu_offload=False, + ) + assert placements == ["cpu", "cpu"] + for name, value in model.state_dict().items(): + assert value.device.type == "cpu" + torch.testing.assert_close(value, weights[name], rtol=0, atol=0) + x = torch.randn(2, 8) + torch.testing.assert_close(model.linear(x), torch.nn.functional.linear(x, weights["linear.weight"], + weights["linear.bias"]), rtol=0, atol=0) From 2c0cd2a18b114b7d6101097c19c4e5740e64f502 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 13:07:35 -0700 Subject: [PATCH 64/68] [bench]: preserve peak memory for failed generations --- .../benchmarks/minimax_h3_4090/bench_pod.py | 31 +++++++---- .../benchmarks/test_h3_benchmark_memory.py | 51 +++++++++++++++++++ 2 files changed, 72 insertions(+), 10 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index 502c59b494..bfd50ef8ef 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -72,6 +72,15 @@ def __exit__(self, *_args): if self._nvml_shutdown is not None: self._nvml_shutdown() + def metrics(self): + return { + "peak_gpu_used_gib": (round(self.peak_gpu_bytes / 2**30, 3) + if self.peak_gpu_bytes is not None else None), + "peak_host_cgroup_gib": (round(self.peak_bytes / 2**30, 3) if self.host_error is None else None), + "peak_host_anon_gib": (round(self.peak_anon_bytes / 2**30, 3) if self.host_error is None else None), + "host_memory_error": self.host_error, + } + def main(): ap = argparse.ArgumentParser() @@ -173,19 +182,21 @@ def main(): "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) + try: + with HostMemoryPeak() as host_peak: + generator.generate(request) + except Exception as exc: + failed = {"prompt": pid, "warmup": i < a.warmup, + "wall_s": round(time.perf_counter() - t, 2), + "error": f"{type(exc).__name__}: {exc}", **host_peak.metrics()} + results.setdefault("failed_runs", []).append(failed) + print("RUN_FAILED", json.dumps(failed), flush=True) + (out_dir / "results.json").write_text(json.dumps(results, indent=1)) + raise 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), - # 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}) + **host_peak.metrics()}) 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/tests/local_tests/benchmarks/test_h3_benchmark_memory.py b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py index 44299b4398..796c93c48e 100644 --- a/tests/local_tests/benchmarks/test_h3_benchmark_memory.py +++ b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py @@ -1,5 +1,6 @@ """Benchmark capacity sampling supports both pod cgroup versions, without CUDA.""" import importlib.util +import json from pathlib import Path import sys from types import SimpleNamespace @@ -63,3 +64,53 @@ def test_malformed_host_counters_do_not_report_a_measured_zero(tmp_path, monkeyp peak = sample(tmp_path, monkeypatch) assert peak.host_error is not None assert peak.peak_gpu_bytes == 8192 + assert peak.metrics()["peak_host_cgroup_gib"] is None + assert peak.metrics()["peak_host_anon_gib"] is None + + +def test_failed_generation_preserves_sampled_memory_and_shuts_down(tmp_path, monkeypatch): + model = tmp_path / "model" + model.mkdir() + (model / "fastvideo_inference.json").write_text("{}") + prompts = tmp_path / "prompts.json" + prompts.write_text('{"ceramics": "test"}') + shutdown = [] + + class FailingGenerator: + def generate(self, _request): + raise RuntimeError("CUDA out of memory") + + def shutdown(self): + shutdown.append(True) + + class Memory: + def __enter__(self): + return self + + def __exit__(self, *_args): + pass + + def metrics(self): + return {"peak_gpu_used_gib": 11.9, "peak_host_cgroup_gib": 28.3, + "peak_host_anon_gib": 20.1, "host_memory_error": None} + + fake_torch = SimpleNamespace(__version__="test", version=SimpleNamespace(cuda="test"), + cuda=SimpleNamespace(get_device_name=lambda _index: "test GPU")) + fake_video = SimpleNamespace(VideoGenerator=SimpleNamespace(from_config=lambda _config: FailingGenerator())) + monkeypatch.setitem(sys.modules, "torch", fake_torch) + monkeypatch.setitem(sys.modules, "fastvideo", fake_video) + monkeypatch.setattr(BENCH, "HostMemoryPeak", Memory) + monkeypatch.setattr(BENCH.subprocess, "check_output", lambda *_args, **_kwargs: "test GPU") + monkeypatch.setattr(sys, "argv", [str(SCRIPT), "failure", str(model), "fp8", "--prompt-file", str(prompts), + "--output-root", str(tmp_path / "outputs")]) + with pytest.raises(RuntimeError, match="CUDA out of memory"): + BENCH.main() + raw = json.loads((tmp_path / "outputs/failure/results.json").read_text()) + assert raw["runs"] == [] + assert len(raw["failed_runs"]) == 1 + failed = raw["failed_runs"][0] + assert failed["peak_gpu_used_gib"] == 11.9 + assert failed["peak_host_cgroup_gib"] == 28.3 + assert failed["warmup"] + assert "CUDA out of memory" in failed["error"] + assert shutdown == [True] From 3c011eff2a92d97d360fc504a8d383cc821693b0 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 13:16:07 -0700 Subject: [PATCH 65/68] [bench]: build exact V2 modulation tables with provenance --- scripts/benchmarks/minimax_h3_4090/README.md | 18 ++++ .../minimax_h3_4090/precompute_v2_adaln.py | 102 ++++++++++++++++++ 2 files changed, 120 insertions(+) create mode 100644 scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py diff --git a/scripts/benchmarks/minimax_h3_4090/README.md b/scripts/benchmarks/minimax_h3_4090/README.md index 23619ec065..4b7132f028 100644 --- a/scripts/benchmarks/minimax_h3_4090/README.md +++ b/scripts/benchmarks/minimax_h3_4090/README.md @@ -23,6 +23,24 @@ The memory sampler supports cgroup v2 (`memory.current`, `anon`) and v1 (`memory.usage_in_bytes`, anonymous RSS). Both report pod-wide usage including file cache separately from anonymous memory. If host counters cannot be read, the host peak is unknown and NVML GPU capacity sampling continues. +Failed generate calls retain sampled memory and an error in `failed_runs`, +separately from completed timing samples. + +Full-rank V2 can replace its large AdaLN projection weights with exact outputs +for the checkpoint's fixed T2AV DMD ladder: + +```bash +python -P scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py \ + /workspace/vol/v2-fp8 /workspace/v2-adaln.pt --source-commit +export FASTVIDEO_H3_ADALN_TABLE=/workspace/v2-adaln.pt +``` + +The helper checks every block/rung against the original modulation module, +writes the exact inputs alongside the table, and records hashes and checkpoint +provenance in a JSON sidecar. Validate the real pipeline timestep embeddings +against the saved inputs before using a new table in headline measurements. +Use the table with its matching checkpoint and T2AV contract; unsupported +timestep keys raise rather than silently changing the schedule. 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 diff --git a/scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py b/scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py new file mode 100644 index 0000000000..af90f211a4 --- /dev/null +++ b/scripts/benchmarks/minimax_h3_4090/precompute_v2_adaln.py @@ -0,0 +1,102 @@ +"""Exact V2 T2AV modulation tables for the checkpoint's fixed DMD ladder.""" +import argparse +import hashlib +import json +from pathlib import Path +from types import SimpleNamespace + +import torch +import torch.nn.functional as F +from safetensors import safe_open +from fastvideo.layers.mlp import MLP +from fastvideo.layers.visual_embedding import Timesteps +from fastvideo.models.dits.minimax_h3 import MiniMaxH3AdaLayerNormModulation +from fastvideo.models.schedulers.scheduling_minimax_h3 import MiniMaxH3Scheduler +from fastvideo.pipelines.basic.minimax_h3.packing import build_row_timesteps +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_denoising import MiniMaxH3DenoisingStage + + +@torch.inference_mode() +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('model', type=Path) + parser.add_argument('output', type=Path) + parser.add_argument('--source-commit', help='FastVideo revision used to build the table') + args = parser.parse_args() + config = json.loads((args.model / 'transformer/config.json').read_text()) + contract = json.loads((args.model / 'fastvideo_inference.json').read_text()) + if config.get('adaln_rank') is not None or contract.get('task') != 't2av': + raise ValueError('This tool supports full-rank V2 T2AV only') + files = list((args.model / 'transformer').glob('*.safetensors')) + locations = {} + for path in files: + with safe_open(path, framework='pt', device='cpu') as reader: + for key in reader.keys(): + if 'adaln_proj.linear.' in key or key.startswith('time_embedder.'): + locations[key] = path + + def tensor(key, dtype=None): + with safe_open(locations[key], framework='pt', device='cpu') as reader: + return reader.get_tensor(key).to(device='cuda', dtype=dtype) + + time_proj = Timesteps(config['freq_dim'], flip_sin_to_cos=True, downscale_freq_shift=0) + embedder = MLP(config['freq_dim'], config['time_embed_hidden_dim'], config['time_embed_dim'], + act_type='silu', dtype=torch.float32).cuda().eval() + for destination, source in (('fc_in', 'linear_1'), ('fc_out', 'linear_2')): + layer = getattr(embedder, destination) + layer.weight.data = tensor(f'time_embedder.{source}.weight', torch.float32) + layer.bias.data = tensor(f'time_embedder.{source}.bias', torch.float32) + video = MiniMaxH3Scheduler(shift=contract['video_scheduler_shift']) + audio = MiniMaxH3Scheduler(shift=contract['audio_scheduler_shift']) + stage = MiniMaxH3DenoisingStage(None, video, audio) + stage._set_dmd_schedule(contract['dmd_denoising_steps'], contract['num_inference_steps'], torch.device('cuda')) + # No reference rows; text rows share the video timestep in the real packer. + layout = SimpleNamespace(sequence_length=3, video_indices=torch.tensor([0]), + audio_indices=torch.tensor([1]), num_condition_video_rows=0, + num_condition_audio_rows=0) + inputs, embeddings = {}, {} + for vt, at in zip(video.timesteps, audio.timesteps, strict=True): + unique, _ = build_row_timesteps(layout, float(vt), float(at), float(vt), 1.0) + unique = unique.cuda() + temb = embedder(time_proj(unique).to(embedder.fc_in.weight.dtype)) + key = repr((tuple(unique.reshape(-1).tolist()), tuple(temb.shape), str(temb.dtype))) + embeddings[key] = temb + inputs[key] = F.silu(temb) + if len(inputs) != contract['transformer_forwards']: + raise ValueError('The contract must supply one unique timestep key per transformer forward') + tables = {} + for block in range(config['num_layers']): + prefix = f'transformer_blocks.{block}.adaln_proj.linear' + weight, bias = tensor(prefix + '.weight'), tensor(prefix + '.bias') + # Use the original release module as an independent arithmetic check. + with torch.device('meta'): + reference = MiniMaxH3AdaLayerNormModulation(config['time_embed_dim'], config['hidden_size']) + reference.linear.weight = torch.nn.Parameter(weight, requires_grad=False) + reference.linear.bias = torch.nn.Parameter(bias, requires_grad=False) + tables[block] = {} + for key, x in inputs.items(): + result = F.linear(x.to(weight.dtype), weight, bias) + expected = torch.cat(reference(embeddings[key]), dim=-1).view_as(result) + torch.testing.assert_close(result, expected, atol=0, rtol=0) + tables[block][key] = result.cpu() + del reference, weight, bias + print('BLOCK_VALIDATED', block, flush=True) + args.output.parent.mkdir(parents=True, exist_ok=True) + torch.save(tables, args.output) + torch.save({key: value.cpu() for key, value in inputs.items()}, str(args.output) + '.inputs') + revision_file = args.model / '.cache/huggingface/download/fastvideo_inference.json.metadata' + metadata = {'config_sha256': hashlib.sha256((args.model / 'transformer/config.json').read_bytes()).hexdigest(), + 'contract_sha256': hashlib.sha256((args.model / 'fastvideo_inference.json').read_bytes()).hexdigest(), + 'model_revision': revision_file.read_text().splitlines()[0] if revision_file.is_file() else None, + 'source_commit': args.source_commit, 'torch': str(torch.__version__), 'cuda': torch.version.cuda, + 'gpu': torch.cuda.get_device_name(0), + 'helper_sha256': hashlib.sha256(Path(__file__).read_bytes()).hexdigest(), + 'contract': contract, 'blocks': len(tables), 'timestep_keys': list(inputs), + 'table_sha256': hashlib.sha256(args.output.read_bytes()).hexdigest(), + 'validation': 'Every block/rung exactly equals the original release modulation module on sm89.'} + Path(str(args.output) + '.json').write_text(json.dumps(metadata, indent=2) + '\n') + print('TABLE_DONE', args.output, flush=True) + + +if __name__ == '__main__': + main() From 1299c7d084483c20a6e3c66e90a7711523ecf9ce Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 13:26:02 -0700 Subject: [PATCH 66/68] [bench]: render a showcase prompt once without extra warmup clips --- .../benchmarks/minimax_h3_4090/bench_pod.py | 14 +++--- .../benchmarks/test_h3_benchmark_memory.py | 45 ++++++++++++++----- 2 files changed, 42 insertions(+), 17 deletions(-) diff --git a/scripts/benchmarks/minimax_h3_4090/bench_pod.py b/scripts/benchmarks/minimax_h3_4090/bench_pod.py index bfd50ef8ef..4882b9995d 100644 --- a/scripts/benchmarks/minimax_h3_4090/bench_pod.py +++ b/scripts/benchmarks/minimax_h3_4090/bench_pod.py @@ -105,8 +105,9 @@ def main(): 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)") + ap.add_argument("--once", action="store_true", help="render each selected prompt once, without extra warmup clips") a = ap.parse_args() - if a.timed < 2 or a.warmup < 1: + if not a.once and (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") @@ -168,13 +169,15 @@ def main(): "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": []} + "gpu": torch.cuda.get_device_name(0), "config": config, "sampling": sampling, + "mode": "showcase_once" if a.once else "benchmark", "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)] + order = ids if a.once else [ids[i % len(ids)] for i in range(a.warmup + a.timed)] + warmup = 1 if a.once else a.warmup try: for i, pid in enumerate(order): request = {"prompt": texts[pid], "negative_prompt": "", @@ -186,7 +189,7 @@ def main(): with HostMemoryPeak() as host_peak: generator.generate(request) except Exception as exc: - failed = {"prompt": pid, "warmup": i < a.warmup, + failed = {"prompt": pid, "warmup": i < warmup, "wall_s": round(time.perf_counter() - t, 2), "error": f"{type(exc).__name__}: {exc}", **host_peak.metrics()} results.setdefault("failed_runs", []).append(failed) @@ -194,12 +197,13 @@ def main(): (out_dir / "results.json").write_text(json.dumps(results, indent=1)) raise wall = round(time.perf_counter() - t, 2) - results["runs"].append({"prompt": pid, "warmup": i < a.warmup, "wall_s": wall, + results["runs"].append({"prompt": pid, "warmup": i < warmup, "wall_s": wall, "clip": request["output"]["output_path"], **host_peak.metrics()}) timed = [run["wall_s"] for run in results["runs"] if not run["warmup"]] if timed: results["median_e2e_s"] = statistics.median(timed) + results["mean_e2e_s"] = statistics.mean(timed) print("RUN", json.dumps(results["runs"][-1]), flush=True) (out_dir / "results.json").write_text(json.dumps(results, indent=1)) finally: diff --git a/tests/local_tests/benchmarks/test_h3_benchmark_memory.py b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py index 796c93c48e..72e41829c2 100644 --- a/tests/local_tests/benchmarks/test_h3_benchmark_memory.py +++ b/tests/local_tests/benchmarks/test_h3_benchmark_memory.py @@ -68,17 +68,21 @@ def test_malformed_host_counters_do_not_report_a_measured_zero(tmp_path, monkeyp assert peak.metrics()["peak_host_anon_gib"] is None -def test_failed_generation_preserves_sampled_memory_and_shuts_down(tmp_path, monkeypatch): +@pytest.mark.parametrize("fail,once", [(True, False), (False, False), (False, True)]) +def test_generation_receipts_preserve_memory_and_once_never_adds_clips(tmp_path, monkeypatch, fail, once): model = tmp_path / "model" model.mkdir() (model / "fastvideo_inference.json").write_text("{}") prompts = tmp_path / "prompts.json" prompts.write_text('{"ceramics": "test"}') shutdown = [] + requests = [] class FailingGenerator: def generate(self, _request): - raise RuntimeError("CUDA out of memory") + requests.append(_request) + if fail: + raise RuntimeError("CUDA out of memory") def shutdown(self): shutdown.append(True) @@ -101,16 +105,33 @@ def metrics(self): monkeypatch.setitem(sys.modules, "fastvideo", fake_video) monkeypatch.setattr(BENCH, "HostMemoryPeak", Memory) monkeypatch.setattr(BENCH.subprocess, "check_output", lambda *_args, **_kwargs: "test GPU") - monkeypatch.setattr(sys, "argv", [str(SCRIPT), "failure", str(model), "fp8", "--prompt-file", str(prompts), - "--output-root", str(tmp_path / "outputs")]) - with pytest.raises(RuntimeError, match="CUDA out of memory"): + argv = [str(SCRIPT), "failure", str(model), "fp8", "--prompt-file", str(prompts), + "--output-root", str(tmp_path / "outputs")] + if once: + argv.append("--once") + monkeypatch.setattr(sys, "argv", argv) + if fail: + with pytest.raises(RuntimeError, match="CUDA out of memory"): + BENCH.main() + else: BENCH.main() raw = json.loads((tmp_path / "outputs/failure/results.json").read_text()) - assert raw["runs"] == [] - assert len(raw["failed_runs"]) == 1 - failed = raw["failed_runs"][0] - assert failed["peak_gpu_used_gib"] == 11.9 - assert failed["peak_host_cgroup_gib"] == 28.3 - assert failed["warmup"] - assert "CUDA out of memory" in failed["error"] + if fail: + assert raw["runs"] == [] + assert len(raw["failed_runs"]) == 1 + failed = raw["failed_runs"][0] + assert failed["peak_gpu_used_gib"] == 11.9 + assert failed["peak_host_cgroup_gib"] == 28.3 + assert failed["warmup"] + assert "CUDA out of memory" in failed["error"] + else: + assert len(requests) == len(raw["runs"]) == (1 if once else 3) + assert raw["runs"][0]["warmup"] + assert all(run["peak_gpu_used_gib"] == 11.9 for run in raw["runs"]) + if once: + assert "median_e2e_s" not in raw + assert "mean_e2e_s" not in raw + else: + assert not any(run["warmup"] for run in raw["runs"][1:]) + assert "mean_e2e_s" in raw assert shutdown == [True] From 73ee8c2012ac60c34ea870c76f95d543fad354ae Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 13:41:23 -0700 Subject: [PATCH 67/68] [perf]: retain mapped H3 encoder weights on pageable hosts --- .../encoders/minimax_h3_checkpoint_nvfp4.py | 1 + .../models/encoders/minimax_h3_qwen3_vl.py | 20 +++++++++++++++- fastvideo/models/loader/component_loader.py | 12 +++++++++- .../test_minimax_h3_encoder_layerwise.py | 23 ++++++++++++++++-- ...st_minimax_h3_qwen3_vl_checkpoint_nvfp4.py | 24 ++++++++++++++++++- 5 files changed, 75 insertions(+), 5 deletions(-) diff --git a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py index 933041e88d..27925846d6 100644 --- a/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py +++ b/fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py @@ -346,6 +346,7 @@ def load(param: torch.Tensor, loaded_weight: torch.Tensor, *args: Any, **kwargs: f"got {loaded_weight.dtype} for a {param.dtype} parameter of shape {tuple(param.shape)}") return base_loader(param, loaded_weight, *args, **kwargs) + load._h3_base_loader = base_loader return load diff --git a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py index ee34c47cd8..6d8a634395 100644 --- a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py +++ b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py @@ -805,10 +805,28 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}") parameter = parameters[name] loader = getattr(parameter, "weight_loader", default_weight_loader) - loader(parameter, tensor) + base_loader = getattr(loader, "_h3_base_loader", loader) + copy_only = (base_loader is default_weight_loader or getattr(base_loader, "__func__", None) in ( + ColumnParallelLinear.weight_loader, RowParallelLinear.weight_loader, + VocabParallelEmbedding.weight_loader)) + if (getattr(self, "_h3_checkpoint_backed_cpu", False) and copy_only + and parameter.device.type == tensor.device.type == "cpu" + and parameter.shape == tensor.shape and parameter.dtype == tensor.dtype): + # TP=1 loaders only copy already matching tensors. Keep the mapping + # so the OS can reclaim encoder checkpoint pages during denoising. + # Preserve the Parameter and its loader/quantization attributes. + parameter.data = tensor.detach() + else: + loader(parameter, tensor) loaded.add(name) return loaded + def enable_checkpoint_backed_cpu_load(self) -> None: + """Retain immutable CPU checkpoint storage for single-GPU streamed inference.""" + if get_tp_world_size() != 1: + raise ValueError("Checkpoint-backed H3 encoder requires tensor parallel size 1") + self._h3_checkpoint_backed_cpu = True + def _is_omitted_checkpoint_key(self, name: str) -> bool: """Return whether a valid checkpoint key belongs to an unbuilt layer.""" language_model = self.language_model diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 106e8d7962..fcaed55507 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -422,6 +422,12 @@ def load_model( with target_device: model = model_cls(model_config) # type: ignore + retain_checkpoint = getattr(model, "enable_checkpoint_backed_cpu_load", None) + checkpoint_backed_cpu = (target_device.type == "cpu" and envs.FASTVIDEO_H3_ENCODER_LAYERWISE.get() + and callable(retain_checkpoint) and not fastvideo_args.pin_cpu_memory) + if checkpoint_backed_cpu: + retain_checkpoint() + weights_to_load = {name for name, _ in model.named_parameters()} if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None): if os.path.isdir(checkpoint_path): @@ -462,7 +468,11 @@ def load_model( f"checkpoint: {weights_not_loaded}") if checkpoint_quant_config is not None: - processed_linears = _process_quantized_text_encoder_weights(model, runtime_device) + # NVFP4 validation and scalar derivation work on the host. Moving + # packed layers to CUDA and back would discard checkpoint mappings. + process_device = (target_device if checkpoint_backed_cpu and checkpoint_quant_config.get_name() == "nvfp4" + else runtime_device) + processed_linears = _process_quantized_text_encoder_weights(model, process_device) logger.info("Validated %d serialized %s text-encoder linears", processed_linears, checkpoint_quant_config.get_name()) diff --git a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py index ecc0c51eee..abd6e386a0 100644 --- a/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py +++ b/fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py @@ -20,9 +20,9 @@ @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)]) -@pytest.mark.parametrize("pin_cpu_memory", [True, False]) +@pytest.mark.parametrize("pin_cpu_memory,checkpoint_backed", [(True, False), (False, False), (False, True)]) def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, env_overrides, quantized, - fused, pin_cpu_memory): + fused, pin_cpu_memory, checkpoint_backed, tmp_path): # DiT residency must not accidentally keep encoder layers resident too. env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.override(6)) env_overrides.enter_context(envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.override(False)) @@ -55,6 +55,11 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup 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) + checkpoint = tmp_path / "encoder.safetensors" + if checkpoint_backed: + from safetensors.torch import save_file + save_file({name: parameter.detach().cpu().contiguous() for name, parameter in model.named_parameters()}, + checkpoint) ids = torch.tensor([1, 7, 4, 21, 5, 31, 18], device="cuda") model.to("cuda") expected = model.encode_ids(ids) @@ -64,6 +69,16 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup if hasattr(layer, "_nvfp4_fused_dequant"): layer._nvfp4_fused_dequant = True model.to("cpu") + mapped = {} + if checkpoint_backed: + from safetensors.torch import load_file + mapped = load_file(checkpoint, device="cpu") + model.enable_checkpoint_backed_cpu_load() + model.load_weights(mapped.items()) + if quantized: + _process_quantized_text_encoder_weights(model, torch.device("cpu")) + for name, parameter in model.named_parameters(): + assert parameter.data_ptr() == mapped[name].data_ptr() model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) # repeated setup is harmless assert model.language_model.embed_tokens.weight.device.type == "cpu" @@ -79,6 +94,10 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup assert not state.gpu_named_parameters assert state.pin_cpu_memory == pin_cpu_memory assert (state.cpu_arena is not None) == pin_cpu_memory + if checkpoint_backed: + for name, tensor in state.cpu_named_parameters.items(): + layer_name = next(key for key, value in model.named_modules() if value is layer) + assert tensor.data_ptr() == mapped[f"{layer_name}.{name}"].data_ptr() 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)) 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 e3d2e4921f..ed539e0738 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 @@ -454,10 +454,13 @@ def _tiny_conditioner_config(keep_bf16: tuple[str, ...] = ("mlp.down_proj", )) - return config -def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) -> None: +@pytest.mark.parametrize("checkpoint_backed", [False, True]) +def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup, checkpoint_backed) -> None: """The real chain: a conditioner built with the NVFP4 config, checkpoint keys spelled the way the converter writes them, ``load_weights``, the strict missing-tensor check, then the post-load hook.""" model = MiniMaxH3Qwen3VLConditioner(_tiny_conditioner_config()) + if checkpoint_backed: + model.enable_checkpoint_backed_cpu_load() layer = model.language_model.layers[0] assert isinstance(layer.self_attn.q_proj.quant_method, MiniMaxH3SerializedNVFP4LinearMethod) assert isinstance(layer.self_attn.o_proj.quant_method, MiniMaxH3SerializedNVFP4LinearMethod) @@ -482,6 +485,9 @@ def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) loaded = model.load_weights(iter(checkpoint.items())) assert loaded == expected assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 6 + if checkpoint_backed: + for name, parameter in model.named_parameters(): + assert parameter.data_ptr() == checkpoint[f"model.{name}"].data_ptr() assert layer.self_attn.q_proj._nvfp4_alpha.item() == pytest.approx(0.5) assert layer.mlp.up_proj._nvfp4_alpha.item() == pytest.approx(0.5) @@ -504,3 +510,19 @@ def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) stray["model.language_model.layers.0.self_attn.q_proj.weight"] = torch.zeros(128, 128) with pytest.raises(ValueError, match="Unexpected"): model.load_weights(iter(stray.items())) + + +def test_checkpoint_backed_load_keeps_custom_loader_semantics(distributed_setup) -> None: + model = MiniMaxH3Qwen3VLConditioner(_tiny_conditioner_config()) + model.enable_checkpoint_backed_cpu_load() + name = "language_model.layers.0.input_layernorm.weight" + parameter = dict(model.named_parameters())[name] + source = torch.ones_like(parameter) + + def custom_loader(target, value): + target.data.copy_(value + 1) + + parameter.weight_loader = custom_loader + model.load_weights([(name, source)]) + torch.testing.assert_close(parameter, source + 1, rtol=0, atol=0) + assert parameter.data_ptr() != source.data_ptr() From cda05792603657954575d32d9a8917a0ea3d9686 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Mon, 5 Oct 2026 14:05:21 -0700 Subject: [PATCH 68/68] [bugfix]: preserve padded vocabulary checkpoint validation --- fastvideo/models/encoders/minimax_h3_qwen3_vl.py | 6 ++++++ .../test_minimax_h3_qwen3_vl_checkpoint_nvfp4.py | 13 +++++++++++++ 2 files changed, 19 insertions(+) diff --git a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py index 6d8a634395..3f958ec1b1 100644 --- a/fastvideo/models/encoders/minimax_h3_qwen3_vl.py +++ b/fastvideo/models/encoders/minimax_h3_qwen3_vl.py @@ -809,6 +809,12 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: copy_only = (base_loader is default_weight_loader or getattr(base_loader, "__func__", None) in ( ColumnParallelLinear.weight_loader, RowParallelLinear.weight_loader, VocabParallelEmbedding.weight_loader)) + if getattr(base_loader, "__func__", None) is VocabParallelEmbedding.weight_loader: + output_dim = getattr(parameter, "output_dim", None) + copy_only = (not getattr(parameter, "is_gguf_weight_type", False) + and getattr(parameter, "packed_dim", None) is None + and (output_dim is None or (tensor.ndim > output_dim + and tensor.shape[output_dim] == base_loader.__self__.org_vocab_size))) if (getattr(self, "_h3_checkpoint_backed_cpu", False) and copy_only and parameter.device.type == tensor.device.type == "cpu" and parameter.shape == tensor.shape and parameter.dtype == tensor.dtype): 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 ed539e0738..7fa85c1f97 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 @@ -526,3 +526,16 @@ def custom_loader(target, value): model.load_weights([(name, source)]) torch.testing.assert_close(parameter, source + 1, rtol=0, atol=0) assert parameter.data_ptr() != source.data_ptr() + + +def test_checkpoint_backed_embedding_rejects_padded_checkpoint_rows(distributed_setup) -> None: + config = _tiny_conditioner_config() + config.arch_config.vocab_size = 63 + model = MiniMaxH3Qwen3VLConditioner(config) + model.enable_checkpoint_backed_cpu_load() + weight = model.language_model.embed_tokens.weight + assert weight.shape[0] == 64 + # Matching the allocated padded shape must not bypass the loader's check + # that checkpoint rows match the original unpadded vocabulary. + with pytest.raises(AssertionError): + model.load_weights([("language_model.embed_tokens.weight", torch.ones_like(weight))])