Compare commits

..
Author SHA1 Message Date
SolitaryThinker 825c6365e0 bundle variant -> preset resolution, pipeline pin, and LTX-2.5 examples
A bundle's training variant (header-declared `variant`, filename token as
fallback) now selects its sampling preset: distilled -> ltx2_distilled
(8 steps, cfg 1.0, distilled sigmas engage at <=8 steps), sft/base ->
ltx2_base. The bundle table also pins the pipeline class, so
VideoGenerator.from_pretrained(<bundle>) resolves without an override.
Adds basic_ltx2_5{,_distilled}.py examples (--model-path takes the FILE,
--gemma-root declares the paired encoder root), a README section, and CPU
tests for variant/preset resolution, table integrity, and the examples.
2026-08-11 17:03:42 -07:00
SolitaryThinker bb1fb7c149 WIP: route single-file bundles through the modular train stack
carried over the SWE session's uncommitted work across the rebase onto
main@8208536c: model_index_and_component_path helper feeds the train
moduleloader, trainer + legacy distillation pipeline pick up bundle-aware
component paths, nvfp4_config/registry adjustments, and the ltx2_5
fine-tuning example configs (untracked until now).

(--no-verify: pre-commit not configured in this worktree)
2026-08-10 13:06:06 -07:00
William Lin 436c44e583 bundle route through the real pipeline: module map, component paths, loader branches, connector factorization 2026-08-10 12:52:52 -07:00
William Lin 6343eadbec resolve pipeline config for single-file bundles from declared transformer class
A bundle path is a file, so config resolution took the directory branch and
was rejected for having no model_index.json. Resolve it from the transformer
_class_name the checkpoint declares about itself, via an explicit alias table.
Keyed on the class name, not the registry model_id (a positional index that
rebinds on registration-order changes). No wildcard fallback and no file-name
inference; an unmapped class raises naming the table and the override.

verify_model_config_and_directory is unchanged.
2026-08-10 12:51:08 -07:00
William Lin 3f5614fd03 single-file bundle loader: metadata config, prefix routing, per-stream ff bias, encoder-root resolution 2026-08-10 12:51:08 -07:00
185 changed files with 2037 additions and 23456 deletions
-149
View File
@@ -1,149 +0,0 @@
name: macOS MLX Smoke
on:
pull_request:
branches: [main]
paths:
- ".github/workflows/ci-macos-mlx.yml"
- "fastvideo/mlx_runtime/**"
- "fastvideo/tests/mlx/**"
- "fastvideo/tests/platforms/test_mps_vsa_error.py"
- "fastvideo/platforms/mps.py"
- "fastvideo/platforms/__init__.py"
- "fastvideo/__init__.py"
- "examples/inference/basic/mlx_*.py"
- "fastvideo/benchmarks/mlx_*.py"
- "pyproject.toml"
workflow_dispatch:
permissions:
contents: read
concurrency:
group: macos-mlx-${{ github.ref }}
cancel-in-progress: true
jobs:
mlx-smoke:
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
runs-on: macos-15
timeout-minutes: 25
env:
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
TOKENIZERS_PARALLELISM: "false"
MASTER_ADDR: localhost
MASTER_PORT: "29513"
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
cache: pip
- uses: astral-sh/setup-uv@v3
- name: Install lightweight MLX smoke dependencies
run: |
uv pip install --system \
--index-url https://download.pytorch.org/whl/cpu \
torch==2.11.0 torchvision torchaudio
uv pip install --system \
pytest numpy scipy pillow imageio einops cloudpickle filelock \
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
- name: Show Apple runtime
run: |
python - <<'PY'
import platform
import mlx.core as mx
import torch
print("machine:", platform.machine())
print("processor:", platform.processor())
print("mlx default device:", mx.default_device())
memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable"
print("mlx memory_size:", memory_size)
print("torch:", torch.__version__)
print("torch mps available:", torch.backends.mps.is_available())
PY
- name: Run MLX smoke tests
run: |
python -m pytest \
fastvideo/tests/mlx/test_dmd_sampling.py \
fastvideo/tests/mlx/test_memory_limits.py \
fastvideo/tests/mlx/test_quant_capability.py \
fastvideo/tests/mlx/test_mlx_dit_parity.py \
fastvideo/tests/mlx/test_mlx_compile_parity.py \
fastvideo/tests/mlx/test_mlx_checkpoint.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
fastvideo/tests/mlx/test_mlx_refine.py \
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
fastvideo/tests/mlx/test_wan22_sample.py \
fastvideo/tests/mlx/test_windowed_attention.py \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
fastvideo/tests/platforms/test_mps_vsa_error.py \
-q
# Same tests on MLX's CPU backend. Hosted macOS runners are scarce and
# slower to schedule; this Linux job gives fast PR signal on the identical
# graph (the parity tests were designed to be backend-agnostic), while the
# macOS job above stays the source of truth for Metal behavior.
mlx-smoke-linux-cpu:
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
runs-on: ubuntu-latest
timeout-minutes: 20
env:
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
TOKENIZERS_PARALLELISM: "false"
MASTER_ADDR: localhost
MASTER_PORT: "29513"
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.12"
cache: pip
- uses: astral-sh/setup-uv@v3
- name: Install lightweight MLX smoke dependencies (CPU backend)
run: |
uv pip install --system \
--index-url https://download.pytorch.org/whl/cpu \
torch==2.11.0 torchvision torchaudio
uv pip install --system \
pytest numpy scipy pillow imageio einops cloudpickle filelock \
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
- name: Run MLX smoke tests (CPU backend)
run: |
python -m pytest \
fastvideo/tests/mlx/test_dmd_sampling.py \
fastvideo/tests/mlx/test_memory_limits.py \
fastvideo/tests/mlx/test_quant_capability.py \
fastvideo/tests/mlx/test_mlx_dit_parity.py \
fastvideo/tests/mlx/test_mlx_compile_parity.py \
fastvideo/tests/mlx/test_mlx_checkpoint.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
fastvideo/tests/mlx/test_mlx_refine.py \
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
fastvideo/tests/mlx/test_wan22_sample.py \
fastvideo/tests/mlx/test_windowed_attention.py \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
fastvideo/tests/platforms/test_mps_vsa_error.py \
-q
-2
View File
@@ -6,7 +6,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
@@ -17,7 +16,6 @@ on:
paths:
- 'docs/**'
- 'examples/**'
- 'scripts/inference/**'
- 'mkdocs.yml'
- 'requirements-mkdocs.in'
- 'requirements-mkdocs.txt'
-6
View File
@@ -9,7 +9,6 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
@@ -63,11 +62,6 @@ UV_TORCH_BACKEND=cu126 uv pip install fastvideo
Use `UV_TORCH_BACKEND=cu130` on CUDA 13. Apple silicon users should follow the
[MPS installation guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
> **On an Apple Silicon Mac?** FastVideo runs FastWan text-to-video natively
> through an MLX runtime — a 5-second 480p clip generated locally, no cloud,
> no discrete GPU. Install with `uv pip install -e '.[mlx]'` and follow the
> [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
> **On an NVIDIA DGX Spark (GB10 / ARM64 + CUDA 13)?** There's no prebuilt ARM wheel for the FastVideo CUDA kernel, so it's an editable from-source install (`UV_TORCH_BACKEND=cu130 uv pip install -e .`, which compiles that kernel for you) rather than `UV_TORCH_BACKEND=cu130 uv pip install fastvideo`. A compatible prebuilt ARM64 FlashAttention wheel is available separately. Follow the [DGX Spark install guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/spark/).
+1 -1
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
-52
View File
@@ -1,52 +0,0 @@
{
"recipes": [
{
"id": "fastwan21-t2v",
"task": "Text to video",
"label": "FastWan2.1 1.3B (distilled + VSA)",
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
},
{
"id": "wan22-t2v",
"task": "Text to video",
"label": "Wan2.2 A14B",
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2.py",
"command": "python examples/inference/basic/basic_wan2_2.py"
},
{
"id": "wan21-i2v",
"task": "Image to video",
"label": "Wan2.1 14B 480P",
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
"source": "scripts/inference/inference_wan_i2v.yaml",
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
},
{
"id": "turbowan22-i2v",
"task": "Image to video",
"label": "TurboWan2.2 A14B",
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
},
{
"id": "wan22-ti2v",
"task": "Text or image to video",
"label": "Wan2.2 TI2V 5B",
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
},
{
"id": "matrix-game-2",
"task": "Interactive world",
"label": "Matrix Game 2.0",
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
"source": "examples/inference/basic/basic_matrixgame2.py",
"command": "python examples/inference/basic/basic_matrixgame2.py"
}
]
}
-60
View File
@@ -1,60 +0,0 @@
(() => {
let recipesPromise;
const loadRecipes = (url) => {
recipesPromise ||= fetch(url).then((response) => {
if (!response.ok) throw new Error(`HTTP ${response.status}`);
return response.json();
});
return recipesPromise;
};
const init = () => {
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
if (root.dataset.initialized) return;
root.dataset.initialized = "true";
const select = root.querySelector("[data-cookbook-recipe]");
const model = root.querySelector("[data-cookbook-model]");
const source = root.querySelector("[data-cookbook-source]");
const command = root.querySelector("[data-cookbook-command]");
const status = root.querySelector("[data-cookbook-status]");
try {
const { recipes } = await loadRecipes(root.dataset.recipes);
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
const groups = new Map();
select.replaceChildren();
recipes.forEach((recipe) => {
if (!groups.has(recipe.task)) {
const group = document.createElement("optgroup");
group.label = recipe.task;
groups.set(recipe.task, group);
select.append(group);
}
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
});
const render = () => {
const recipe = byId.get(select.value);
model.textContent = recipe.model;
source.textContent = recipe.source;
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
command.textContent = recipe.command;
status.textContent = `${recipe.label} selected.`;
};
select.addEventListener("change", render);
select.disabled = false;
render();
} catch (error) {
status.textContent = "Recipes could not be loaded. Use the examples link below.";
console.error("Failed to load FastVideo cookbook recipes", error);
}
});
};
if (window.document$) window.document$.subscribe(init);
else document.addEventListener("DOMContentLoaded", init);
})();
-40
View File
@@ -42,46 +42,6 @@ img {
margin: 0 auto;
}
.cookbook-picker {
padding: 1rem;
border: 0.05rem solid var(--md-default-fg-color--lightest);
border-radius: 0.2rem;
}
.cookbook-picker select {
width: 100%;
padding: 0.6rem;
color: var(--md-default-fg-color);
background: var(--md-default-bg-color);
border: 0.05rem solid var(--md-default-fg-color--lighter);
border-radius: 0.2rem;
}
.cookbook-picker dl {
display: grid;
grid-template-columns: max-content 1fr;
gap: 0.25rem 1rem;
}
.cookbook-picker dt {
font-weight: 700;
}
.cookbook-picker dd {
margin: 0;
min-width: 0;
overflow-wrap: anywhere;
}
.cookbook-picker__status {
position: absolute;
width: 1px;
height: 1px;
overflow: hidden;
clip: rect(0, 0, 0, 0);
white-space: nowrap;
}
.md-typeset .copy-page-button.md-button {
float: right;
margin: 0 0 1rem 1rem;
-44
View File
@@ -1,44 +0,0 @@
# Inference Cookbook
Choose a complete recipe maintained in the FastVideo repository. Each command
runs its checked-in source directly, so coupled model, GPU, offload, and
attention settings do not drift into unsupported combinations.
The commands expect a local clone:
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git
cd FastVideo
```
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
<label for="cookbook-recipe"><strong>Recipe</strong></label>
<select id="cookbook-recipe" data-cookbook-recipe disabled>
<option>Loading recipes…</option>
</select>
<dl>
<dt>Model</dt>
<dd data-cookbook-model>Loading…</dd>
<dt>Source</dt>
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
</dl>
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
<noscript>
JavaScript is needed for the recipe picker. Browse the
<a href="../inference/examples/examples_inference_index/">inference examples</a>
instead.
</noscript>
</div>
## Customize a recipe
Start from the checked-in source, then change only the settings your model
supports:
- [Configuration](../inference/configuration.md) covers the Python and CLI
config surfaces.
- [Optimizations](../inference/optimizations.md) covers attention backends,
compilation, and memory tradeoffs.
- [Support matrix](../inference/support_matrix.md) lists supported models and
optimizations.
-128
View File
@@ -1,128 +0,0 @@
# Fast mode (RIFE) — Apple Silicon
`--fast` makes local generation ~2.7× faster by **generating fewer frames and
interpolating the rest** with an Apple-Silicon-native RIFE model, instead of
denoising every frame. Video-diffusion denoise is dominated by self-attention,
which is O(tokens²); halving the frames cuts the token count ~2× and the denoise
compute ~3.7×, so the wall-clock drops far more than 2×. RIFE (which estimates
its own optical flow — no motion vectors needed) fills the dropped frames back
in for ~1.4 s, and a light unsharp pass counters its softening.
Measured on the 1.3B INT8 QAD model (fox, 480×832×81, M4): generate 41 + RIFE→81
runs in ~35 s of denoise vs ~90 s full, at reconstruction MS-SSIM **0.97**.
Reproduce with `python -m fastvideo.benchmarks.eval_metalfx_rife --mode int8`.
> **Note:** Apple's *MetalFX* frame interpolation is **not** usable here — it
> requires game-engine motion vectors + depth, which diffusion output lacks. We
> use the video-native **`rife-mlx`** model instead (Metal-backed, torch-free).
## Install
```bash
uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
```
## Use
```bash
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--mlx-checkpoint <FastWan2.1-T2V-1.3B-INT8-QAD> \
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
--num-frames 81 --fast \
--output-path video_samples/fox_fast.mp4
```
`--num-frames` stays the *target* length; fast mode generates the smallest
VAE-aligned keyframe count that RIFE can interpolate to that target.
| Flag | Default | Meaning |
|---|---|---|
| `--fast` / `--no-fast` | off | enable fast mode |
| `--fast-factor` | 2 | generate 1/factor of the frames (2 = half) |
| `--fast-sharpen` | 0.6 | light unsharp strength to counter RIFE softness (0 disables) |
Fast mode composes with everything else (`--mlx-quantization int8`,
`--mlx-compile`, TAEHV vs `--decode-backend wan-vae`). Keep `--fast-factor` at 2
for quality — larger temporal gaps are where RIFE starts inventing motion.
## Spatial fast mode (`--fast-spatial`)
The spatial twin of `--fast`: instead of dropping frames, drop pixels. Denoise
*and decode* at `height/width // fast-spatial-scale`, then resample the decoded
frames up to the requested size. Self-attention is O(tokens²), so halving each
spatial axis cuts the token count 4× and the denoise time far more than that —
measured on the 1.3B INT8 QAD model at 480×832×81, M4 Max: **86.1 s → 10.3 s**
of denoise. It composes with `--fast`; both together run the same clip in
**4.5 s** of denoise.
```bash
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
--height 480 --width 832 --num-frames 81 --fast-spatial \
--output-path video_samples/fox_fast_spatial.mp4
```
| Flag | Default | Meaning |
|---|---|---|
| `--fast-spatial` / `--no-fast-spatial` | off | enable spatial fast mode |
| `--fast-spatial-scale` | 2 | denoise at 1/scale of each spatial axis |
| `--fast-spatial-upsample-mode` | `lanczos` | pixel interpolation kernel (`lanczos`, `cubic`, `bilinear`, `nearest`) |
| `--fast-spatial-sharpen` | 0.4 | light unsharp strength to counter resampling softness (0 disables) |
### The upsample must happen in pixel space
This is the one thing to get right. The obvious implementation — bilinearly
upsample the finished latents and decode at the target size — **does not work**,
and produces a distinctive failure: correct composition and silhouette under a
smeared, hazy veil, with ringing along strong edges.
A Wan latent cell is a *learned code* for an 8×8 (Wan2.1) or 16×16 (Wan2.2)
pixel block, not a low-pass sample of the image. The average of two adjacent
codes is not the code of the averaged blocks; it is a vector the decoder was
never trained on. Measured on Wan2.1-1.3B at 480×832, a 2× bilinear latent
upsample destroys **62%** of the latent's high-frequency energy while leaving
its overall magnitude intact — exactly the signature of that veil. At Wan2.2-5B
the same operation degrades to black or noise.
Decoded RGB frames have no such problem: an image *is* a sampled 2-D signal, so
Lanczos interpolation is the operation it was defined for. The result is soft —
it carries stage-1's real detail budget and no more — but clean and coherent.
`--refine` gets away with a latent-space upsample only because a second DMD pass
re-denoises the hand-off; spatial fast mode passes the latent straight to the
decoder, so it cannot.
## Refine (`--refine`) stage-2 timesteps
`--refine` hands stage 1 to stage 2 as `(1 - sigma) * upsampled + sigma * noise`,
where `sigma` comes from the *first* stage-2 timestep. FastWan's DMD grid opens
at `t=1000`, which is `sigma == 1` exactly — so a stage-2 grid that starts there
weights the stage-1 result at zero and refine silently degrades into a plain
full-resolution run at twice the cost.
Left unset, `--refine-dmd-denoising-steps` now derives the stage-2 grid from the
stage-1 one with leading full-noise steps dropped (`1000,757,522` → `757,522`).
That keeps the pass on timesteps the distilled student was trained on while
letting stage-1 structure through: hand-off `sigma = 0.757`, stage-1 weight
`0.243`. Passing a grid that starts at full noise is now an error rather than a
silently wasted pass.
The run prints the resolved hand-off so it is visible:
```
[refine] stage-2 hand-off sigma=0.7568 (stage-1 weight 0.2432)
```
There is a trade-off in choosing that grid. Later start = more of the draft
survives, but fewer stage-2 steps. On Wan2.1 the default `757,522` gives weight
0.243 with two steps; `--refine-dmd-denoising-steps 522` gives weight 0.478 with
one. `--refine-sigma` decouples the noise level from the timestep entirely — it
logs a warning, because the DiT is then told a timestep that does not match the
noise it receives.
**Wan2.2-5B has a lower ceiling.** Its warped schedule maps `1000,757,522` to
sigmas `1.000, 0.940, 0.845`, so the best available stage-1 weight is **0.060**
(vs 0.243 at 1.3B). Un-warped (`--no-warp`) the same grid gives `1.000, 0.757,
0.522` and a weight of 0.243 — but warping is what matches the FastVideo
sampling schedule, so turning it off changes the timesteps the distilled student
sees. Which is better at 5B is unresolved and needs a run on real 5B weights.
@@ -76,10 +76,6 @@ surfaces:
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
-37
View File
@@ -2,7 +2,6 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import json
import os
import re
from dataclasses import dataclass, field
@@ -20,40 +19,6 @@ GENERATED_DOC_PREFIXES = (
"training/examples/",
"distillation/examples/",
)
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
COOKBOOK_SOURCE_ROOTS = (
ROOT_DIR / "examples/inference",
ROOT_DIR / "scripts/inference",
)
def validate_cookbook() -> None:
"""Keep cookbook entries tied to checked-in runnable sources."""
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
if not isinstance(recipes, list) or not recipes:
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
seen: set[str] = set()
for recipe in recipes:
required = ("id", "task", "label", "model", "source", "command")
missing = {key for key in required if not recipe.get(key)}
if missing:
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
if recipe["id"] in seen:
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
seen.add(recipe["id"])
source = (ROOT_DIR / recipe["source"]).resolve()
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
if not source.is_file():
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
source_text = source.read_text(encoding="utf-8")
if recipe["model"] not in source_text:
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
if recipe["source"] not in recipe["command"]:
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
def fix_case(text: str) -> str:
@@ -571,7 +536,6 @@ def on_pre_build(config, **kwargs):
MkDocs hook to generate examples before building the documentation.
This function is called automatically by MkDocs' native hook system.
"""
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
@@ -585,7 +549,6 @@ def on_page_context(context, page, **kwargs):
if __name__ == "__main__":
validate_cookbook()
print("Generating example documentation...")
generate_examples(generate_main_index=True)
print("Example documentation generated successfully!")
+2 -3
View File
@@ -65,7 +65,6 @@ uv pip install flash-attn --no-build-isolation -v
## Next Steps
- [Quick Start](quick_start.md) - Generate your first video
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
+4 -6
View File
@@ -49,12 +49,10 @@ brew install ffmpeg
### Installation
FastWan's native Apple Silicon runtime requires the `mlx` extra.
#### With uv (recommended)
```bash
uv pip install "fastvideo[mlx]"
uv pip install fastvideo
```
#### With Conda environment (alternative)
@@ -62,7 +60,7 @@ uv pip install "fastvideo[mlx]"
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
```bash
uv pip install "fastvideo[mlx]"
uv pip install fastvideo
```
### Installation from Source
@@ -78,13 +76,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
Basic installation:
```bash
uv pip install -e ".[mlx]"
uv pip install -e .
```
Alternative with Conda environment:
```bash
uv pip install -e ".[mlx]"
uv pip install -e .
```
## Development Environment Setup
+49 -9
View File
@@ -23,21 +23,61 @@ Also optionally install flash-attn:
uv pip install flash-attn --no-build-isolation -v
```
## Choose a maintained recipe
## Basic Usage
The cookbook selects complete, checked-in recipes instead of mixing model,
parallelism, offload, and attention settings independently.
### Text-to-Video Generation
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
```python
from fastvideo import VideoGenerator
!!! tip "Need more control?"
Start from a maintained recipe, then use the
[configuration](../inference/configuration.md) and
[optimization](../inference/optimizations.md) guides for supported changes.
def main():
# Create a video generator with a pre-trained model
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1, # Adjust based on your hardware
)
# Define a prompt for your video
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
# Generate the video
video = generator.generate_video(
prompt,
output_path="my_videos/", # Controls where videos are saved
save_video=True
)
if __name__ == '__main__':
main()
```
### Image-to-Video Generation
```python
from fastvideo import VideoGenerator, SamplingParam
def main():
# Create the generator
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
# Set up parameters with an initial image
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.num_frames = 107
# Generate video based on the image
prompt = "A photograph coming to life with gentle movement"
generator.generate_video(prompt, sampling_param=sampling_param,
output_path="my_videos/",
save_video=True)
if __name__ == '__main__':
main()
```
## Next Steps
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
- [Installation Guide](installation.md) - Detailed installation instructions
- [Configuration](../inference/configuration.md) - Learn about configuration options
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
-11
View File
@@ -178,17 +178,6 @@ optimizations: absence means **untested**, not incompatible.
| Matrix Game 3.0 Base Distilled | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | 720x1280 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
## Apple Silicon native runtime
| Release path | Model | Mode | Validated hardware | Status |
| --- | --- | --- | --- | --- |
| MLX FastWan T2V | FastWan-QAD-INT8-1.3B `[release model ID pending]` | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB unified-memory class, MLX 0.31.2 | Release candidate; requires release-owner visual sign-off |
This is a text-to-video-only source-install release. It is validated on the
hardware listed above; MLX allocator caps are not evidence of support for a
physical 16 GB Mac. See [Apple Silicon FastWan](../getting_started/installation/mps.md)
for the supported command and release gates.
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
+27 -6
View File
@@ -33,13 +33,34 @@ For the typed config/request path added during the inference API refactor:
python examples/inference/basic/basic_dmd_new_api.py
```
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
```
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
```
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
## LTX-2.5 single-file bundles
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
These examples load LTX-2.5 from a single-file bundle: one `.safetensors`
file carrying every component (transformer, video VAE, audio VAE, vocoder,
text projection). (The Hugging Face release ships as a split pack — one
file per component — which loads through the standard repo path instead.)
Pass the bundle FILE path directly:
```bash
python examples/inference/basic/basic_ltx2_5_distilled.py \
--model-path /path/to/bundle.safetensors \
--gemma-root /path/to/gemma_root
```
| bundle variant | preset | sampling defaults |
|---|---|---|
| distilled | `ltx2_distilled` | one stage, 8 steps, CFG 1.0 (no CFG/STG, no spatial upscaler) |
| sft / base | `ltx2_base` | 40 steps, CFG 3.0 with STG |
- The variant is read from the bundle header when declared; otherwise the
`distilled` token in the file name decides.
- The Gemma text encoder is NOT in the bundle. Declare its root with
`--gemma-root` (or `FASTVIDEO_LTX_ENCODER_ROOT`), and use the root shipped
with the same transformer variant — the roots differ in prompt templating.
- Decoder caveat: bundles that declare a diffusion VAE decoder
(`CausalDiffusionVAE`) are not supported yet; loading them currently
fails at VAE build time. Bundles with the classic `CausalVideoAutoencoder`
run end to end.
## Basic Walkthrough
-180
View File
@@ -1,180 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
sampler's shift-12 schedule instead of the base model's 50 steps, generating
synchronized video and audio in one pipeline call.
The student was trained with block-sparse video attention (VSA, 64-token
tiles) and its checkpoint carries the trained sparse-gate parameters
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
dense (every tile is selected); raise the sparsity for additional speedup.
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
# The HF repo is private while the MiniMax H3 Community License review
# completes; until it flips public, pass --model-path with a local
# snapshot of the release instead (e.g. the team export at
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", default="outputs/fasth3")
parser.add_argument("--height", type=int, default=768)
parser.add_argument("--width", type=int, default=1344)
parser.add_argument("--num-frames", type=int, default=124)
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
# default here is 5. Other grids are off-distribution.
parser.add_argument("--steps",
type=int,
default=5,
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
"forwards. 5 (default) is the distilled 4-forward grid")
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--vsa-sparsity",
type=float,
default=0.0,
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
"exactly dense attention; the student was trained at 0.9")
# 64 is the trained contract: the student was TRAINED with 64-token
# (4,4,4) tiles, and its to_gate_compress gates were learned against
# pooling at that granularity — keep 64 unless you are ablating.
parser.add_argument("--vsa-tile-size",
type=int,
choices=(64, 256),
default=64,
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
"geometry for ablations")
parser.add_argument("--vsa-kernel",
choices=("triton", "sm100a"),
default="triton",
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
"fastvideo-kernel build that carries the extension; if a precondition fails at "
"run time the attention layer logs one warning and falls back to Triton. Only "
"meaningful with --vsa-tile-size 64")
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--repeats",
type=int,
default=1,
help="generate N times; with --torch-compile the first run pays "
"compilation, so steady-state is the last repeat")
return parser.parse_args()
def main() -> None:
args = parse_args()
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
if args.vsa_kernel == "sm100a":
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
# before the pipeline boots so spawned GPU workers inherit it. The
# kernel is forward-only and inference runs under no-grad, so every
# denoising forward qualifies for the CUDA route.
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
# Boot-time run configuration, folded into FastVideoArgs (the same route
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
# - attention_backend: the checkpoint carries trained to_gate_compress
# gates, which only exist under the VSA-H3 backend — a dense-backend
# load would reject them as unexpected weights. Layers that do not
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
# branch pools per tile, and the gates were trained at 64 tokens/tile.
experimental: dict[str, object] = {
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
"VSA_tile_size": args.vsa_tile_size,
}
if args.vsa_sparsity > 0.0:
experimental["VSA_sparsity"] = args.vsa_sparsity
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
pipeline=PipelineSelection(experimental=experimental),
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=args.num_gpus > 1,
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
offload=OffloadConfig(
dit=False,
dit_layerwise=False,
text_encoder=True,
vae=True,
pin_cpu_memory=False,
),
compile=CompileConfig(
enabled=args.torch_compile,
mode=args.compile_mode,
),
),
))
try:
request = GenerationRequest(
prompt=args.prompt,
negative_prompt="",
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=args.num_frames,
fps=24,
num_inference_steps=args.steps,
# the base model is guidance-distilled; the student inherits it
guidance_scale=1.0,
batch_cfg=False,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output_dir / "fasth3.mp4"),
save_video=True,
return_frames=False,
),
)
result = generator.generate(request)
print(f"Output written to: {result.video_path}")
if result.generation_time is not None:
# machine-readable: benchmark harnesses parse this line to separate
# generation from model-load time (last occurrence = steady state)
print(f"Generation time: {result.generation_time:.2f}s")
for _ in range(args.repeats - 1):
result = generator.generate(request)
if result.generation_time is not None:
print(f"Generation time: {result.generation_time:.2f}s")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+103
View File
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2.5 (sft/base) text-to-video from a single-file bundle.
Loads LTX-2.5 from a single-file bundle: ONE ``.safetensors`` file carrying
every component (transformer, video VAE, audio VAE, vocoder, text
projection), so ``--model-path`` takes the bundle FILE, not a repo
directory. (The split-pack Hugging Face release loads through the standard
repo path instead.)
The Gemma text encoder is NOT in the bundle: pass its root directory with
``--gemma-root`` (or set ``FASTVIDEO_LTX_ENCODER_ROOT``). Use the encoder
root shipped WITH this transformer variant -- the published roots share
weights but differ in prompt templating, so a mismatched root silently
changes prompting.
A non-distilled bundle resolves to the standard ``ltx2_base`` preset
(40 steps, CFG 3.0 with STG); run without sampling flags to use it as-is.
For the 8-step distilled recipe, see ``basic_ltx2_5_distilled.py``.
Audio: the bundle carries an audio VAE + vocoder, so generated videos get
an audio track. A bundle that declares no audio decoder skips audio
decoding and still produces video.
Decoder caveat: bundles that declare a *diffusion* VAE decoder
(``CausalDiffusionVAE``) are not implemented yet -- loading such a bundle
currently fails at VAE build time with an unsupported-architecture error
(no classic-decoder fallback is wired). Bundles with the classic
``CausalVideoAutoencoder`` decode end to end.
"""
import argparse
import os
from fastvideo import VideoGenerator
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
# Sampling flags default to None: unset flags are NOT passed to
# `generate_video`, so the bundle's preset supplies them via
# `SamplingParam.from_pretrained` (ltx2_base: 40 steps, cfg 3.0, 512x768,
# 121 frames).
_SAMPLING_FLAGS = ("height", "width", "num_frames", "num_inference_steps", "guidance_scale", "seed")
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="LTX-2.5 (sft/base) inference from a single-file bundle.")
parser.add_argument(
"--model-path",
required=True,
help="Path to the LTX-2.5 bundle (.safetensors FILE).",
)
parser.add_argument(
"--gemma-root",
default=None,
help="Gemma text-encoder root directory. Must be the root paired with "
"this transformer variant; the roots differ in prompt templating.",
)
parser.add_argument("--prompt", default=PROMPT)
parser.add_argument("--output-path", default="outputs_video/ltx2_5/output_ltx2_5_t2v.mp4")
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--height", type=int, default=None)
parser.add_argument("--width", type=int, default=None)
parser.add_argument("--num-frames", type=int, default=None)
parser.add_argument("--num-inference-steps", type=int, default=None)
parser.add_argument("--guidance-scale", type=float, default=None)
parser.add_argument("--seed", type=int, default=None)
return parser.parse_args(argv)
def sampling_overrides(args: argparse.Namespace) -> dict:
"""Only sampling flags the user actually passed; the preset supplies the rest."""
return {name: getattr(args, name) for name in _SAMPLING_FLAGS if getattr(args, name) is not None}
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
if args.gemma_root:
os.environ["FASTVIDEO_LTX_ENCODER_ROOT"] = args.gemma_root
generator = VideoGenerator.from_pretrained(
args.model_path,
num_gpus=args.num_gpus,
)
generator.generate_video(
prompt=args.prompt,
output_path=args.output_path,
save_video=True,
**sampling_overrides(args),
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
"""LTX-2.5 distilled text-to-video from a single-file bundle.
Loads LTX-2.5 from a single-file bundle: ONE ``.safetensors`` file carrying
every component (transformer, video VAE, audio VAE, vocoder, text
projection), so ``--model-path`` takes the bundle FILE, not a repo
directory. (The split-pack Hugging Face release loads through the standard
repo path instead.)
The Gemma text encoder is NOT in the bundle: pass its root directory with
``--gemma-root`` (or set ``FASTVIDEO_LTX_ENCODER_ROOT``). Use the encoder
root shipped WITH this transformer variant -- the published roots share
weights but differ in prompt templating, so a mismatched root silently
changes prompting.
Distilled recipe: ONE stage -- 8 steps on the distilled sigma schedule,
CFG 1.0, no CFG/STG, and no spatial upscaler needed. All of it comes from
the ``ltx2_distilled`` preset the bundle resolves to; run without sampling
flags to use it as-is.
Audio: the bundle carries an audio VAE + vocoder, so generated videos get
an audio track. A bundle that declares no audio decoder skips audio
decoding and still produces video.
Decoder caveat: bundles that declare a *diffusion* VAE decoder
(``CausalDiffusionVAE``) are not implemented yet -- loading such a bundle
currently fails at VAE build time with an unsupported-architecture error
(no classic-decoder fallback is wired). Bundles with the classic
``CausalVideoAutoencoder`` decode end to end.
"""
import argparse
import os
from fastvideo import VideoGenerator
PROMPT = ("A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic.")
# Sampling flags default to None: unset flags are NOT passed to
# `generate_video`, so the bundle's preset supplies them via
# `SamplingParam.from_pretrained` (ltx2_distilled: 8 steps, cfg 1.0,
# 1024x1536, 121 frames).
_SAMPLING_FLAGS = ("height", "width", "num_frames", "num_inference_steps", "guidance_scale", "seed")
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="LTX-2.5 distilled inference from a single-file bundle.")
parser.add_argument(
"--model-path",
required=True,
help="Path to the LTX-2.5 distilled bundle (.safetensors FILE).",
)
parser.add_argument(
"--gemma-root",
default=None,
help="Gemma text-encoder root directory. Must be the root paired with "
"this transformer variant; the roots differ in prompt templating.",
)
parser.add_argument("--prompt", default=PROMPT)
parser.add_argument("--output-path", default="outputs_video/ltx2_5/output_ltx2_5_distilled_t2v.mp4")
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--height", type=int, default=None)
parser.add_argument("--width", type=int, default=None)
parser.add_argument("--num-frames", type=int, default=None)
parser.add_argument("--num-inference-steps", type=int, default=None)
parser.add_argument("--guidance-scale", type=float, default=None)
parser.add_argument("--seed", type=int, default=None)
return parser.parse_args(argv)
def sampling_overrides(args: argparse.Namespace) -> dict:
"""Only sampling flags the user actually passed; the preset supplies the rest."""
return {name: getattr(args, name) for name in _SAMPLING_FLAGS if getattr(args, name) is not None}
def main(argv: list[str] | None = None) -> None:
args = parse_args(argv)
if args.gemma_root:
os.environ["FASTVIDEO_LTX_ENCODER_ROOT"] = args.gemma_root
generator = VideoGenerator.from_pretrained(
args.model_path,
num_gpus=args.num_gpus,
)
generator.generate_video(
prompt=args.prompt,
output_path=args.output_path,
save_video=True,
**sampling_overrides(args),
)
generator.shutdown()
if __name__ == "__main__":
main()
@@ -27,8 +27,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
@@ -27,8 +27,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
@@ -24,8 +24,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
# --model-path noctuashap/MiniMax-H3-pruned-r16
# (or a local dir produced by
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
# adaln_rank is read from the checkpoint config; no other flags needed.
# Rank-reduced checkpoints are inference-only: training needs the
# full-rank release.
@@ -1,46 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Tiny MLX RIFE frame-interpolation smoke test."""
from __future__ import annotations
import argparse
import time
import numpy as np
from fastvideo.mlx_runtime.rife_interp import interpolate, load_model
def main() -> None:
parser = argparse.ArgumentParser(
description="MLX RIFE 4.25 frame interpolation smoke test."
)
parser.add_argument(
"--self-test",
action="store_true",
help="Run a tiny two-frame interpolation test.",
)
args = parser.parse_args()
if not args.self_test:
raise SystemExit("Nothing to do; pass --self-test")
frame0 = np.zeros((64, 96, 3), dtype=np.uint8)
frame1 = np.zeros((64, 96, 3), dtype=np.uint8)
frame1[:, :, 0] = 255
start = time.perf_counter()
model = load_model()
load_s = time.perf_counter() - start
start = time.perf_counter()
frames = interpolate([frame0, frame1], factor=2, model=model)
interp_s = time.perf_counter() - start
assert len(frames) == 3
assert frames[1].shape == frame0.shape
assert frames[1].dtype == np.uint8
print(
"MLX RIFE self-test passed: "
f"load_s={load_s:.3f} interp_s={interp_s:.3f} shape={frames[1].shape}"
)
if __name__ == "__main__":
main()
@@ -1,510 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""End-to-end Wan2.2-TI2V-5B generation on Apple Silicon (MLX DiT + MLX TAEHV).
Pipeline: torch/MPS UMT5 encode (shared with 1.3B) → MLXWan22DiT 3-step DMD
(warped schedule, flow_shift=5) → MLX TAEHV decode (taew2_2.pth). Fully MLX
on the heavy DiT + decode path.
PYTHONPATH=$PWD python examples/inference/basic/mlx_wan22_generate.py \
--prompt "A red fox trotting through a snowy pine forest at golden hour" \
--output-path video_samples/demo_5b/fox_5b_mlx.mp4
Decoder backends: ``taehv`` (default, MLX, ~seconds), ``taehv-torch`` (parity),
``wan-vae`` (full AutoencoderKLWan on MPS, slow).
"""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
import numpy as np
from fastvideo.mlx_runtime.fast_spatial import DEFAULT_FAST_SPATIAL_SHARPEN
from fastvideo.mlx_runtime.frame_upsample import DEFAULT_PIXEL_UPSAMPLE_MODE, PIXEL_UPSAMPLE_MODES
from fastvideo.mlx_runtime.memory import cleanup_mlx
from fastvideo.mlx_runtime.prompt_cache import (
fingerprint_digest,
load_prompt_cache,
save_prompt_cache,
text_encoder_fingerprint,
)
from fastvideo.mlx_runtime.rife_interp import aligned_keyframe_count
FASTWAN21_MODEL_ID = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
FASTWAN22_MODEL_ID = "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
DEFAULT_HEIGHT = 448
DEFAULT_WIDTH = 832
DEFAULT_NUM_FRAMES = 121
def _resolve_model_paths(
*,
text_encoder_root: Path | None,
dit_checkpoint: Path | None,
dit_config: Path | None,
vae_root: Path | None,
mlx_checkpoint: Path | None,
decode_backend: str,
) -> tuple[Path, Path | None, Path | None, Path | None]:
"""Download only the missing assets required by the selected Wan2.2 path."""
from huggingface_hub import snapshot_download
if text_encoder_root is None:
text_encoder_root = Path(snapshot_download(
FASTWAN21_MODEL_ID,
allow_patterns=["tokenizer/*", "text_encoder/*"],
))
if mlx_checkpoint is None and (dit_checkpoint is None or dit_config is None):
patterns = []
if dit_checkpoint is None:
patterns.append("transformer/diffusion_pytorch_model.safetensors")
if dit_config is None:
patterns.append("transformer/config.json")
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=patterns))
dit_checkpoint = dit_checkpoint or model_root / "transformer/diffusion_pytorch_model.safetensors"
dit_config = dit_config or model_root / "transformer/config.json"
if decode_backend == "wan-vae" and vae_root is None:
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=["vae/*"]))
vae_root = model_root / "vae"
return text_encoder_root, dit_checkpoint, dit_config, vae_root
def _prompt_cache_fingerprint(
*,
prompt: str,
prompt_used: str,
enhance_prompt: bool,
enhance_prompt_backend: str,
text_encoder_root: Path,
max_sequence_length: int,
dtype: str,
) -> dict[str, object]:
return {
"prompt": prompt,
"prompt_used": prompt_used,
"enhance_prompt": enhance_prompt,
"enhance_prompt_backend": enhance_prompt_backend,
"text_encoder": text_encoder_fingerprint(text_encoder_root),
"max_sequence_length": max_sequence_length,
"dtype": dtype,
}
def _default_prompt_cache_path(fingerprint: dict[str, object]) -> Path:
"""Content-addressed default cache file for a prompt fingerprint.
The Wan2.1 entrypoint caches prompt embeddings by default; this one only
did so when handed an explicit ``--prompt-embeds-cache`` path, so every 5B
run paid a full UMT5 encode (~45s on an M4 Max) even for a repeat prompt.
The fingerprint already covers everything that changes the embedding, so
hash it for the filename.
"""
digest = fingerprint_digest(fingerprint)[:32]
return Path.home() / ".cache" / "fastvideo" / "prompt_embeds" / f"wan22_{digest}.npy"
def main() -> None:
parser = argparse.ArgumentParser(
description="MLX Wan2.2-5B T2V (encode → DiT DMD → TAEHV/VAE decode)"
)
parser.add_argument(
"--prompt",
default="A red fox trotting through a snowy pine forest at golden hour, cinematic",
)
parser.add_argument(
"--output-path",
type=Path,
default=Path("video_samples/demo_5b/fox_5b_mlx.mp4"),
)
parser.add_argument(
"--text-encoder-root",
type=Path,
default=None,
help="Root with text_encoder/ + tokenizer/",
)
parser.add_argument(
"--prompt-embeds-cache",
type=Path,
default=None,
help="Explicit .npy UMT5 embedding cache file. Overrides the automatic "
"content-addressed cache (--prompt-cache).",
)
parser.add_argument(
"--prompt-cache",
action=argparse.BooleanOptionalAction,
default=True,
help="Cache prompt embeddings under ~/.cache/fastvideo/prompt_embeds so "
"repeat runs skip the text encoder entirely. Default: on.",
)
parser.add_argument(
"--text-encoder-device",
choices=("auto", "cpu", "mps"),
default="cpu",
help="Device for UMT5 encoding. CPU is safest beside the 5B MLX DiT.",
)
parser.add_argument(
"--enhance-prompt",
action="store_true",
help="Apply deterministic local cinematic prompt enrichment before UMT5.",
)
parser.add_argument(
"--enhance-prompt-backend",
choices=("template",),
default="template",
help="Prompt enrichment backend.",
)
parser.add_argument(
"--dit-checkpoint",
type=Path,
default=None,
)
parser.add_argument("--dit-config", type=Path, default=None)
parser.add_argument(
"--mlx-checkpoint",
type=Path,
default=None,
help="Pre-quantized MLX DiT checkpoint directory. Rewrapped with Wan2.2 per-token conditioning.",
)
parser.add_argument("--vae-root", type=Path, default=None)
parser.add_argument("--height", type=int, default=DEFAULT_HEIGHT)
parser.add_argument("--width", type=int, default=DEFAULT_WIDTH)
parser.add_argument(
"--num-frames",
type=int,
default=DEFAULT_NUM_FRAMES,
help="Pixel frames (121 at 24fps = 5.04 seconds)",
)
parser.add_argument("--seed", type=int, default=1234)
parser.add_argument("--renoise-seed", type=int, default=0)
parser.add_argument("--fps", type=int, default=24)
parser.add_argument("--flow-shift", type=float, default=5.0)
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
parser.add_argument(
"--no-warp",
action="store_true",
help="Disable schedule warping (debug only).",
)
parser.add_argument(
"--fast",
action="store_true",
help="Generate fewer frames then RIFE-interpolate to --num-frames.",
)
parser.add_argument("--fast-factor", type=int, default=2)
parser.add_argument("--fast-sharpen", type=float, default=0.6)
parser.add_argument(
"--fast-spatial",
action="store_true",
help="Denoise and decode at reduced spatial resolution, then resample "
"the decoded frames up to the target size.",
)
parser.add_argument("--fast-spatial-scale", type=int, default=2)
parser.add_argument(
"--fast-spatial-upsample-mode",
choices=PIXEL_UPSAMPLE_MODES,
default=DEFAULT_PIXEL_UPSAMPLE_MODE,
)
parser.add_argument("--fast-spatial-sharpen", type=float, default=DEFAULT_FAST_SPATIAL_SHARPEN)
parser.add_argument(
"--refine",
action="store_true",
help="Two-pass DMD: coarse denoise, upsample/re-noise, full-res denoise.",
)
parser.add_argument("--refine-scale", type=int, default=2)
parser.add_argument(
"--refine-upsample-mode",
choices=("bilinear", "nearest"),
default="bilinear",
)
parser.add_argument("--no-refine-add-noise", action="store_true")
parser.add_argument(
"--decode-backend",
choices=("taehv", "taehv-torch", "wan-vae"),
default="taehv",
)
parser.add_argument("--save-latents", type=Path, default=None)
parser.add_argument("--metrics-json", type=Path, default=None,
help="Write measured run metadata as JSON for reports or galleries.")
parser.add_argument(
"--compile",
action="store_true",
help="Compile the DiT forward with mx.compile; fallback to eager on failure.",
)
args = parser.parse_args()
if args.fast_factor < 2:
parser.error("--fast-factor must be at least 2")
# --fast-spatial used to be rejected here because it upsampled the completed
# 48-channel latent, which is out of distribution for the decoder and gave
# black or noisy video. The upsample now runs on decoded frames, so the
# latent never leaves the grid it was denoised on and the mode is usable.
if args.refine and args.fast_spatial:
print("[wan22] --refine takes precedence over --fast-spatial")
args.text_encoder_root, args.dit_checkpoint, args.dit_config, args.vae_root = _resolve_model_paths(
text_encoder_root=args.text_encoder_root,
dit_checkpoint=args.dit_checkpoint,
dit_config=args.dit_config,
vae_root=args.vae_root,
mlx_checkpoint=args.mlx_checkpoint,
decode_backend=args.decode_backend,
)
target_frames = args.num_frames
if args.fast:
args.num_frames = aligned_keyframe_count(target_frames, args.fast_factor)
print(
f"[wan22 fast] generating {args.num_frames} frames, "
f"RIFE {args.fast_factor}x -> {target_frames}"
)
import mlx.core as mx
import torch
from examples.inference.basic.mlx_wan_prompt_to_video import (
_postprocess_video,
encode_prompt,
make_rotary_embeddings,
)
from fastvideo.mlx_runtime.fast_spatial import plan_fast_spatial
from fastvideo.mlx_runtime.refine import (
default_refine_timesteps,
plan_refine_resolutions,
prepare_refine_latents,
)
from fastvideo.mlx_runtime.wan22 import (
mlx_wan22_dit_from_diffusers_safetensors,
mlx_wan22_dit_from_mlx_checkpoint,
)
from fastvideo.mlx_runtime.wan22_sample import build_wan22_dmd_schedule, sample_wan22_dmd
from fastvideo.mlx_runtime.wan_vae import decode_latents_to_video
if args.mlx_checkpoint is not None:
config = json.loads((args.mlx_checkpoint / "mlx_dit.json").read_text())["config"]
else:
config = json.loads(args.dit_config.read_text())
patch_size = tuple(config.get("patch_size", (1, 2, 2)))
if args.refine:
active_plan = plan_refine_resolutions(
height=args.height, width=args.width, num_frames=args.num_frames,
spatial_scale=args.refine_scale, vae_spatial_compression=16,
vae_temporal_compression=4, patch_size=patch_size, enabled=True,
)
spatial_mode = "refine"
elif args.fast_spatial:
fast_spatial_plan = plan_fast_spatial(
height=args.height, width=args.width, num_frames=args.num_frames,
spatial_scale=args.fast_spatial_scale, vae_spatial_compression=16,
vae_temporal_compression=4, patch_size=patch_size,
upsample_mode=args.fast_spatial_upsample_mode,
sharpen=args.fast_spatial_sharpen, enabled=True,
)
active_plan = fast_spatial_plan.plan
spatial_mode = "fast_spatial"
else:
active_plan = plan_refine_resolutions(
height=args.height, width=args.width, num_frames=args.num_frames,
spatial_scale=1, vae_spatial_compression=16, vae_temporal_compression=4,
patch_size=patch_size, enabled=False,
)
spatial_mode = "off"
lat_h, lat_w = active_plan.stage1_latent_height, active_plan.stage1_latent_width
lat_t = active_plan.latent_frames
in_ch = int(config["in_channels"])
print(f"[5B] latent {in_ch}x{lat_t}x{lat_h}x{lat_w}", flush=True)
total_start = time.perf_counter()
prompt_for_encode = args.prompt
enhance_backend = None
enhance_elapsed_s = 0.0
if args.enhance_prompt:
from fastvideo.mlx_runtime.prompt_enhance import enhance_prompt
enhancement = enhance_prompt(args.prompt, backend=args.enhance_prompt_backend)
prompt_for_encode = enhancement.enhanced
enhance_backend = enhancement.backend
enhance_elapsed_s = enhancement.elapsed_s
print(f"[enhance] backend={enhance_backend} in {enhance_elapsed_s:.2f}s", flush=True)
print(f"[enhance] prompt: {prompt_for_encode}", flush=True)
t0 = time.perf_counter()
prompt_cache_fingerprint = _prompt_cache_fingerprint(
prompt=args.prompt,
prompt_used=prompt_for_encode,
enhance_prompt=args.enhance_prompt,
enhance_prompt_backend=args.enhance_prompt_backend,
text_encoder_root=args.text_encoder_root,
max_sequence_length=512,
dtype="fp16",
)
prompt_cache_path = args.prompt_embeds_cache
if prompt_cache_path is None and args.prompt_cache:
prompt_cache_path = _default_prompt_cache_path(prompt_cache_fingerprint)
cached_embeds = load_prompt_cache(
prompt_cache_path,
prompt_cache_fingerprint,
)
if cached_embeds is not None:
embeds = torch.from_numpy(cached_embeds).contiguous()
else:
embeds = encode_prompt(
model_root=args.text_encoder_root,
prompt=prompt_for_encode,
max_sequence_length=512,
device_arg=args.text_encoder_device,
dtype_arg="fp16",
)
save_prompt_cache(
prompt_cache_path,
embeds.cpu().numpy(),
prompt_cache_fingerprint,
)
ehs = mx.array(embeds.numpy()).astype(mx.float16)
prompt_encode_s = time.perf_counter() - t0
print(f"[5B] prompt encoded {tuple(ehs.shape)} in {prompt_encode_s:.1f}s", flush=True)
t1 = time.perf_counter()
if args.mlx_checkpoint is not None:
dit = mlx_wan22_dit_from_mlx_checkpoint(
args.mlx_checkpoint,
compile=args.compile,
)
else:
dit = mlx_wan22_dit_from_diffusers_safetensors(
args.dit_checkpoint,
args.dit_config,
dtype="fp16",
compile=args.compile,
)
dit_load_s = time.perf_counter() - t1
print(f"[5B] DiT loaded in {dit_load_s:.1f}s", flush=True)
freqs = make_rotary_embeddings(config, latent_frames=lat_t, latent_height=lat_h, latent_width=lat_w)
gen = torch.Generator().manual_seed(args.seed)
noise = mx.array(
torch.randn(1, in_ch, lat_t, lat_h, lat_w, generator=gen, dtype=torch.float32).numpy()).astype(mx.float16)
steps = [int(s) for s in args.dmd_denoising_steps.split(",") if s.strip()]
t2 = time.perf_counter()
mx.reset_peak_memory()
latents = sample_wan22_dmd(
dit,
ehs,
noise,
freqs,
dmd_denoising_steps=steps,
flow_shift=args.flow_shift,
warp_denoising_step=not args.no_warp,
seed=args.renoise_seed,
)
if spatial_mode == "refine":
schedule, warped_steps = build_wan22_dmd_schedule(
steps, flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
)
# The grid opens at sigma == 1, where the hand-off
# `(1 - sigma) * upsampled + sigma * noise` weights stage 1 at zero and
# refine silently becomes a plain full-res run. Drop the leading
# full-noise steps so stage 1 actually reaches stage 2.
stage2_warped = default_refine_timesteps(schedule, warped_steps)
stage2_steps = steps[len(warped_steps) - len(stage2_warped):]
sigma = schedule.sigma_for(stage2_warped[0])
print(f"[5B refine] stage-2 steps={stage2_steps} sigma={sigma:.4f} "
f"(stage-1 weight {1.0 - sigma:.4f})", flush=True)
latents = prepare_refine_latents(
latents, scale=args.refine_scale, sigma=sigma,
add_noise_flag=not args.no_refine_add_noise,
upsample_mode=args.refine_upsample_mode, seed=args.renoise_seed + 1,
)
freqs_stage2 = make_rotary_embeddings(
config, latent_frames=lat_t,
latent_height=active_plan.stage2_latent_height,
latent_width=active_plan.stage2_latent_width,
)
latents = sample_wan22_dmd(
dit, ehs, latents, freqs_stage2, dmd_denoising_steps=stage2_steps,
flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
seed=args.renoise_seed + 2,
)
# spatial_mode == "fast_spatial" leaves the latents on the stage-1 grid;
# the resample happens after decode, in _postprocess_video.
denoise_s = time.perf_counter() - t2
peak = mx.get_peak_memory() / (1024**3)
print(f"[5B] denoise {len(steps)} steps in {denoise_s:.1f}s, peak {peak:.2f} GiB", flush=True)
latents_np = np.array(latents.astype(mx.float32))
if args.save_latents is not None:
args.save_latents.parent.mkdir(parents=True, exist_ok=True)
np.savez(args.save_latents, latents=latents_np, prompt=args.prompt, seed=args.seed)
print(f"[5B] wrote latents {args.save_latents}", flush=True)
if spatial_mode == "refine":
del freqs_stage2
del dit, latents, ehs, noise, freqs
cleanup_mlx()
metrics = decode_latents_to_video(
latents_np,
args.output_path,
fps=args.fps,
backend=args.decode_backend,
vae_dir=args.vae_root if args.decode_backend == "wan-vae" else None,
z_dim=in_ch,
)
# One h264 round-trip for both post-decode passes (see _postprocess_video).
rife_s = 0.0
rife_request = ({
"factor": args.fast_factor,
"target_frames": target_frames,
"sharpen": args.fast_sharpen,
} if args.fast else None)
spatial_request = fast_spatial_plan if spatial_mode == "fast_spatial" else None
if rife_request is not None or spatial_request is not None:
rife_start = time.perf_counter()
_postprocess_video(
video_path=args.output_path, fps=args.fps,
rife=rife_request, spatial=spatial_request,
)
rife_s = time.perf_counter() - rife_start
print(f"[5B] decoded via {metrics['backend']} in {metrics['decode_s']:.1f}s → {args.output_path}", flush=True)
summary = {
"output_path": str(args.output_path.resolve()),
"prompt": args.prompt,
"prompt_used": prompt_for_encode,
"enhance_prompt": args.enhance_prompt,
"enhance_backend": enhance_backend,
"enhance_elapsed_s": round(enhance_elapsed_s, 3),
"height": args.height,
"width": args.width,
"fps": args.fps,
"target_frames": target_frames,
"generated_frames": args.num_frames,
"seed": args.seed,
"renoise_seed": args.renoise_seed,
"dmd_denoising_steps": steps,
"flow_shift": args.flow_shift,
"warp": not args.no_warp,
"spatial_mode": spatial_mode,
"fast": args.fast,
"fast_factor": args.fast_factor if args.fast else None,
"fast_spatial_scale": args.fast_spatial_scale if args.fast_spatial else None,
"refine_scale": args.refine_scale if args.refine else None,
"decode_backend": args.decode_backend,
"prompt_encode_s": round(prompt_encode_s, 3),
"dit_load_s": round(dit_load_s, 3),
"denoise_s": round(denoise_s, 3),
"decode_s": round(metrics["decode_s"], 3),
"rife_s": round(rife_s, 3),
"wall_total_s": round(time.perf_counter() - total_start, 3),
"peak_gib": round(peak, 3),
"latent_shape": [in_ch, lat_t, lat_h, lat_w],
"stage2_latent_shape": [in_ch, lat_t, active_plan.stage2_latent_height, active_plan.stage2_latent_width],
"mlx_checkpoint": str(args.mlx_checkpoint.resolve()) if args.mlx_checkpoint else None,
}
if args.metrics_json is not None:
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
args.metrics_json.write_text(json.dumps(summary, indent=2) + "\n")
print(f"[5B] wrote metrics {args.metrics_json}", flush=True)
print(json.dumps(summary, indent=2), flush=True)
if __name__ == "__main__":
main()
@@ -1,130 +0,0 @@
"""Compare Wan VAE and TAEHV decode on saved FastWan latents."""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import DEFAULT_MODEL_ROOT, decode_latents_to_video
def _torch_mps_memory() -> dict[str, int | None]:
"""
Report current and recommended memory usage for the MPS backend.
Returns:
dict[str, int | None]: Memory metrics in bytes, or `None` values when
PyTorch or MPS is unavailable.
"""
try:
import torch
except ImportError:
return {
"current_allocated_bytes": None,
"driver_allocated_bytes": None,
"recommended_max_bytes": None,
}
if not torch.backends.mps.is_available():
return {
"current_allocated_bytes": None,
"driver_allocated_bytes": None,
"recommended_max_bytes": None,
}
return {
"current_allocated_bytes": int(torch.mps.current_allocated_memory()),
"driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
"recommended_max_bytes": int(torch.mps.recommended_max_memory()),
}
def _parse_backends(raw: str) -> list[str]:
"""
Parse and validate a comma-separated list of decoding backends.
Parameters:
raw (str): Comma-separated backend names.
Returns:
list[str]: Trimmed, supported backend names in input order.
Raises:
ValueError: If the input contains an unsupported backend.
"""
backends = [backend.strip() for backend in raw.split(",") if backend.strip()]
allowed = {"wan-vae", "taehv"}
unknown = sorted(set(backends) - allowed)
if unknown:
raise ValueError(f"Unsupported decode backends: {unknown}")
return backends
def main() -> None:
"""
Benchmark selected Wan latent decoding backends and record their performance metrics.
Loads the specified latent array, decodes it with each selected backend, exports the
results as MP4 files, and writes per-backend timing and Torch MPS memory metrics to
`metrics.json`.
"""
parser = argparse.ArgumentParser(description="Benchmark decode backends on saved Wan/FastWan latents.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--latents-path", type=Path, required=True)
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_decode_benchmark"))
parser.add_argument("--backends", default="wan-vae,taehv")
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--taehv-source-path", type=Path, default=None)
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
parser.add_argument("--taehv-parallel", action="store_true")
args = parser.parse_args()
latents = np.load(args.latents_path)
args.output_dir.mkdir(parents=True, exist_ok=True)
rows = []
for backend in _parse_backends(args.backends):
print(f"=== Decode backend: {backend} ===")
before = _torch_mps_memory()
start = time.perf_counter()
output_path = args.output_dir / f"{args.latents_path.stem}_{backend}.mp4"
decode_latents_to_video(
model_root=args.model_root,
latents_np=latents,
output_path=output_path,
fps=args.fps,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
backend=backend,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
)
elapsed = time.perf_counter() - start
after = _torch_mps_memory()
metrics = {
"backend": backend,
"latents_path": str(args.latents_path),
"latents_shape": list(latents.shape),
"decode_export_s": elapsed,
"torch_mps_current_before_bytes": before["current_allocated_bytes"],
"torch_mps_current_after_bytes": after["current_allocated_bytes"],
"torch_mps_driver_before_bytes": before["driver_allocated_bytes"],
"torch_mps_driver_after_bytes": after["driver_allocated_bytes"],
"torch_mps_recommended_max_bytes": after["recommended_max_bytes"],
"output_path": str(output_path),
}
rows.append(metrics)
print(json.dumps(metrics, indent=2))
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(rows, indent=2))
print(f"Wrote decode metrics to: {metrics_path}")
if __name__ == "__main__":
main()
File diff suppressed because it is too large Load Diff
@@ -1,365 +0,0 @@
"""Benchmark MLX FastWan quantization modes with one shared prompt encode."""
from __future__ import annotations
import argparse
import json
import time
from pathlib import Path
from typing import cast
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import (
DEFAULT_MODEL_ROOT,
decode_latents_to_video,
encode_prompt,
make_rotary_embeddings,
)
from fastvideo.mlx_runtime.memory import cleanup_mlx
def _parse_modes(raw: str) -> list[str]:
"""
Parse and validate a comma-separated list of quantization modes.
Parameters:
raw (str): Comma-separated mode names.
Returns:
list[str]: Normalized, whitespace-trimmed mode names.
Raises:
ValueError: If any mode is unsupported.
"""
modes = [mode.strip() for mode in raw.split(",") if mode.strip()]
allowed = {"none", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"}
unknown = sorted(set(modes) - allowed)
if unknown:
raise ValueError(f"Unsupported modes: {unknown}")
return modes
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
"""
Compare candidate and baseline latent arrays using error and signal-quality metrics.
Parameters:
candidate (np.ndarray): Latent array to evaluate.
baseline (np.ndarray): Reference latent array for comparison.
Returns:
dict[str, float]: Mean squared error, mean absolute error, maximum absolute
error, and signal-to-noise ratio in decibels between the arrays.
"""
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
mse = float(np.mean(np.square(diff)))
mae = float(np.mean(np.abs(diff)))
max_abs = float(np.max(np.abs(diff)))
signal = float(np.mean(np.square(baseline.astype(np.float32))))
return {
"latent_mse_vs_fp16": mse,
"latent_mae_vs_fp16": mae,
"latent_max_abs_vs_fp16": max_abs,
"latent_snr_db_vs_fp16": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
}
def _torch_mps_memory() -> dict[str, int | None]:
"""
Report PyTorch MPS memory statistics when PyTorch MPS is available.
Returns:
dict[str, int | None]: A mapping of MPS memory metric names to byte counts, or `None` values when PyTorch or MPS is unavailable.
"""
try:
import torch
except ImportError:
return {
"torch_mps_current_allocated_bytes": None,
"torch_mps_driver_allocated_bytes": None,
"torch_mps_recommended_max_bytes": None,
}
if not torch.backends.mps.is_available():
return {
"torch_mps_current_allocated_bytes": None,
"torch_mps_driver_allocated_bytes": None,
"torch_mps_recommended_max_bytes": None,
}
return {
"torch_mps_current_allocated_bytes": int(torch.mps.current_allocated_memory()),
"torch_mps_driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
"torch_mps_recommended_max_bytes": int(torch.mps.recommended_max_memory()),
}
def _decode_with_metrics(*, args, latents: np.ndarray, output_path: Path) -> dict[str, float | int | None | str]:
"""
Decode latents to a video and collect export timing and PyTorch MPS memory metrics.
Parameters:
args: Configuration values for decoding and video export.
latents (np.ndarray): Latent representation to decode.
output_path (Path): Destination path for the exported video.
Returns:
dict[str, float | int | None | str]: Video export duration and PyTorch MPS memory measurements.
"""
before = _torch_mps_memory()
decode_start = time.perf_counter()
decode_latents_to_video(
model_root=args.model_root,
latents_np=latents,
output_path=output_path,
fps=args.fps,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
backend=args.decode_backend,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
)
decode_time = time.perf_counter() - decode_start
after = _torch_mps_memory()
return {
"decode_export_s": decode_time,
"decode_torch_mps_current_before_bytes": before["torch_mps_current_allocated_bytes"],
"decode_torch_mps_current_after_bytes": after["torch_mps_current_allocated_bytes"],
"decode_torch_mps_driver_before_bytes": before["torch_mps_driver_allocated_bytes"],
"decode_torch_mps_driver_after_bytes": after["torch_mps_driver_allocated_bytes"],
"decode_torch_mps_recommended_max_bytes": after["torch_mps_recommended_max_bytes"],
}
def _run_one_mode(
*,
mode: str,
args,
config: dict,
checkpoint_path: Path,
config_path: Path,
prompt_embeds,
freqs_cis,
):
"""
Run denoising for one quantization mode and collect performance and memory metrics.
Parameters:
mode (str): Quantization mode to benchmark.
args: Benchmark configuration, including dtype, dimensions, seed, scheduler, and denoising settings.
config (dict): Model configuration containing the input channel count.
checkpoint_path (Path): Path to the transformer checkpoint.
config_path (Path): Path to the transformer configuration.
prompt_embeds: Encoded prompt embeddings shared across benchmark modes.
freqs_cis: Rotary positional embeddings used during denoising.
Returns:
dict: The mode name, generated latent array, and metrics for model loading,
denoising, step timing, and MLX memory usage.
"""
import mlx.core as mx
import torch
from fastvideo.benchmarks.mlx_fastwan_bench import denoise_dmd_on_device
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
mx_dtype = mx.float16 if args.mlx_dtype == "fp16" else mx.float32
quantization = None if mode == "none" else mode
latent_frames = (args.num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
load_start = time.perf_counter()
mx.clear_cache()
mx.reset_peak_memory()
dit = mlx_dit_from_diffusers_safetensors(
checkpoint_path,
config_path,
dtype=args.mlx_dtype,
quantization=quantization,
)
load_time = time.perf_counter() - load_start
load_peak_memory = mx.get_peak_memory()
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
# Same torch generator sequence as the original host-round-trip loop
# (initial latents first, then one re-noise draw per intermediate step),
# so every mode still shares identical stochasticity.
generator = torch.Generator(device="cpu").manual_seed(args.seed)
latents_seed = torch.randn(
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
generator=generator,
dtype=torch.float32,
).numpy()
renoise_by_step = [
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
for _ in range(max(0, len(timesteps) - 1))
]
latents = mx.array(latents_seed).astype(mx_dtype)
encoder_hidden_states = mx.array(prompt_embeds.numpy()).astype(mx_dtype)
denoise_start = time.perf_counter()
mx.reset_peak_memory()
latents_np, step_times = denoise_dmd_on_device(
mx=mx,
dit=dit,
latents=latents,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis,
timesteps=timesteps,
renoise_by_step=renoise_by_step,
schedule=schedule,
dmd_step=dmd_step,
mx_dtype=mx_dtype,
)
denoise_time = time.perf_counter() - denoise_start
denoise_peak_memory = mx.get_peak_memory()
active_memory = mx.get_active_memory()
return {
"mode": mode,
"latents": latents_np,
"metrics": {
"mlx_dit_load_s": load_time,
"mlx_denoise_s": denoise_time,
"mlx_denoise_first_step_s": step_times[0] if step_times else None,
"mlx_load_peak_bytes": int(load_peak_memory),
"mlx_denoise_peak_bytes": int(denoise_peak_memory),
"mlx_active_after_denoise_bytes": int(active_memory),
},
}
def main() -> None:
"""
Run the MLX FastWan quantization benchmark for the selected modes and write latency, memory, output, and latent-difference metrics to the output directory.
"""
parser = argparse.ArgumentParser(description="Benchmark MLX FastWan quantization modes.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--prompt", default="A snow leopard walks across a windy mountain ridge.")
parser.add_argument("--height", type=int, default=192)
parser.add_argument("--width", type=int, default=320)
parser.add_argument("--num-frames", type=int, default=17)
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
parser.add_argument("--flow-shift", type=float, default=8.0)
parser.add_argument("--max-sequence-length", type=int, default=256)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--mlx-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--modes", default="none,int8,int4,mxfp8,mxfp4,nvfp4")
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_quant_benchmark"))
parser.add_argument("--decode-backend", choices=("none", "wan-vae", "taehv"), default="taehv")
parser.add_argument("--taehv-source-path", type=Path, default=None)
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
parser.add_argument("--taehv-parallel", action="store_true")
args = parser.parse_args()
import mlx.core as mx
import torch
mx.random.seed(args.seed)
torch.manual_seed(args.seed)
args.output_dir.mkdir(parents=True, exist_ok=True)
config_path = args.model_root / "transformer/config.json"
checkpoint_path = args.model_root / "transformer/diffusion_pytorch_model.safetensors"
config = json.loads(config_path.read_text())
latent_frames = (args.num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
prompt_start = time.perf_counter()
prompt_embeds = encode_prompt(
model_root=args.model_root,
prompt=args.prompt,
max_sequence_length=args.max_sequence_length,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
)
prompt_time = time.perf_counter() - prompt_start
freqs_cis = make_rotary_embeddings(
config,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
)
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
baseline_latents = None
rows = []
for mode in _parse_modes(args.modes):
print(f"=== MLX quant mode: {mode} ===")
mode_start = time.perf_counter()
try:
result = _run_one_mode(
mode=mode,
args=args,
config=config,
checkpoint_path=checkpoint_path,
config_path=config_path,
prompt_embeds=prompt_embeds,
freqs_cis=freqs_cis,
)
except UnsupportedMLXQuantizationError as exc:
print(f"skipping mode (unsupported by this MLX build): {exc}")
rows.append({"mode": mode, "status": "unsupported_by_mlx", "error": str(exc)})
continue
cleanup_mlx(mx)
latents = result["latents"]
if baseline_latents is None:
baseline_latents = latents
latent_path = args.output_dir / f"latents_{mode}.npy"
np.save(latent_path, latents)
decode_time = 0.0
decode_metrics = {}
output_path = None
if args.decode_backend != "none":
output_path = args.output_dir / f"video_{mode}_{args.decode_backend}_{args.height}x{args.width}x{args.num_frames}.mp4"
decode_metrics = _decode_with_metrics(args=args, latents=latents, output_path=output_path)
decode_time = cast(float, decode_metrics["decode_export_s"])
mode_total = time.perf_counter() - mode_start
mlx_denoise_peak_bytes = int(result["metrics"]["mlx_denoise_peak_bytes"])
mlx_active_bytes = int(result["metrics"]["mlx_active_after_denoise_bytes"])
metrics = {
"mode": mode,
"status": "ok",
"prompt_encode_shared_s": prompt_time,
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"decode_backend": args.decode_backend,
"decode_export_s": decode_time,
"mode_total_excluding_shared_prompt_s": mode_total,
"mode_total_including_shared_prompt_s": mode_total + prompt_time,
"latents_path": str(latent_path),
"output_path": str(output_path) if output_path else None,
"mlx_denoise_peak_gib": mlx_denoise_peak_bytes / (1024**3),
"mlx_active_after_denoise_gib": mlx_active_bytes / (1024**3),
"mlx_dit_peak_under_16gb": mlx_denoise_peak_bytes < 16 * 1024**3,
"mlx_dit_active_under_16gb": mlx_active_bytes < 16 * 1024**3,
"mac_16gb_status": (
"dit_memory_fits_16gb_measured_decode_separately"
if mlx_denoise_peak_bytes < 16 * 1024**3 else "dit_memory_exceeds_16gb"
),
**result["metrics"],
**decode_metrics,
**_latent_delta_metrics(latents, baseline_latents),
}
rows.append(metrics)
print(json.dumps(metrics, indent=2))
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(rows, indent=2))
print(f"Wrote benchmark metrics to: {metrics_path}")
if __name__ == "__main__":
main()
@@ -1,104 +0,0 @@
"""Compare generated MP4s against a reference MP4 with simple pixel metrics."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import numpy as np
def _read_video(path: Path) -> np.ndarray:
"""
Read all frames from a video file as an RGB NumPy array.
Parameters:
path (Path): Path to the video file.
Returns:
np.ndarray: Video frames stacked along the first axis.
Raises:
ValueError: If the video contains no readable frames.
"""
import cv2
cap = cv2.VideoCapture(str(path))
frames = []
try:
while True:
ok, frame_bgr = cap.read()
if not ok:
break
frame_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)
frames.append(frame_rgb)
finally:
cap.release()
if not frames:
raise ValueError(f"No frames read from {path}")
return np.stack(frames, axis=0)
def _metrics(candidate: np.ndarray, reference: np.ndarray) -> dict[str, float | int | list[int]]:
"""
Compute pixel-level comparison metrics between candidate and reference video frames.
Parameters:
candidate (np.ndarray): Candidate video frames in frame, height, width, and channel order.
reference (np.ndarray): Reference video frames with the same shape as the candidate.
Returns:
dict[str, float | int | list[int]]: Frame dimensions and pixel comparison metrics, including MSE, MAE, maximum absolute difference, and PSNR in decibels.
Raises:
ValueError: If the candidate and reference arrays have different shapes.
"""
if candidate.shape != reference.shape:
raise ValueError(f"Shape mismatch: candidate={candidate.shape}, reference={reference.shape}")
candidate_f = candidate.astype(np.float32)
reference_f = reference.astype(np.float32)
diff = candidate_f - reference_f
mse = float(np.mean(np.square(diff)))
mae = float(np.mean(np.abs(diff)))
max_abs = float(np.max(np.abs(diff)))
psnr = float(20.0 * np.log10(255.0 / np.sqrt(mse))) if mse > 0 else float("inf")
return {
"frames": int(candidate.shape[0]),
"height": int(candidate.shape[1]),
"width": int(candidate.shape[2]),
"channels": int(candidate.shape[3]),
"mse_vs_reference": mse,
"mae_vs_reference": mae,
"max_abs_vs_reference": max_abs,
"psnr_db_vs_reference": psnr,
}
def main() -> None:
"""Compare candidate MP4 videos with a reference and write pixel-level metrics to a JSON file."""
parser = argparse.ArgumentParser(description="Compare MP4s against a reference MP4.")
parser.add_argument("--reference", type=Path, required=True)
parser.add_argument("--candidates", type=Path, nargs="+", required=True)
parser.add_argument("--metrics-json", type=Path, required=True)
args = parser.parse_args()
reference = _read_video(args.reference)
rows = []
for candidate_path in args.candidates:
candidate = _read_video(candidate_path)
row = {
"reference_path": str(args.reference),
"candidate_path": str(candidate_path),
**_metrics(candidate, reference),
}
rows.append(row)
print(json.dumps(row, indent=2))
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
args.metrics_json.write_text(json.dumps(rows, indent=2))
print(f"Wrote video quality metrics to: {args.metrics_json}")
if __name__ == "__main__":
main()
@@ -0,0 +1,130 @@
# Dual-stream T2V NVFP4 quantization-aware fine-tune (QAT).
#
# Adapted from fine_tuning/ltx2_3/nvfp4_qat_t2v.yaml. The quantization wiring
# is UNCHANGED and deliberately so: the curated linear prefix set in
# fastvideo/layers/quantization/nvfp4_config.py targets
# ltx2.blocks.{0..47}.* and attaches to this architecture unmodified --
# verified against the constructed module tree, 576 linears / 12.9B params,
# the same 12 suffix classes x 48 blocks as before. Do NOT narrow the set for
# this architecture: the cross-modal projections it targets
# (audio_to_video_attn, video_to_audio_attn) are deliberately enumerated in the
# DEPLOYMENT set, so dropping them in training would break train/deploy
# symmetry.
#
# What is NOT targeted, and why (all deliberate, none accidental):
# * head-dim-64 audio self/cross attention -- same constraint that kept it
# dense before;
# * per-head gate projections (to_gate_logits, [32, width]) -- tiny and
# precision-sensitive;
# * the embeddings connectors -- they live in the TEXT ENCODER, not the DiT,
# so deployment targeting must not reach across the component boundary.
# The coverage audit prints these as skipped-by-rule so a reader can see the
# exclusions were chosen, and FAILS if any declared class matches zero modules.
#
# Video attention uses ATTN_QAT_TRAIN (quantized forward, STE backward).
#
# Data paths, step counts and learning rate below are placeholders to adapt.
models:
student:
_target_: fastvideo.train.models.ltx2.LTX2Model
# REQUIRED: single-file bundle path. The config parser does NOT expand
# environment interpolations, so override this on the command line:
# --models.student.init_from /path/to/bundle.safetensors
init_from: SET_ME_ON_THE_COMMAND_LINE
trainable: true
enable_gradient_checkpointing_type: full
attention_backend: ATTN_QAT_TRAIN
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
# ~19B params: bf16 weights 35.4 GiB + bf16 grads 35.4 GiB + Adam fp32
# m/v 141.5 GiB = ~212 GiB, against ~173 GiB usable per GPU. Single-GPU is
# arithmetically impossible; 4-way shard puts it at ~53 GiB/GPU.
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
# REQUIRED: path to your preprocessed dataset.
data_path: data/your_dataset_preprocessed
dataloader_num_workers: 4
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
training_cfg_rate: 0.0
seed: 42
# Must match your preprocessed resolution/length.
# num_latent_t = (num_frames - 1) / 8 + 1.
num_latent_t: 11
num_height: 480
num_width: 832
num_frames: 81
optimizer:
# Carried from the validated LTX-2.0 QAT overfit run as a starting
# point; tune for your dataset size and batch configuration.
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
# Workload-dependent: set from your dataset size and target epochs.
max_train_steps: 2000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/dual_stream_t2v_nvfp4_qat_finetune
# A full training-state checkpoint is ~150GB for the 13B trainable
# video branch — size the interval and total limit to your storage.
training_state_checkpointing_steps: 500
checkpoints_total_limit: 2
resume_from_checkpoint: latest
tracker:
# No experiment tracking.
# It must be ["none"], NOT []: build_tracker appends wandb whenever the
# list is empty and project_name is non-empty, so [] silently means wandb.
trackers: ["none"]
model:
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
# so the default noise-minus-clean target reproduces the official
# unweighted masked-MSE (mask is all-ones for plain T2V).
precondition_outputs: false
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.ltx2.ltx2_pipeline.LTX2Pipeline
# REQUIRED: prompts to sample during training.
dataset_file: data/your_dataset_preprocessed/validation_prompts.json
every_steps: 250
# 8-step single-pass sampling matches the distilled checkpoint
# (validated in the LTX-2.3 overfit runs).
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Validation reuses the live transformer and temporarily replaces its
# ATTN_QAT_TRAIN implementations with the sm120 inference kernel.
# Set false on GB200 (see header).
attn_qat_infer: true
# quant_config selects NVFP4 QAT for the same curated linear prefixes as
# LTX-2 NVFP4 deployment (see the LTX-2.3 applicability note in the
# header). They fake-quantize through real FP4 GEMMs with an STE
# backward. The string resolves to NVFP4QATTrainConfig at parse.
pipeline:
dit_config:
quant_config: nvfp4_qat_train
-36
View File
@@ -318,38 +318,6 @@ if(BUILD_CXX_KERNELS)
# Combined FastVideo Extension
# Using name 'fastvideo_kernel_ops' to distinguish from the python package namespace
# ---------------------------------------------------------------------------
# VSA block-sparse attention forward, Blackwell (sm_100a) only.
#
# NOTE the "a" suffix: -arch=sm_100a is NOT enough -- it emits a plain sm_100 target and
# ptxas rejects every tcgen05 / setmaxnreg instruction. The explicit gencode spelling
# below is required, and matches the 10.0a entry in TORCH_CUDA_ARCH_LIST.
#
# Built for 64-token sparse blocks -- FastVideo's default (4,4,4) tiling, so no
# tile-size change and no top-k granularity change is needed. VSA_BHSD
# selects [B, H, S, D]. Both are compile-time; the Python is_supported() checks incoming
# tensors against them so callers fall back to Triton rather than getting a wrong answer.
# Read the ENVIRONMENT as well as the cache variable. When TORCH_CUDA_ARCH_LIST is
# exported (build.sh, and `pip install` with it set) the branch above only prints it --
# the cmake variable stays empty, so testing that alone silently skips the kernel and
# leaves a build that succeeds with the op missing.
set(ENABLE_VSA_SM100A OFF)
set(_VSA_ARCH_LIST "${TORCH_CUDA_ARCH_LIST}")
if(NOT _VSA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
set(_VSA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
endif()
if(_VSA_ARCH_LIST MATCHES "(^|[; ,])(10\\.0a|100a|sm_100a)([; ,]|$)")
set(ENABLE_VSA_SM100A ON)
endif()
if(ENABLE_VSA_SM100A)
message(STATUS "fastvideo-kernel: building block_sparse_sm100a (Blackwell, 64- and 128-token blocks)")
list(APPEND EXTENSION_SOURCES csrc/attention/block_sparse_sm100a.cu
csrc/attention/block_sparse_blk128_sm100a.cu)
set_source_files_properties(csrc/attention/block_sparse_sm100a.cu
csrc/attention/block_sparse_blk128_sm100a.cu PROPERTIES
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a;-DVSA_BHSD=true")
endif()
Python_add_library(fastvideo_kernel_ops MODULE USE_SABI ${SKBUILD_SABI_VERSION} WITH_SOABI
${EXTENSION_SOURCES}
)
@@ -365,14 +333,10 @@ if(BUILD_CXX_KERNELS)
# Build compile definitions list
set(COMPILE_DEFS TORCH_EXTENSION_NAME=fastvideo_kernel_ops)
if(ENABLE_VSA_SM100A)
list(APPEND COMPILE_DEFS TK_COMPILE_BLOCK_SPARSE_VSA_SM100A)
endif()
if(ENABLE_TK_KERNELS)
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
endif()
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
target_compile_options(fastvideo_kernel_ops PRIVATE
+11 -40
View File
@@ -17,7 +17,7 @@ Runtime-JIT kernels (no build step, ship in every wheel/image):
| Kernels | Where | Used when |
|---|---|---|
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
| FA4 CuTe-DSL block-sparse forward/backward (VSA-128/256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
## What gets built where, and when
@@ -64,49 +64,28 @@ cd fastvideo-kernel
./build.sh --rocm
```
### Optional: FA4 CuTe block-sparse backend (VSA-128/256 fastpath)
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
The VSA-128/256 fastpaths (tile volume 128 or 256, on NVIDIA Blackwell / sm_100) route to the
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
FlashAttention-4 CuTe-DSL block-sparse kernel exposed as `flash_attn.cute`. This is
an **optional** dependency: it is imported lazily, and `video_sparse_attn`
transparently falls back to the Triton backend when it is absent (so the package is
fully usable without it).
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`
and the public/private forward-backward bridges in `flash_attn.cute.interface`) are provided upstream by
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`,
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
commit `14c377950125c70b7a9dabf9c561fca53715ac7d`, the revision FastVideo pins as
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
the `flash-attn-4` source in the repo-root `pyproject.toml`; other revisions may
have incompatible block-sparse forward/backward interfaces.
Install it under its distribution name so its own runtime stack resolves with it.
Do **not** pre-install `nvidia-cutlass-dsl` by hand: this revision pins
`nvidia-cutlass-dsl==4.6.0.dev0` exactly, and a hand-installed 4.5.x floor either
gets silently upgraded or, if something else holds it back, leaves the CuTe
kernels broken.
have an incompatible `_flash_attn_fwd` signature.
```bash
pip install torchvision
pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@14c377950125c70b7a9dabf9c561fca53715ac7d#subdirectory=flash_attn/cute"
pip install "nvidia-cutlass-dsl>=4.5.0" torchvision
pip install "git+https://github.com/Dao-AILab/flash-attention.git@940cd9680f3315f2f06b43ab5bea2c2cf2d96806#subdirectory=flash_attn/cute"
```
That resolves `nvidia-cutlass-dsl` to 4.6.0.dev0 and `quack-kernels` to 0.5.3, a
combination this revision works with. A mismatched CuTe DSL only surfaces when the
kernel JIT-compiles, so the error points at CuTe internals rather than at the
install:
| Error on first VSA-128/256 CuTe call | Cause |
|---|---|
| `TypeError: fmax() missing 1 required positional argument: 'b'` | `nvidia-cutlass-dsl` 4.5.x |
| `AttributeError: module 'cutlass.cute.core' has no attribute 'ThrMma'` | `quack-kernels` older than 0.5.1 |
| `ImportError: cannot import name 'alloc_reserved_mbarrier'` | `quack-kernels` 0.6.2 or newer |
An environment whose `flash_attn.cute` came from a prebuilt flash-attn wheel rather
than from this pin hits the first row; that is what the overlay step in
`docker/Dockerfile` works around.
The CuTe kernels JIT-compile on first use. Forward and backward are verified on
Blackwell (sm_100) against `tests/test_vsa128_*.py` and `tests/test_vsa256_*.py`.
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
`tests/test_vsa256_forward*.py`.
## Usage
@@ -163,14 +142,6 @@ After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
# VSA-256 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 256 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 20
# VSA-128 FA4 CuTe forward/backward on Blackwell
python benchmarks/bench_vsa.py --block_size 128 --use_cute \
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 40
```
### TurboDiffusion Kernels
+20 -48
View File
@@ -2,9 +2,8 @@
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrappers:
- 64-token TK/Triton: fastvideo_kernel.block_sparse_attn.block_sparse_attn
- 128/256-token Triton/CuTe: fastvideo_kernel.block_sparse_attn_256
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
@@ -24,6 +23,9 @@ try:
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
@@ -39,11 +41,7 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens",
type=int,
nargs="+",
default=[49152],
help="Q sequence lengths (must be divisible by --block_size)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--kv_seq_lens",
type=int,
nargs="+",
@@ -53,13 +51,9 @@ def parse_arguments() -> argparse.Namespace:
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--block_size", type=int, default=64, choices=[64, 128, 256])
p.add_argument("--force_triton",
action="store_true",
help="Force wrapper to use Triton path (if supported by shapes).")
p.add_argument("--use_cute",
action="store_true",
help="Use the optional FA4 CuTe forward/backward path (requires --block_size 128 or 256).")
return p.parse_args()
@@ -90,38 +84,18 @@ def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def _configure_backend(args: argparse.Namespace) -> None:
if args.use_cute and args.block_size not in (128, 256):
raise ValueError("--use_cute requires --block_size 128 or 256")
if args.use_cute and args.force_triton:
raise ValueError("--use_cute and --force_triton are mutually exclusive")
if args.force_triton:
os.environ.pop("FASTVIDEO_VSA_CUTEDSL", None)
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
elif args.use_cute:
os.environ.pop("FASTVIDEO_VSA_TRITON", None)
os.environ.pop("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", None)
os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1"
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
_configure_backend(args)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128, block_sparse_attn_256
bs, h, d = args.batch_size, args.num_heads, args.head_dim
block_size = args.block_size
attention = {
64: block_sparse_attn,
128: block_sparse_attn_128,
256: block_sparse_attn_256,
}[block_size]
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
@@ -131,22 +105,20 @@ def main() -> None:
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"block_size={block_size}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.use_cute:
print("dispatch: FA4 CuTe")
elif args.force_triton:
if args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % block_size != 0 or kv_len % block_size != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by {block_size}")
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
continue
num_q_blocks = q_len // block_size
num_kv_blocks = kv_len // block_size
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
@@ -157,11 +129,11 @@ def main() -> None:
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full logical blocks.
variable_block_sizes = torch.full((num_kv_blocks, ), block_size, dtype=torch.int32, device="cuda")
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks, ), BLOCK_N, dtype=torch.int32, device="cuda")
def _fwd():
return attention(q, k, v, block_map, variable_block_sizes)
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
@@ -170,7 +142,7 @@ def main() -> None:
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = attention(q_, k_, v_, block_map, variable_block_sizes)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
@@ -184,7 +156,7 @@ def main() -> None:
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, block_size)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
@@ -1,7 +0,0 @@
// block_sparse_blk128_sm100a.cu -- the 128-token-block instantiation of the torch binding.
//
// Same source as block_sparse_sm100a.cu with VSA_BLK128 set: the kernel and launch land in
// namespace vsa_blk128 (distinct symbols, no ODR clash with the blk64 objects) and the
// exported entry point becomes block_sparse_sm100a_blk128_fwd.
#define VSA_BLK128 true
#include "block_sparse_sm100a.cu"
File diff suppressed because it is too large Load Diff
@@ -1,201 +0,0 @@
#ifndef BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
#define BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
// Launch surface for the sm_100a VSA block-sparse FMHA forward.
//
// Everything a caller needs: a POD argument struct, a predicate saying whether this build can
// run those arguments, and one launch entry point. The benchmark in
// block_sparse_bench_sm100a.cu and the torch binding both go through here, so there is
// one tensormap construction and one launch configuration rather than two that can drift.
//
// Two compile-time knobs select the four builds:
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
// VSA_BHSD false -> [token][head][dim] (BSHD), true -> [batch][head][token][dim] (BHSD)
#include "block_sparse_kernel_sm100a.cuh"
namespace VSA_NAMESPACE {
struct BlockSparseVsaArgs {
const __nv_bfloat16* q;
const __nv_bfloat16* k;
const __nv_bfloat16* v; // natural layout; only blk128 reads it (blk64 still needs v_t)
const __nv_bfloat16* v_t; // unused: kept so the bench's V_T buffer still binds
__nv_bfloat16* o;
float* lse; // [batch, num_heads, seqlen] fp32, or nullptr
const int* q2k_idx; // [batch*num_heads*num_blocks, max_kv] int32
const int* q2k_num; // [batch*num_heads*num_blocks] int32
const int* variable_block_sizes; // [num_blocks] int32, valid tokens per block
int batch;
int num_heads;
int seqlen;
int head_dim;
int num_blocks;
int max_kv;
float sm_scale;
};
// cudaSuccess iff this build can run `a`. Deliberately conservative: the caller is expected
// to fall back to its own implementation rather than get a wrong answer.
__host__ inline cudaError_t block_sparse_supported(const BlockSparseVsaArgs& a) {
if (a.head_dim != HEAD_DIM) return cudaErrorInvalidValue; // compile-time in the kernel
if (a.num_blocks % 2 != 0) return cudaErrorInvalidValue; // a CTA owns an adjacent pair
if (a.seqlen != a.num_blocks * BLOCK) return cudaErrorInvalidValue;
if (a.max_kv < 1 || a.num_blocks < 1) return cudaErrorInvalidValue;
if (a.q == nullptr || a.k == nullptr || a.o == nullptr) return cudaErrorInvalidValue;
if (a.q2k_idx == nullptr || a.q2k_num == nullptr) return cudaErrorInvalidValue;
// FastVideo always supplies this; without it padded keys would be attended as real zeros.
if (a.variable_block_sizes == nullptr) return cudaErrorInvalidValue;
// V is read MN-major at BOTH block sizes now, so no pre-transposed V_T is ever needed.
if (a.v == nullptr) return cudaErrorInvalidValue;
return cudaSuccess;
}
__host__ inline cudaError_t launch_block_sparse_sm100a(const BlockSparseVsaArgs& a,
cudaStream_t stream) {
const cudaError_t sup = block_sparse_supported(a);
if (sup != cudaSuccess) return sup;
const int B = a.batch, H = a.num_heads, S = a.seqlen, hd = a.head_dim;
const int num_blocks = a.num_blocks, max_kv = a.max_kv;
const long tq = (long)B * S;
const int packed_mtiles_per_seq = num_blocks / 2;
const int total_work = B * H * packed_mtiles_per_seq;
constexpr bool BHSD = VSA_BHSD;
CUtensorMap tq_, tk_, tvt_, tv_, to_;
{
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16, BHSD ? (uint64_t)((long)B * H) : (uint64_t)H,
BHSD ? (uint64_t)S : (uint64_t)tq, (uint64_t)Q_SUBTILES };
uint64_t gs[3] = { BHSD ? (uint64_t)((long)S * hd) * 2u : (uint64_t)hd * 2u,
BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
(uint64_t)SUB_COLS_BF16 * 2u };
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, 1u, (uint32_t)M_TILE, (uint32_t)Q_SUBTILES };
uint32_t es[4] = { 1u, 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tq_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4,
const_cast<__nv_bfloat16*>(a.q), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
if (cuTensorMapEncodeTiled(&to_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, a.o, gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
}
{
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16,
BHSD ? (uint64_t)S : (uint64_t)tq,
BHSD ? (uint64_t)(hd / SUB_COLS_BF16)
: (uint64_t)((long)H * hd / SUB_COLS_BF16),
(uint64_t)((long)B * H) };
uint64_t gs[3] = { BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
(uint64_t)SUB_COLS_BF16 * 2u,
(uint64_t)((long)S * hd) * 2u };
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, (uint32_t)BLOCK,
BLK128 ? (uint32_t)K_SUBTILES : 1u, 1u };
uint32_t es[4] = { 1u, 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tk_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
const_cast<__nv_bfloat16*>(a.k), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
// V map is byte-for-byte the K map over a.v: MN-major V needs no transpose (blk128).
const __nv_bfloat16* vbase = a.v ? a.v : a.k;
if (cuTensorMapEncodeTiled(&tv_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
const_cast<__nv_bfloat16*>(vbase), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
}
// V_T map: blk64 only. Unused at blk128 but must still be a valid tensormap to pass by value.
{
const __nv_bfloat16* vt = a.v_t ? a.v_t : a.k;
if constexpr (BLK128) {
uint64_t gd[3] = { (uint64_t)SUB_COLS_BF16, (uint64_t)((long)H * hd),
(uint64_t)((long)tq / SUB_COLS_BF16) };
uint64_t gs[2] = { (uint64_t)tq * 2u, (uint64_t)SUB_COLS_BF16 * 2u };
uint32_t bd[3] = { (uint32_t)SUB_COLS_BF16, (uint32_t)hd, (uint32_t)V_SUBTILES };
uint32_t es[3] = { 1u, 1u, 1u };
if (cuTensorMapEncodeTiled(&tvt_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 3,
const_cast<__nv_bfloat16*>(vt), gd, gs, bd, es,
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
return cudaErrorInvalidValue;
} else {
if (make_tma_2d_tiled(&tvt_, const_cast<__nv_bfloat16*>(vt), (long)H * hd, (int)tq, hd,
SUB_COLS_BF16, 2, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
CU_TENSOR_MAP_SWIZZLE_128B) != cudaSuccess)
return cudaErrorInvalidValue;
}
}
const size_t smem =
(size_t)2 * Q_TILE_BYTES + NUM_KV_STAGES * KV_RING_SLOT_BYTES
+ (size_t)2 * M_TILE * HEAD_DIM * sizeof(__nv_bfloat16)
+ (2 * NUM_KV_STAGES + 22) * 8
+ (size_t)CLC_STAGES * (2 * 8 + 16) + 16
+ 8
+ (size_t)2 * STAT_REGIONS * STATS * sizeof(float)
+ 256;
#ifndef VSA_NAMED_BAR
#define VSA_NAMED_BAR false
#endif
#ifndef VSA_THROTTLE
#define VSA_THROTTLE false
#endif
#ifndef VSA_USE_CLC
#define VSA_USE_CLC true
#endif
constexpr bool FULL_NAMED_BAR = VSA_NAMED_BAR, EX2_EMU = true, SPLIT_P = true,
SOFTMAX_THROTTLE = VSA_THROTTLE, USE_CLC = VSA_USE_CLC,
Q_RASTER = true, MHA = true;
auto kfn = &fmha_context_bf16_gen_kernel<32, FULL_NAMED_BAR, EX2_EMU, SPLIT_P,
SOFTMAX_THROTTLE, USE_CLC, Q_RASTER, MHA,
/*RESCALE_THRESHOLD=*/8, /*BHSD=*/VSA_BHSD>;
cudaError_t e = cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
if (e != cudaSuccess) return e;
const unsigned long long magic0 = make_magic((unsigned)(H * packed_mtiles_per_seq));
const unsigned long long magic1 = make_magic((unsigned)H);
const unsigned long long magic2 = make_magic((unsigned)packed_mtiles_per_seq);
const float scale_log2 = a.sm_scale * (float)M_LOG2E;
int numSM = 0;
e = cudaDeviceGetAttribute(&numSM, cudaDevAttrMultiProcessorCount, 0);
if (e != cudaSuccess) return e;
const int num_ctas = USE_CLC ? total_work : (total_work < numSM ? total_work : numSM);
dim3 grid(num_ctas, 1, 1), block(N_WARPS * 32, 1, 1);
if (USE_CLC) {
cudaLaunchConfig_t cfg = {};
cfg.gridDim = grid; cfg.blockDim = block; cfg.dynamicSmemBytes = smem; cfg.stream = stream;
cudaLaunchAttribute cfgAttr[1];
cfgAttr[0].id = cudaLaunchAttributeClusterDimension;
cfgAttr[0].val.clusterDim.x = 1; cfgAttr[0].val.clusterDim.y = 1;
cfgAttr[0].val.clusterDim.z = 1;
cfg.attrs = cfgAttr; cfg.numAttrs = 1;
return cudaLaunchKernelEx(&cfg, kfn, tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B,
num_blocks, packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
}
kfn<<<grid, block, smem, stream>>>(tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B, num_blocks,
packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
return cudaGetLastError();
}
} // namespace VSA_NAMESPACE
// Callers (the bench, the torch binding) keep using unqualified names; each translation unit
// only ever sees the one configuration its VSA_BLK128 selected.
using namespace VSA_NAMESPACE;
#endif // BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
@@ -1,114 +0,0 @@
// block_sparse_sm100a.cu -- torch binding for the sm_100a VSA block-sparse FMHA forward.
//
// Forward only: returns (out, lse) so FastVideo's existing Triton backward keeps working
// unchanged. lse is exactly the M tensor triton_block_sparse_attn_forward writes --
// max(qk * qk_scale) + log2(l), [B, H, S] fp32 -- which is what lets
// block_sparse_attn_backward_triton run against our forward untouched.
//
// The build is fixed at compile time by two flags, so one extension carries one configuration:
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
// VSA_BHSD false -> [B, S, H, D], true -> [B, H, S, D]
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include "block_sparse_launch_sm100a.cuh"
namespace {
void check_qkv(const torch::Tensor& t, const char* name, int64_t B, int64_t H, int64_t S,
int64_t D) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.scalar_type() == at::kBFloat16, name, " must be bfloat16, got ", t.scalar_type());
TORCH_CHECK(t.dim() == 4, name, " must be 4-D, got ", t.dim(), " dims");
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
if (VSA_BHSD) {
TORCH_CHECK(t.size(0) == B && t.size(1) == H && t.size(2) == S && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", H, ",", S, ",", D, "]");
} else {
TORCH_CHECK(t.size(0) == B && t.size(1) == S && t.size(2) == H && t.size(3) == D, name,
" has shape ", t.sizes(), ", expected [", B, ",", S, ",", H, ",", D, "]");
}
}
void check_index(const torch::Tensor& t, const char* name) {
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
TORCH_CHECK(t.scalar_type() == at::kInt, name, " must be int32, got ", t.scalar_type());
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
}
} // namespace
// The exported symbol carries the block size: block_sparse_sm100a_fwd is the 64-token build,
// block_sparse_sm100a_blk128_fwd the 128-token one (block_sparse_blk128_sm100a.cu re-includes
// this file with VSA_BLK128 set). The python backend picks by the metadata's block size.
#if VSA_BLK128
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_blk128_fwd
#else
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_fwd
#endif
// Returns {out} or {out, lse}. Layout of out matches the inputs.
std::vector<torch::Tensor> BLOCK_SPARSE_SM100A_FWD(torch::Tensor q, torch::Tensor k,
torch::Tensor v,
c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx,
torch::Tensor q2k_num,
torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse) {
const at::cuda::OptionalCUDAGuard guard(device_of(q));
const int64_t B = q.size(0);
const int64_t H = VSA_BHSD ? q.size(1) : q.size(2);
const int64_t S = VSA_BHSD ? q.size(2) : q.size(1);
const int64_t D = q.size(3);
check_qkv(q, "q", B, H, S, D);
check_qkv(k, "k", B, H, S, D);
check_qkv(v, "v", B, H, S, D);
check_index(q2k_idx, "q2k_idx");
check_index(q2k_num, "q2k_num");
check_index(variable_block_sizes, "variable_block_sizes");
const int64_t num_blocks = variable_block_sizes.numel();
const int64_t max_kv = q2k_idx.size(-1);
TORCH_CHECK(S == num_blocks * BLOCK, "seqlen ", S, " must equal num_blocks (", num_blocks,
") * ", BLOCK, "; FastVideo pads the sequence up to whole blocks");
auto out = torch::empty_like(q);
torch::Tensor lse;
if (need_lse) lse = torch::empty({B, H, S}, q.options().dtype(torch::kFloat32));
BlockSparseVsaArgs a{};
a.q = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr());
a.k = reinterpret_cast<const __nv_bfloat16*>(k.data_ptr());
a.v = reinterpret_cast<const __nv_bfloat16*>(v.data_ptr());
a.v_t = v_t.has_value() ? reinterpret_cast<const __nv_bfloat16*>(v_t->data_ptr()) : nullptr;
a.o = reinterpret_cast<__nv_bfloat16*>(out.data_ptr());
a.lse = need_lse ? lse.data_ptr<float>() : nullptr;
a.q2k_idx = q2k_idx.data_ptr<int>();
a.q2k_num = q2k_num.data_ptr<int>();
a.variable_block_sizes = variable_block_sizes.data_ptr<int>();
a.batch = (int)B;
a.num_heads = (int)H;
a.seqlen = (int)S;
a.head_dim = (int)D;
a.num_blocks = (int)num_blocks;
a.max_kv = (int)max_kv;
a.sm_scale = (float)sm_scale;
// Report an unsupported regime loudly rather than returning plausible-looking wrong values.
TORCH_CHECK(block_sparse_supported(a) == cudaSuccess,
"block_sparse_sm100a: unsupported configuration -- requires head_dim==",
HEAD_DIM, ", an even num_blocks, seqlen == num_blocks*", BLOCK,
", and a variable_block_sizes tensor. Got head_dim=", D, " num_blocks=",
num_blocks, " seqlen=", S);
const cudaError_t err = launch_block_sparse_sm100a(a, at::cuda::getCurrentCUDAStream());
TORCH_CHECK(err == cudaSuccess,
"block_sparse_sm100a launch failed: ", cudaGetErrorString(err));
if (need_lse) return {out, lse};
return {out};
}
@@ -1,877 +0,0 @@
// primitives.cuh -- device primitives for the sm_100a VSA block-sparse attention
// forward: tcgen05 (alloc / mma / ld / st / commit / wait / fence), TMA load / store /
// tensormap, mbarrier, cluster launch control, setmaxnreg, fast math, and the FMHA helpers.
//
// Generated and pruned to what the kernel reaches -- do not edit by hand.
#pragma once
#include <cstdint>
#include <cstdio>
#include <cstdlib>
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cmath>
#include <cassert>
#include <cstring>
#include <vector_types.h>
#ifndef CUDA_CHECK
#define CUDA_CHECK(stmt) do { \
cudaError_t _e = (stmt); \
if (_e != cudaSuccess) { \
fprintf(stderr, "CUDA error %s:%d: %s -> %s\n", \
__FILE__, __LINE__, #stmt, cudaGetErrorString(_e)); \
std::exit(1); \
} \
} while (0)
#endif
__device__ __forceinline__
uint64_t mbarrier_arrive(uint32_t mbar_smem) {
uint64_t state;
asm volatile("mbarrier.arrive.shared::cta.b64 %0, [%1];\n"
: "=l"(state) : "r"(mbar_smem) : "memory");
return state;
}
__device__ __forceinline__
void mbarrier_arrive_nostate(uint32_t mbar_smem) {
asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n"
:: "r"(mbar_smem) : "memory");
}
__device__ __forceinline__
void mbarrier_arrive_cluster_default(uint32_t cluster_smem_addr) {
asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];\n"
:: "r"(cluster_smem_addr) : "memory");
}
__device__ __forceinline__
void mbarrier_arrive_expect_tx(uint32_t mbar_smem, uint32_t expected_bytes) {
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;\n"
:: "r"(mbar_smem), "r"(expected_bytes) : "memory");
}
__device__ __forceinline__
void mbarrier_wait_parity_suspend(uint32_t mbar_smem, uint32_t phase_parity) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1, 10000000;\n"
"@!P1 bra.uni LAB_WAIT;\n"
"}\n"
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
}
__device__ __forceinline__
void mbarrier_wait_parity(uint32_t mbar_smem, uint32_t phase_parity) {
asm volatile(
"{\n"
".reg .pred P1;\n"
"LAB_WAIT_HOT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
"@!P1 bra.uni LAB_WAIT_HOT;\n"
"}\n"
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
}
__device__ __forceinline__
void fence_proxy_async_shared_cta() {
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__
void fence_proxy_async_shared() {
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
}
__device__ __forceinline__ void clc_try_cancel_async(
uint32_t smem_dst, uint32_t mbar_smem) {
asm volatile(
"clusterlaunchcontrol.try_cancel.async.shared::cta.mbarrier::complete_tx::bytes.b128"
" [%0], [%1];\n"
:: "r"(smem_dst), "r"(mbar_smem) : "memory");
}
__device__ __forceinline__ void clc_load_response(
uint32_t smem_slot, uint32_t& r0, uint32_t& r1,
uint32_t& r2, uint32_t& r3) {
asm volatile("ld.shared::cta.v4.b32 {%0, %1, %2, %3}, [%4];\n"
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
: "r"(smem_slot));
}
template <int NUM_STAGES>
struct MbarrierPhaseTracker {
uint32_t phase[NUM_STAGES];
int idx;
__device__ __forceinline__
void init() {
for (int i = 0; i < NUM_STAGES; ++i) phase[i] = 0;
idx = 0;
}
__device__ __forceinline__
uint32_t current_phase() const { return phase[idx]; }
__device__ __forceinline__
void advance() {
phase[idx] ^= 1u;
idx = (idx + 1) % NUM_STAGES;
}
__device__ __forceinline__
int stage() const { return idx; }
};
template <int NUM_STAGES>
struct PhaseTracker {
int stage;
uint32_t phase;
__device__ __forceinline__
PhaseTracker() : stage(0), phase(0) {}
__device__ __forceinline__
void advance() {
stage++;
if (stage == NUM_STAGES) {
stage = 0;
phase ^= 1;
}
}
__device__ __forceinline__
int get_stage() const { return stage; }
__device__ __forceinline__
uint32_t get_phase() const { return phase; }
};
template <int NUM_STAGES>
struct EmptyPhaseTracker {
int stage;
uint32_t phase;
__device__ __forceinline__
EmptyPhaseTracker() : stage(0), phase(1) {}
__device__ __forceinline__
void advance() {
stage++;
if (stage == NUM_STAGES) {
stage = 0;
phase ^= 1;
}
}
__device__ __forceinline__
int get_stage() const { return stage; }
__device__ __forceinline__
uint32_t get_phase() const { return phase; }
};
template <int STAGES>
__device__ __forceinline__
void advance_stage_phase(int& stage, uint32_t& phase) {
++stage;
if (stage == STAGES) {
stage = 0;
phase ^= 1u;
}
}
static constexpr uint32_t SM100_CLC_PEER_MASK = 0xFEFFFFFF;
struct ClcTileInfo {
int m_tile;
int n_tile;
bool valid;
};
enum class ClcRasterOrder { AlongN, AlongM };
__device__ __forceinline__
void clc_arrive_expect_tx_cta(uint32_t clc_full_local_addr, uint32_t tx_bytes) {
if ((threadIdx.x & 31) == 0) {
mbarrier_arrive_expect_tx(clc_full_local_addr, tx_bytes);
}
}
__device__ __forceinline__
void clc_consumer_release(uint32_t clc_empty_local_addr) {
uint32_t peer0_addr = clc_empty_local_addr & SM100_CLC_PEER_MASK;
mbarrier_arrive_cluster_default(peer0_addr);
}
__device__ __forceinline__
void clc_consumer_release_cta(uint32_t clc_empty_local_addr) {
mbarrier_arrive_nostate(clc_empty_local_addr);
}
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER>
__device__ __forceinline__
ClcTileInfo clc_parse_response(uint32_t resp_smem_addr) {
uint32_t d0, d1, d2, d3;
fence_proxy_async_shared_cta();
clc_load_response(resp_smem_addr, d0, d1, d2, d3);
const int ctaid_x = static_cast<int>(d0);
const int ctaid_y = static_cast<int>(d1 & 0xFFFFu);
const bool valid = (d2 & 1u) != 0u;
(void)d3;
ClcTileInfo info;
info.valid = valid;
if constexpr (ORDER == ClcRasterOrder::AlongN) {
info.m_tile = ctaid_y / CLUSTER_SHAPE_M;
info.n_tile = ctaid_x / CLUSTER_SHAPE_N;
} else {
info.m_tile = ctaid_x / CLUSTER_SHAPE_M;
info.n_tile = ctaid_y / CLUSTER_SHAPE_N;
}
return info;
}
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER,
int CTA_GROUP = 2, bool SUSPEND = false>
__device__ __forceinline__
ClcTileInfo clc_fetch_next_tile(
uint64_t* clc_full_bar, uint64_t* clc_empty_bar, uint32_t* clc_response,
int clc_cons_stage, uint32_t clc_cons_phase, bool do_release) {
uint32_t full_addr = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_full_bar[clc_cons_stage]));
if constexpr (SUSPEND) mbarrier_wait_parity_suspend(full_addr, clc_cons_phase);
else mbarrier_wait_parity(full_addr, clc_cons_phase);
uint32_t resp_addr = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_response[clc_cons_stage * 4]));
ClcTileInfo t = clc_parse_response<
CLUSTER_SHAPE_M, CLUSTER_SHAPE_N, ORDER>(resp_addr);
if (do_release) {
uint32_t empty_local = static_cast<uint32_t>(
__cvta_generic_to_shared(&clc_empty_bar[clc_cons_stage]));
if constexpr (CTA_GROUP == 1) {
clc_consumer_release_cta(empty_local);
} else {
clc_consumer_release(empty_local);
}
}
return t;
}
template <int STAGES = 2>
__device__ __forceinline__
void clc_fetch_next_tile_advance(int& clc_cons_stage,
uint32_t& clc_cons_phase) {
advance_stage_phase<STAGES>(clc_cons_stage, clc_cons_phase);
}
__device__ __forceinline__ unsigned fdiv(unsigned n, unsigned long long pk) {
unsigned M = (unsigned)pk;
if (M == 0u) return n;
return __umulhi(n, M) >> (unsigned)(pk >> 32);
}
__host__ inline unsigned long long make_magic(unsigned d) {
if (d <= 1u) return 0ULL;
unsigned l = 0; while ((1u << (l + 1)) <= d) ++l;
unsigned p = 31u + l;
unsigned long long m = ((1ull << p) + (unsigned long long)d - 1ull) / d;
return (m & 0xffffffffULL) | ((unsigned long long)(p - 32u) << 32);
}
template <int BARRIER_ID>
__device__ __forceinline__
void bar_sync(uint32_t thread_count) {
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
"bar.sync: BARRIER_ID must be in [0, 15]");
asm volatile("bar.sync %0, %1;\n"
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
}
template <int BARRIER_ID>
__device__ __forceinline__
void bar_arrive(uint32_t thread_count) {
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
"bar.arrive: BARRIER_ID must be in [0, 15]");
asm volatile("bar.arrive %0, %1;\n"
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
}
__device__ __forceinline__ void full_bar_arrive(int m_tile, int band) {
switch (1 + m_tile * 4 + band) {
case 1: bar_arrive<1>(64); break;
case 2: bar_arrive<2>(64); break;
case 3: bar_arrive<3>(64); break;
case 4: bar_arrive<4>(64); break;
case 5: bar_arrive<5>(64); break;
case 6: bar_arrive<6>(64); break;
case 7: bar_arrive<7>(64); break;
case 8: bar_arrive<8>(64); break;
}
}
__device__ __forceinline__ void full_bar_wait(int m_tile, int band) {
switch (1 + m_tile * 4 + band) {
case 1: bar_sync<1>(64); break;
case 2: bar_sync<2>(64); break;
case 3: bar_sync<3>(64); break;
case 4: bar_sync<4>(64); break;
case 5: bar_sync<5>(64); break;
case 6: bar_sync<6>(64); break;
case 7: bar_sync<7>(64); break;
case 8: bar_sync<8>(64); break;
}
}
template <bool IS_CAUSAL, int K_TILE>
__device__ __forceinline__ void mask_s_row_r2p(float* scores, int k_offset, int q_pos, int seqlen_k) {
int n_keep = seqlen_k - k_offset;
if constexpr (IS_CAUSAL) {
const int causal = q_pos - k_offset + 1;
n_keep = n_keep < causal ? n_keep : causal;
}
#pragma unroll
for (int s = 0; s < K_TILE / 32; ++s) {
int m = (s + 1) * 32 - n_keep;
m = m < 0 ? 0 : (m > 32 ? 32 : m);
const uint32_t keep = (m >= 32) ? 0u : (0xFFFFFFFFu >> m);
#pragma unroll
for (int i = 0; i < 32; ++i)
if (!(keep & (1u << i))) scores[s * 32 + i] = -INFINITY;
}
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_alloc(uint32_t smem_dst_ptr,
uint32_t n_cols) {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_alloc: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile(
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
:: "r"(smem_dst_ptr), "r"(n_cols));
} else {
asm volatile(
"tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;\n"
:: "r"(smem_dst_ptr), "r"(n_cols));
}
}
__device__ __forceinline__ void tcgen05_st_32x32b_x16(
uint32_t tmem_addr, const uint32_t (&r)[16]) {
asm volatile("tcgen05.st.sync.aligned.32x32b.x16.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
"%11,%12,%13,%14,%15,%16};\n"
:: "r"(tmem_addr),
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]));
}
__device__ __forceinline__ void tcgen05_st_32x32b_x32(
uint32_t tmem_addr, const uint32_t (&r)[32]) {
asm volatile("tcgen05.st.sync.aligned.32x32b.x32.b32 "
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
"%11,%12,%13,%14,%15,%16,%17,%18,%19,%20,"
"%21,%22,%23,%24,%25,%26,%27,%28,%29,%30,"
"%31,%32};\n"
:: "r"(tmem_addr),
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]),
"r"(r[16]),"r"(r[17]),"r"(r[18]),"r"(r[19]),
"r"(r[20]),"r"(r[21]),"r"(r[22]),"r"(r[23]),
"r"(r[24]),"r"(r[25]),"r"(r[26]),"r"(r[27]),
"r"(r[28]),"r"(r[29]),"r"(r[30]),"r"(r[31]));
}
__device__ __forceinline__ void tcgen05_commit1_lead(uint32_t lead, uint32_t mbar_smem_addr) {
asm volatile(
"{\n\t"
".reg .pred q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"@q tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%1];\n\t"
"}\n"
:: "r"(lead), "r"(mbar_smem_addr));
}
__device__ __forceinline__ void tcgen05_wait_st() {
asm volatile("tcgen05.wait::st.sync.aligned;\n" ::: "memory");
}
__device__ __forceinline__ void tcgen05_fence_before_thread_sync() {
asm volatile("tcgen05.fence::before_thread_sync;\n" ::: "memory");
}
__device__ __forceinline__
void tma_load_2d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int coord_x, int coord_y) {
asm volatile(
"cp.async.bulk.tensor.2d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(coord_x), "r"(coord_y)
: "memory");
}
__device__ __forceinline__
void tma_load_3d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int c0, int c1, int c2) {
asm volatile(
"cp.async.bulk.tensor.3d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2)
: "memory");
}
__device__ __forceinline__
void tma_load_4d(uint32_t smem_dst, const void* tensormap_ptr,
uint32_t mbar_smem, int c0, int c1, int c2, int c3) {
asm volatile(
"cp.async.bulk.tensor.4d.shared::cluster.global.tile"
".mbarrier::complete_tx::bytes"
" [%0], [%1, {%3, %4, %5, %6}], [%2];\n"
:: "r"(smem_dst), "l"(tensormap_ptr),
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2), "r"(c3)
: "memory");
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_dealloc(uint32_t tmem_addr,
uint32_t n_cols) {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_dealloc: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile(
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols));
} else {
asm volatile(
"tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;\n"
:: "r"(tmem_addr), "r"(n_cols));
}
}
__device__ __forceinline__
void tma_store_2d(const void* tensormap_ptr, int coord_x, int coord_y,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.2d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2}], [%3];\n"
:: "l"(tensormap_ptr), "r"(coord_x), "r"(coord_y),
"r"(smem_src)
: "memory");
}
__device__ __forceinline__
void tma_store_3d(const void* tensormap_ptr, int c0, int c1, int c2,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.3d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2, %3}], [%4];\n"
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(smem_src)
: "memory");
}
__device__ __forceinline__
void tma_store_4d(const void* tensormap_ptr, int c0, int c1, int c2, int c3,
uint32_t smem_src) {
asm volatile(
"cp.async.bulk.tensor.4d.global.shared::cta.tile.bulk_group"
" [%0, {%1, %2, %3, %4}], [%5];\n"
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(c3), "r"(smem_src)
: "memory");
}
inline cudaError_t make_tma_2d_tiled(
CUtensorMap* out,
const void* ptr, int rows, int cols, int box_rows, int box_cols,
int elem_bytes, CUtensorMapDataType dtype,
CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B,
CUtensorMapL2promotion l2 = CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CUtensorMapFloatOOBfill oob = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) {
uint64_t globalDim[2] = { (uint64_t)cols, (uint64_t)rows };
uint64_t globalStrides[1] = { (uint64_t)cols * (uint64_t)elem_bytes };
uint32_t boxDim[2] = { (uint32_t)box_cols, (uint32_t)box_rows };
uint32_t elemStrides[2] = { 1u, 1u };
CUresult r = cuTensorMapEncodeTiled(
out, dtype, 2,
const_cast<void*>(ptr), globalDim, globalStrides,
boxDim, elemStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle, l2, oob);
return (r == CUDA_SUCCESS) ? cudaSuccess : cudaErrorInvalidValue;
}
__device__ __forceinline__
void cp_async_bulk_commit_group() {
asm volatile("cp.async.bulk.commit_group;\n" ::: "memory");
}
template <int N>
__device__ __forceinline__
void cp_async_bulk_wait_group_read() {
asm volatile("cp.async.bulk.wait_group.read %0;\n" :: "n"(N) : "memory");
}
__device__ __forceinline__
void mbarrier_init(uint32_t mbar_smem, uint32_t arrive_count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n"
:: "r"(mbar_smem), "r"(arrive_count) : "memory");
}
template <int CTA_GROUP>
__device__ __forceinline__ void tcgen05_relinquish_alloc_permit() {
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
"tcgen05_relinquish_alloc_permit: CTA_GROUP must be 1 or 2");
if constexpr (CTA_GROUP == 1) {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;\n" ::);
} else {
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::2.sync.aligned;\n" ::);
}
}
__device__ __forceinline__
void fence_mbarrier_init_release_cluster() {
asm volatile("fence.mbarrier_init.release.cluster;\n" ::: "memory");
}
__device__ __forceinline__ void tcgen05_mma_f16_ss_lead(uint32_t lead,
uint32_t tmem_c, uint64_t desc_a, uint64_t desc_b, uint32_t idesc,
bool enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p, q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"setp.ne.b32 p, %5, 0;\n\t"
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], %2, %3, %4, {%6, %7, %8, %9}, p;\n\t"
"}\n"
:: "r"(lead), "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(idesc),
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
}
__device__ __forceinline__ void tcgen05_mma_f16_ts_1sm_lead(uint32_t lead,
uint32_t tmem_c, uint32_t tmem_a, uint64_t desc_b, uint32_t idesc,
bool enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p, q;\n\t"
"setp.ne.b32 q, %0, 0;\n\t"
"setp.ne.b32 p, %5, 0;\n\t"
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], [%2], %3, %4, {%6, %7, %8, %9}, p;\n\t"
"}\n"
:: "r"(lead), "r"(tmem_c), "r"(tmem_a), "l"(desc_b), "r"(idesc),
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
}
enum class SmemSwizzleBlackwell : uint32_t {
None = 0,
B128_32atom = 1,
B128 = 2,
B64 = 4,
B32 = 6,
};
__device__ __host__ __forceinline__ uint64_t build_smem_desc_blackwell(
uint32_t smem_addr,
uint32_t stride_byte_offset,
uint32_t leading_byte_offset,
SmemSwizzleBlackwell swizzle = SmemSwizzleBlackwell::B128,
uint32_t base_offset = 0) {
uint64_t d = 0;
d |= static_cast<uint64_t>((smem_addr >> 4) & 0x3FFF);
d |= static_cast<uint64_t>((leading_byte_offset >> 4) & 0x3FFF) << 16;
d |= static_cast<uint64_t>((stride_byte_offset >> 4) & 0x3FFF) << 32;
d |= static_cast<uint64_t>(1) << 46;
d |= (static_cast<uint64_t>(base_offset) & 0x7) << 49;
d |= static_cast<uint64_t>(static_cast<uint32_t>(swizzle) & 0x7) << 61;
return d;
}
__device__ __forceinline__
uint32_t elect_one_sync() {
uint32_t elected;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync %0|p, 0xffffffff;\n\t"
"selp.b32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(elected));
return elected;
}
__device__ __forceinline__
uint32_t elect_one_sync(uint32_t membermask) {
uint32_t elected;
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"elect.sync %0|p, %1;\n\t"
"selp.b32 %0, 1, 0, p;\n\t"
"}\n"
: "=r"(elected) : "r"(membermask));
return elected;
}
template <int N>
__device__ __forceinline__
void setmaxnreg_dec() {
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
"setmaxnreg_dec: N must be in [24, 256] and a multiple of 8");
asm volatile("setmaxnreg.dec.sync.aligned.u32 %0;\n"
:: "n"(N) : "memory");
}
template <int N>
__device__ __forceinline__
void setmaxnreg_inc() {
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
"setmaxnreg_inc: N must be in [24, 256] and a multiple of 8");
asm volatile("setmaxnreg.inc.sync.aligned.u32 %0;\n"
:: "n"(N) : "memory");
}
__device__ __forceinline__
uint32_t cvt_f32x2_to_bf16x2(float a, float b) {
uint32_t r;
asm volatile("cvt.rn.bf16x2.f32 %0, %2, %1;\n"
: "=r"(r) : "f"(a), "f"(b));
return r;
}
namespace {
__device__ __forceinline__ uint64_t f32x2_bits(float2 v) {
uint64_t b; __builtin_memcpy(&b, &v, 8); return b;
}
__device__ __forceinline__ float2 f32x2_make(uint64_t b) {
float2 v; __builtin_memcpy(&v, &b, 8); return v;
}
}
__device__ __forceinline__ float2 fmul2(float2 a, float2 b) {
uint64_t d;
asm volatile("mul.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 fadd2(float2 a, float2 b) {
uint64_t d;
asm volatile("add.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 ffma2(float2 a, float2 b, float2 c) {
uint64_t d;
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
: "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)), "l"(f32x2_bits(c)));
return f32x2_make(d);
}
__device__ __forceinline__ float2 f32x2_splat(float s) { return make_float2(s, s); }
__device__ __forceinline__ float ex2_approx_f32(float z) {
float d;
asm volatile("ex2.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(z));
return d;
}
__device__ __forceinline__ float2 ex2_emu_f32x2(float x, float y) {
uint32_t ox, oy;
asm volatile(
"{\n\t"
".reg .f32 f1,f2,f3,f4,f5,f6,f7;\n\t"
".reg .b64 l1,l2,l3,l4,l5,l6,l7,l8,l9,l10;\n\t"
".reg .s32 r1,r2,r3,r4,r5,r6,r7,r8;\n\t"
"max.f32 f1, %2, 0fC2FE0000;\n\t"
"max.f32 f2, %3, 0fC2FE0000;\n\t"
"mov.b64 l1, {f1, f2};\n\t"
"mov.f32 f3, 0f4B400000;\n\t"
"mov.b64 l2, {f3, f3};\n\t"
"add.rm.f32x2 l7, l1, l2;\n\t"
"sub.rn.f32x2 l8, l7, l2;\n\t"
"sub.rn.f32x2 l9, l1, l8;\n\t"
"mov.f32 f7, 0f3D9DF09D;\n\t"
"mov.b64 l6, {f7, f7};\n\t"
"mov.f32 f6, 0f3E6906A4;\n\t"
"mov.b64 l5, {f6, f6};\n\t"
"mov.f32 f5, 0f3F31F519;\n\t"
"mov.b64 l4, {f5, f5};\n\t"
"mov.f32 f4, 0f3F800000;\n\t"
"mov.b64 l3, {f4, f4};\n\t"
"fma.rn.f32x2 l10, l9, l6, l5;\n\t"
"fma.rn.f32x2 l10, l10, l9, l4;\n\t"
"fma.rn.f32x2 l10, l10, l9, l3;\n\t"
"mov.b64 {r1, r2}, l7;\n\t"
"mov.b64 {r3, r4}, l10;\n\t"
"shl.b32 r5, r1, 23;\n\t"
"add.s32 r7, r5, r3;\n\t"
"shl.b32 r6, r2, 23;\n\t"
"add.s32 r8, r6, r4;\n\t"
"mov.b32 %0, r7;\n\t"
"mov.b32 %1, r8;\n\t"
"}\n"
: "=r"(ox), "=r"(oy) : "f"(x), "f"(y));
float2 r; __builtin_memcpy(&r.x, &ox, 4); __builtin_memcpy(&r.y, &oy, 4);
return r;
}
__device__ __forceinline__ float rcp_approx_ftz_f32(float x) {
float d;
asm("rcp.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(x));
return d;
}
__device__ __forceinline__ uint32_t make_idesc_table44(
int M, int N,
uint32_t dtype, uint32_t atype, uint32_t btype,
bool transpose_a = false, bool transpose_b = false,
bool negate_a = false, bool negate_b = false) {
uint32_t idesc = 0;
idesc |= (dtype & 0x3) << 4;
idesc |= (atype & 0x7) << 7;
idesc |= (btype & 0x7) << 10;
idesc |= (negate_a ? 1u : 0u) << 13;
idesc |= (negate_b ? 1u : 0u) << 14;
idesc |= (transpose_a ? 1u : 0u) << 15;
idesc |= (transpose_b ? 1u : 0u) << 16;
idesc |= ((static_cast<uint32_t>(N) >> 3) & 0x3F) << 17;
idesc |= ((static_cast<uint32_t>(M) >> 4) & 0x1F) << 24;
return idesc;
}
__device__ __forceinline__ uint32_t make_idesc_bf16_f32(
int M, int N, bool ta = false, bool tb = false) {
return make_idesc_table44(M, N, 1,
1, 1, ta, tb);
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x16(
uint32_t tmem_addr, uint32_t (&r)[16]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15}, [%16];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x32(
uint32_t tmem_addr, uint32_t (&r)[32]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x32.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31}, [%32];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x64(
uint32_t tmem_addr, uint32_t (&r)[64]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x64.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
"%60,%61,%62,%63}, [%64];\n"
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31]),
"=r"(r[32]),"=r"(r[33]),"=r"(r[34]),"=r"(r[35]),
"=r"(r[36]),"=r"(r[37]),"=r"(r[38]),"=r"(r[39]),
"=r"(r[40]),"=r"(r[41]),"=r"(r[42]),"=r"(r[43]),
"=r"(r[44]),"=r"(r[45]),"=r"(r[46]),"=r"(r[47]),
"=r"(r[48]),"=r"(r[49]),"=r"(r[50]),"=r"(r[51]),
"=r"(r[52]),"=r"(r[53]),"=r"(r[54]),"=r"(r[55]),
"=r"(r[56]),"=r"(r[57]),"=r"(r[58]),"=r"(r[59]),
"=r"(r[60]),"=r"(r[61]),"=r"(r[62]),"=r"(r[63])
: "r"(tmem_addr));
}
__device__ __forceinline__ void tcgen05_ld_32x32b_x128(
uint32_t tmem_addr, uint32_t (&r)[128]) {
asm volatile("tcgen05.ld.sync.aligned.32x32b.x128.b32 "
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
"%60,%61,%62,%63,%64,%65,%66,%67,%68,%69,"
"%70,%71,%72,%73,%74,%75,%76,%77,%78,%79,"
"%80,%81,%82,%83,%84,%85,%86,%87,%88,%89,"
"%90,%91,%92,%93,%94,%95,%96,%97,%98,%99,"
"%100,%101,%102,%103,%104,%105,%106,%107,%108,%109,"
"%110,%111,%112,%113,%114,%115,%116,%117,%118,%119,"
"%120,%121,%122,%123,%124,%125,%126,%127}, [%128];\n"
: "=r"(r[ 0]),"=r"(r[ 1]),"=r"(r[ 2]),"=r"(r[ 3]),
"=r"(r[ 4]),"=r"(r[ 5]),"=r"(r[ 6]),"=r"(r[ 7]),
"=r"(r[ 8]),"=r"(r[ 9]),"=r"(r[ 10]),"=r"(r[ 11]),
"=r"(r[ 12]),"=r"(r[ 13]),"=r"(r[ 14]),"=r"(r[ 15]),
"=r"(r[ 16]),"=r"(r[ 17]),"=r"(r[ 18]),"=r"(r[ 19]),
"=r"(r[ 20]),"=r"(r[ 21]),"=r"(r[ 22]),"=r"(r[ 23]),
"=r"(r[ 24]),"=r"(r[ 25]),"=r"(r[ 26]),"=r"(r[ 27]),
"=r"(r[ 28]),"=r"(r[ 29]),"=r"(r[ 30]),"=r"(r[ 31]),
"=r"(r[ 32]),"=r"(r[ 33]),"=r"(r[ 34]),"=r"(r[ 35]),
"=r"(r[ 36]),"=r"(r[ 37]),"=r"(r[ 38]),"=r"(r[ 39]),
"=r"(r[ 40]),"=r"(r[ 41]),"=r"(r[ 42]),"=r"(r[ 43]),
"=r"(r[ 44]),"=r"(r[ 45]),"=r"(r[ 46]),"=r"(r[ 47]),
"=r"(r[ 48]),"=r"(r[ 49]),"=r"(r[ 50]),"=r"(r[ 51]),
"=r"(r[ 52]),"=r"(r[ 53]),"=r"(r[ 54]),"=r"(r[ 55]),
"=r"(r[ 56]),"=r"(r[ 57]),"=r"(r[ 58]),"=r"(r[ 59]),
"=r"(r[ 60]),"=r"(r[ 61]),"=r"(r[ 62]),"=r"(r[ 63]),
"=r"(r[ 64]),"=r"(r[ 65]),"=r"(r[ 66]),"=r"(r[ 67]),
"=r"(r[ 68]),"=r"(r[ 69]),"=r"(r[ 70]),"=r"(r[ 71]),
"=r"(r[ 72]),"=r"(r[ 73]),"=r"(r[ 74]),"=r"(r[ 75]),
"=r"(r[ 76]),"=r"(r[ 77]),"=r"(r[ 78]),"=r"(r[ 79]),
"=r"(r[ 80]),"=r"(r[ 81]),"=r"(r[ 82]),"=r"(r[ 83]),
"=r"(r[ 84]),"=r"(r[ 85]),"=r"(r[ 86]),"=r"(r[ 87]),
"=r"(r[ 88]),"=r"(r[ 89]),"=r"(r[ 90]),"=r"(r[ 91]),
"=r"(r[ 92]),"=r"(r[ 93]),"=r"(r[ 94]),"=r"(r[ 95]),
"=r"(r[ 96]),"=r"(r[ 97]),"=r"(r[ 98]),"=r"(r[ 99]),
"=r"(r[100]),"=r"(r[101]),"=r"(r[102]),"=r"(r[103]),
"=r"(r[104]),"=r"(r[105]),"=r"(r[106]),"=r"(r[107]),
"=r"(r[108]),"=r"(r[109]),"=r"(r[110]),"=r"(r[111]),
"=r"(r[112]),"=r"(r[113]),"=r"(r[114]),"=r"(r[115]),
"=r"(r[116]),"=r"(r[117]),"=r"(r[118]),"=r"(r[119]),
"=r"(r[120]),"=r"(r[121]),"=r"(r[122]),"=r"(r[123]),
"=r"(r[124]),"=r"(r[125]),"=r"(r[126]),"=r"(r[127])
: "r"(tmem_addr));
}
__device__ __forceinline__
uint32_t smem_ptr_u32(const void* ptr) {
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
}
__device__ __forceinline__
void sts_f32(uint32_t smem_addr, float val) {
asm volatile("st.shared.f32 [%0], %1;" :: "r"(smem_addr), "f"(val) : "memory");
}
@@ -28,31 +28,10 @@ void register_rms_norm(pybind11::module_ &);
void register_layer_norm(pybind11::module_ &);
void register_gemm(pybind11::module_ &);
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
extern std::vector<torch::Tensor> block_sparse_sm100a_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse);
extern std::vector<torch::Tensor> block_sparse_sm100a_blk128_fwd(
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
double sm_scale, bool need_lse);
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "FastVideo CUDA Kernels";
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
m.def("block_sparse_sm100a_fwd",
torch::wrap_pybind_function(block_sparse_sm100a_fwd),
"VSA block-sparse attention forward, 64-token blocks (Blackwell sm100a)");
m.def("block_sparse_sm100a_blk128_fwd",
torch::wrap_pybind_function(block_sparse_sm100a_blk128_fwd),
"VSA block-sparse attention forward, 128-token blocks (Blackwell sm100a)");
#endif
#ifdef TK_COMPILE_ST_ATTN
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention (Hopper)");
#endif
@@ -1,4 +1,4 @@
"""VSA-128/256 block-sparse attention wrappers.
"""VSA-256 block-sparse attention wrapper.
The default 256-block path is Triton: it expands the logical 256-block map
to the existing 64-block Triton kernel via a dense 4x4 expansion per logical
@@ -7,8 +7,8 @@ edge ("route A"), and requires no optional dependencies.
The FA4 CuTe block-sparse fastpath (intended for Blackwell sm_100+) is
*opt-in* via ``FASTVIDEO_VSA_CUTEDSL=1``. It routes to
:mod:`fastvideo_kernel.block_sparse_attn_cute_fwd`, which natively operates
on 128-token Q/KV blocks (the 256 wrapper expands its logical KV map and
sizes into that physical representation). The CuTe kernel
on 128-token KV blocks (this wrapper expands the logical 256-block map /
sizes into that physical 128-block representation). The CuTe kernel
(``flash_attn.cute`` with block-sparsity) is an optional dependency,
imported lazily only when this fastpath is selected.
@@ -35,7 +35,7 @@ _KV_BLOCK_TRITON = 64 # Existing Triton path uses 64-token KV blocks.
def _resolve_backend() -> str:
"""Pick the backend for the 128/256-block VSA paths.
"""Pick the backend for the 256-block VSA path.
Default is Triton (no optional deps). The FA4 CuTe fastpath is opt-in
via ``FASTVIDEO_VSA_CUTEDSL=1`` and requires the optional FA4 CuTe
@@ -49,26 +49,6 @@ def _resolve_backend() -> str:
return "triton"
def _expand_mask_and_sizes_128_to_64(
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Expand a [B, H, Qb128, KVb128] map to 64-token Triton tiles."""
expanded_mask = logical_mask_128.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
sizes_i32 = logical_kv_sizes_128.to(torch.int32)
offsets = torch.tensor(
[0, _KV_BLOCK_TRITON],
dtype=torch.int32,
device=sizes_i32.device,
)
expanded_sizes = torch.clamp(
sizes_i32[:, None] - offsets[None, :],
min=0,
max=_KV_BLOCK_TRITON,
).reshape(-1)
return expanded_mask, expanded_sizes
def _expand_mask_and_sizes_256_to_128(
logical_mask_256: torch.Tensor,
logical_kv_sizes_256: torch.Tensor,
@@ -132,63 +112,6 @@ def _triton_via_route_a(
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def _triton_via_route_a_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_mask_128: torch.Tensor,
logical_kv_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from .triton_kernels.index import map_to_index as triton_map_to_index
mask_64, sizes_64 = _expand_mask_and_sizes_128_to_64(logical_mask_128, logical_kv_sizes_128)
q2k_idx, q2k_num = triton_map_to_index(mask_64.to(torch.bool))
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
def block_sparse_attn_128(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, H, S, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
return _triton_via_route_a_128(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd
return block_sparse_attn_cute_fwd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_128_bshd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
logical_block_map_128: torch.Tensor,
logical_variable_block_sizes_128: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""VSA-128 sparse-branch entrypoint for [B, S, H, D] inputs."""
if logical_block_map_128.dim() == 3:
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
if _resolve_backend() == "triton":
out_bhsd, aux = _triton_via_route_a_128(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
logical_block_map_128,
logical_variable_block_sizes_128,
)
return out_bhsd.transpose(1, 2).contiguous(), aux
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
return block_sparse_attn_cute_fwd_bshd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
def block_sparse_attn_256(
q: torch.Tensor,
k: torch.Tensor,
@@ -1,18 +1,18 @@
"""FA4 CuTe-DSL block-sparse attention adapter.
"""CuTe-DSL block-sparse attention forward kernel.
This module adapts VSA's ``(block_map, variable_block_sizes)`` inputs into
FA4's forward and backward ``BlockSparseTensorsTorch`` representations.
FA4's public ``flash_attn_func`` owns the forward/backward autograd bridge.
Thin wrapper around `flash_attn.cute.interface._flash_attn_fwd` that adapts
VSA's `(block_map, variable_block_sizes)` inputs into FA4's
`BlockSparseTensorsTorch` representation and the per-KV-block validity mask.
Both [B, H, S, D] (BHSD) and [B, S, H, D] (BSHD) entrypoints are provided.
The BSHD variant is preferred from VSA-128/256 callers to avoid layout
The BSHD variant is preferred from VSA-256 callers to avoid layout
round-trips on the hot path.
The FA4 CuTe block-sparse kernel (``flash_attn.cute`` with
``block_sparsity``) is an *optional* dependency: it is imported lazily and
only exercised when the VSA-128/256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default path is Triton and does not require
it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
only exercised when the VSA-256 CuTe fastpath is explicitly selected
(``FASTVIDEO_VSA_CUTEDSL=1``). The default VSA-256 path is Triton and does
not require it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
"""
from __future__ import annotations
@@ -22,14 +22,13 @@ from typing import Tuple
import torch
_FA4_IMPORT_HINT = ("VSA-128/256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
"provides `flash_attn.cute` with block-sparsity support (plus "
"`nvidia-cutlass-dsl` and `quack-kernels`). This is an optional "
"dependency; the default path is Triton. Install the FA4 CuTe "
"dependency; the default VSA-256 path is Triton. Install the FA4 CuTe "
"build and set FASTVIDEO_VSA_CUTEDSL=1 to enable the CuTe fastpath.")
@functools.lru_cache(maxsize=1)
def _load_fa4_cute():
"""Lazily import the optional FA4 CuTe block-sparse symbols.
@@ -39,39 +38,14 @@ def _load_fa4_cute():
"""
try:
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
from flash_attn.cute.interface import (
_flash_attn_bwd,
_flash_attn_fwd,
flash_attn_func,
)
from flash_attn.cute.interface import _flash_attn_fwd
except ImportError as exc: # pragma: no cover - optional dependency
raise ImportError(_FA4_IMPORT_HINT) from exc
return BlockSparseTensorsTorch, flash_attn_func, _flash_attn_fwd, _flash_attn_bwd
return BlockSparseTensorsTorch, _flash_attn_fwd
# FA4's physical Q tile size; KV block size comes from the VSA caller.
_FA4_Q_BLOCK_SIZE = 128
class _SingleQStageLength(int):
"""Keep the real length while selecting FA4's one-stage Q128 path.
On sm_100 FA4 derives ``q_stage`` from ``max_seqlen_q > tile_m``. Its
kernel supports one 128-token Q stage, but the fixed-length public wrapper
does not expose that choice. VSA-128 must select it explicitly; otherwise
adjacent logical Q blocks are merged into a 256-token sparse block.
"""
def __mul__(self, other):
return type(self)(int(self) * int(other))
def __rmul__(self, other):
return type(self)(int(other) * int(self))
def __gt__(self, other):
if int(other) == _FA4_Q_BLOCK_SIZE:
return False
return int(self) > int(other)
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
_M_BLOCK_SIZE_DEFAULT = 128
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
@@ -90,12 +64,12 @@ def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
return triton_map_to_index(block_map)
def _choose_q_sparse_block_size(q_len: int, q_tile_size: int = _FA4_Q_BLOCK_SIZE) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > q_tile_size.
def _choose_q_sparse_block_size(q_len: int, m_block_size: int = _M_BLOCK_SIZE_DEFAULT) -> int:
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > m_block_size.
major, _ = torch.cuda.get_device_capability()
if major >= 10 and q_len > q_tile_size:
return 2 * q_tile_size
return q_tile_size
if major >= 10 and q_len > m_block_size:
return 2 * m_block_size
return m_block_size
def _aggregate_q_block_map(
@@ -160,35 +134,23 @@ def _build_vbs_mask_mod(kv_block_size: int):
return _vbs_mask_mod
def _build_sparse_tensors(
def _cute_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
q_len: int,
q_block_size: int,
kv_block_size: int,
need_backward: bool,
force_q_sparse_block_size: int | None = None,
) -> Tuple[object, object | None]:
"""Build the Q-owned forward and KV-owned backward sparse metadata.
``need_backward`` is False on inference-only calls: the backward metadata
is a pair of dense ``[B, H, kv_blocks, q_blocks]`` int32 index tensors that
FA4 keeps alive on its autograd ctx until backward runs, so building it
when nothing requires grad is pure overhead (~80 MiB per call at Wan-14B
720p shape).
"""
BlockSparseTensorsTorch, _, _, _ = _load_fa4_cute()
if force_q_sparse_block_size is None:
q_sparse_candidate = _choose_q_sparse_block_size(q_len)
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
else:
q_sparse_block_size = force_q_sparse_block_size
if q_sparse_block_size < q_block_size or q_sparse_block_size % q_block_size != 0:
raise ValueError("force_q_sparse_block_size must be a positive multiple of q_block_size")
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Internal: FA4 CuTe BSA fwd with BSHD inputs."""
BlockSparseTensorsTorch, _flash_attn_fwd = _load_fa4_cute()
q_sparse_candidate = _choose_q_sparse_block_size(q_bshd.shape[1])
q_sparse_block_size = max(
q_block_size,
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
)
sparse_map = _aggregate_q_block_map(
block_map,
q_sparse_block_size=q_sparse_block_size,
@@ -196,166 +158,35 @@ def _build_sparse_tensors(
)
kv_full = (variable_block_sizes == kv_block_size).view(1, 1, 1, -1)
kv_partial = ((variable_block_sizes > 0) & (variable_block_sizes < kv_block_size)).view(1, 1, 1, -1)
full_map = sparse_map & kv_full
mask_map = sparse_map & kv_partial
def from_maps(full_map: torch.Tensor, mask_map: torch.Tensor) -> object:
full_block_idx, full_block_cnt = _map_to_index(full_map.contiguous())
mask_block_idx, mask_block_cnt = _map_to_index(mask_map.contiguous())
return BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
)
full_block_idx, full_block_cnt = _map_to_index(full_map)
mask_block_idx, mask_block_cnt = _map_to_index(mask_map)
forward_sparse_tensors = from_maps(
sparse_map & kv_full,
sparse_map & kv_partial,
sparse_tensors = BlockSparseTensorsTorch(
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
block_size=(q_sparse_block_size, kv_block_size),
)
if not need_backward:
return forward_sparse_tensors, None
# FA4 backward is KV-owned: for each physical KV tile, list the sparse
# query tiles that selected it. Full and partial KV tiles stay separate
# so the token-level validity mask only runs for padded tiles.
backward_sparse_tensors = from_maps(
(sparse_map & kv_full).transpose(2, 3),
(sparse_map & kv_partial).transpose(2, 3),
)
return forward_sparse_tensors, backward_sparse_tensors
def _cute_attention_q128_forward(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
*,
need_backward: bool,
) -> Tuple[torch.Tensor, torch.Tensor, object | None]:
"""Run FA4 with one physical Q stage per logical VSA-128 block."""
_, _, flash_attn_fwd, _ = _load_fa4_cute()
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=_FA4_Q_BLOCK_SIZE,
kv_block_size=_FA4_Q_BLOCK_SIZE,
need_backward=need_backward,
force_q_sparse_block_size=_FA4_Q_BLOCK_SIZE,
)
out, lse = flash_attn_fwd(
# _flash_attn_fwd returns (out, lse, p, row_max); keep the first two.
out, lse = _flash_attn_fwd(
q_bshd,
k_bshd,
v_bshd,
tile_mn=(_FA4_Q_BLOCK_SIZE, _FA4_Q_BLOCK_SIZE),
max_seqlen_q=_SingleQStageLength(q_bshd.shape[1]),
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
block_sparse_tensors=forward_sparse_tensors,
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
mask_mod=_build_vbs_mask_mod(kv_block_size),
block_sparse_tensors=sparse_tensors,
aux_tensors=[variable_block_sizes],
causal=False,
return_lse=True,
)[:2]
return out, lse, backward_sparse_tensors
class _CuteAttentionQ128(torch.autograd.Function):
@staticmethod
def forward(ctx, q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes):
out, lse, backward_sparse_tensors = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=True,
)
ctx.save_for_backward(q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes)
ctx.backward_sparse_tensors = backward_sparse_tensors
ctx.mark_non_differentiable(lse)
ctx.set_materialize_grads(False)
return out, lse
@staticmethod
def backward(ctx, grad_out, grad_lse):
del grad_lse
q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes = ctx.saved_tensors
if grad_out is None:
grad_out = torch.zeros_like(out)
_, _, _, flash_attn_bwd = _load_fa4_cute()
dq, dk, dv = flash_attn_bwd(
q_bshd,
k_bshd,
v_bshd,
out,
grad_out.contiguous(),
lse,
softmax_scale=q_bshd.shape[-1]**-0.5,
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=ctx.backward_sparse_tensors,
)
return dq, dk, dv, None, None
def _cute_attention_q128(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
if need_backward:
return _CuteAttentionQ128.apply(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
out, lse, _ = _cute_attention_q128_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
need_backward=False,
)
return out, lse
def _cute_attention(
q_bshd: torch.Tensor,
k_bshd: torch.Tensor,
v_bshd: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Run FA4's autograd-enabled block-sparse attention with BSHD inputs."""
_, flash_attn_func, _, _ = _load_fa4_cute()
q_block_size = q_bshd.shape[1] // block_map.shape[2]
kv_block_size = k_bshd.shape[1] // block_map.shape[3]
if q_block_size == kv_block_size == _FA4_Q_BLOCK_SIZE:
return _cute_attention_q128(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
block_map,
variable_block_sizes,
q_len=q_bshd.shape[1],
q_block_size=q_block_size,
kv_block_size=kv_block_size,
need_backward=need_backward,
)
return flash_attn_func(
q_bshd,
k_bshd,
v_bshd,
mask_mod=_build_vbs_mask_mod(kv_block_size),
aux_tensors=[variable_block_sizes],
block_sparse_tensors=forward_sparse_tensors,
block_sparse_tensors_bwd=backward_sparse_tensors,
return_lse=True,
)
def block_sparse_attn_cute_fwd(
q: torch.Tensor,
k: torch.Tensor,
@@ -363,25 +194,34 @@ def block_sparse_attn_cute_fwd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Autograd-enabled CuTe block-sparse attention for [B, H, S, D]."""
"""CuTe forward-only block-sparse attention with [B, H, S, D] inputs."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[2] // block_map.shape[2]
kv_block_size = k.shape[2] // block_map.shape[3]
q_bshd = q.transpose(1, 2).contiguous()
k_bshd = k.transpose(1, 2).contiguous()
v_bshd = v.transpose(1, 2).contiguous()
out_bshd, lse = _cute_attention(
out_bshd, lse_bshd = _cute_forward(
q_bshd,
k_bshd,
v_bshd,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
out = out_bshd.transpose(1, 2).contiguous()
# FA4 already returns lse as [B, H, S], matching the Triton path's aux
# contract, so it needs no transpose. Detach before any further op: the
# value is informational and callers never backprop through it.
return out, lse.detach()
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
def block_sparse_attn_cute_fwd_bshd(
@@ -391,16 +231,27 @@ def block_sparse_attn_cute_fwd_bshd(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Autograd-enabled CuTe block-sparse attention for [B, S, H, D]."""
"""CuTe forward-only block-sparse attention with [B, S, H, D] inputs."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
q_block_size = q.shape[1] // block_map.shape[2]
kv_block_size = k.shape[1] // block_map.shape[3]
out, lse = _cute_attention(
out, lse_bshd = _cute_forward(
q,
k,
v,
block_map,
variable_block_sizes,
q_block_size=q_block_size,
kv_block_size=kv_block_size,
)
# lse is [B, H, S] regardless of the q/k/v layout; see above.
return out, lse.detach()
if lse_bshd is None:
lse = torch.empty(
(q.shape[0], q.shape[2], q.shape[1]),
dtype=torch.float32,
device=q.device,
)
else:
lse = lse_bshd.transpose(1, 2).contiguous()
return out, lse
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""sm_100a (Blackwell) CUDA block-sparse VSA forward.
A third backend behind the same VSA op as the Triton and CuTe-DSL paths. Forward only: it
returns ``(out, lse)`` with ``lse`` in exactly the form ``triton_block_sparse_attn_forward``
writes -- ``max(qk * qk_scale) + log2(l)``, ``[B, H, S]`` fp32 -- so
``block_sparse_attn_backward_triton`` runs against it unchanged.
The extension carries TWO instantiations of the kernel, for 64- and 128-token sparse blocks
(tile volumes 64 and 128 in ``build_vsa_metadata``); the block size is inferred from the
tensors and picks the op. Anything else falls back to Triton via ``is_supported``.
"""
from typing import Tuple
import torch
try:
# The pybind symbols live on fastvideo_kernel_ops, NOT on the _C package that contains it.
# `import fastvideo_kernel._C as _C` resolves to the namespace package, whose __init__ is
# empty, so hasattr() fails on a wheel install and the caller silently falls back with the
# kernel built and present.
from fastvideo_kernel._C import fastvideo_kernel_ops as _C
_FWD_BY_BLOCK = {
64: getattr(_C, "block_sparse_sm100a_fwd", None),
128: getattr(_C, "block_sparse_sm100a_blk128_fwd", None),
}
_HAS_VSA_SM100A = any(_FWD_BY_BLOCK.values())
except ImportError: # pragma: no cover - extension not built
_C = None
_FWD_BY_BLOCK = {}
_HAS_VSA_SM100A = False
_SM100 = (10, 0)
HEAD_DIM = 128
# Must match the -DVSA_BHSD the extension was compiled with (see CMakeLists).
BHSD = True
def _block_size(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> int:
num_blocks = variable_block_sizes.numel()
seqlen = q.shape[2] if BHSD else q.shape[1]
return 0 if num_blocks == 0 or seqlen % num_blocks else seqlen // num_blocks
def is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
"""True iff this build can run these tensors; otherwise the caller uses Triton.
Static facts only -- shapes, dtypes, arch, layout. Deliberately NO reads of tensor
contents: the previous ``int(variable_block_sizes.min())`` was a GPU->CPU sync on every
call, and the kernel no longer needs it (see below). This predicate must stay cheap
enough to sit on a per-layer dispatch path.
What the kernel accepts (and is tested to handle):
* q/k/v: contiguous 4-D bf16 CUDA tensors on an sm_100 device, head_dim 128, laid out
as compiled (BHSD here); seqlen == num_blocks * block with an EVEN num_blocks (a CTA
owns an adjacent pair of query blocks) and a 64- or 128-token build present.
* q2k_num: any per-row counts in [0, max_kv], NON-uniform across rows included. Rows
with count 0 produce exactly-zero output rows (and a finite LSE sentinel) rather
than attending anywhere -- so no ``.min()`` floor is required of the caller.
* q2k_idx: rows only need valid entries (in [0, num_blocks)) BELOW that row's count;
padding past the count (e.g. map_to_index's -1 fill) is never dereferenced. max_kv
(= q2k_idx.shape[-1]) must be >= 1, which the host launcher re-checks.
* variable_block_sizes: per-KV-block valid-token counts in [0, block]; keys at or past
a block's count are masked. Integer metadata is converted to int32/contiguous by
``block_sparse_attn_sm100a`` itself, so int64 inputs merely cost a cast.
"""
if not _HAS_VSA_SM100A or not q.is_cuda:
return False
if torch.cuda.get_device_capability(q.device) != _SM100:
return False
if q.dtype != torch.bfloat16 or q.dim() != 4 or q.shape[-1] != HEAD_DIM:
return False
if not q.is_contiguous():
return False
if _FWD_BY_BLOCK.get(_block_size(q, variable_block_sizes)) is None:
return False
# A CTA owns an adjacent pair of query blocks.
if variable_block_sizes.numel() % 2 != 0:
return False
# Metadata must be integer-typed so the wrapper's int32 conversion is value-preserving.
if not variable_block_sizes.is_cuda or variable_block_sizes.dtype not in (torch.int32, torch.int64):
return False
return True
def block_sparse_attn_sm100a(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
need_lse: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass. Returns ``(out, lse)``; ``out`` has q's layout."""
fwd = _FWD_BY_BLOCK[_block_size(q, variable_block_sizes)]
idx = q2k_idx.to(torch.int32).contiguous()
num = q2k_num.to(torch.int32).contiguous()
vbs = variable_block_sizes.to(torch.int32).contiguous()
sm_scale = 1.0 / (q.shape[-1]**0.5)
res = fwd(q.contiguous(), k.contiguous(), v.contiguous(), None,
idx, num, vbs, sm_scale, need_lse)
return (res[0], res[1]) if need_lse else (res[0], None)
+22 -24
View File
@@ -2,8 +2,6 @@ import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn_256 import (
block_sparse_attn_128,
block_sparse_attn_128_bshd,
block_sparse_attn_256,
block_sparse_attn_256_bshd,
)
@@ -76,13 +74,12 @@ def video_sparse_attn(
Dispatches the sparse branch by ``block_elements = prod(block_size)``:
- 64 -> existing TK/Triton path (see ``block_sparse_attn_from_indices``).
- 128 -> Triton fallback or CuTe FA4 block-sparse attention.
- 256 -> CuTe FA4 block-sparse attention (see ``block_sparse_attn_256``).
Backend overrides:
- ``FASTVIDEO_VSA_TRITON=1`` forces Triton in either path.
- ``FASTVIDEO_VSA_TK=1`` prefers the sm_90 TK kernel in the 64-block path.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 128/256-block paths.
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 256-block path.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
@@ -122,9 +119,8 @@ def video_sparse_attn(
# Sparse branch (fused Triton topk mask)
mask = fused_topk_mask(scores, topk)
if block_elements in (128, 256):
attention = block_sparse_attn_128 if block_elements == 128 else block_sparse_attn_256
out_s = attention(q, k, v, mask, variable_block_sizes)[0]
if block_elements == 256:
out_s = block_sparse_attn_256(q, k, v, mask, variable_block_sizes)[0]
else:
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
@@ -146,14 +142,14 @@ def video_sparse_attn_bshd(
"""VSA entrypoint for [B, S, H, D] tensors.
Avoids the BHSD<->BSHD round-trip that ``video_sparse_attn`` performs on
the CuTe 128/256-block paths; the 64-block path still expects BHSD and is not
the CuTe 256-block path; the 64-block path still expects BHSD and is not
supported here.
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
if block_elements not in (128, 256):
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=128 or 256 "
if block_elements != 256:
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=256 "
f"(got {block_elements}); use video_sparse_attn for the 64-block path.")
batch, q_seq_len, heads, dim = q.shape
@@ -175,15 +171,19 @@ def video_sparse_attn_bshd(
raise ValueError(f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
f"got {q_variable_block_sizes.numel()}")
# Compression branch (BSHD-native: match fused_block_mean's semantics).
# Padding values are expected to be zero; gradients are broadcast across
# the full padded block, just like the BHSD fused common path.
# Compression branch (BSHD-native: mean over the 256-token axis).
token_idx = torch.arange(block_elements, device=q.device, dtype=torch.int32)
q_token_valid = (token_idx.view(1, -1) < q_variable_block_sizes.view(-1,
1)).view(1, q_num_blocks, block_elements, 1, 1)
kv_token_valid = (token_idx.view(1, -1) < variable_block_sizes.view(-1,
1)).view(1, kv_num_blocks, block_elements, 1, 1)
q_c = q.view(batch, q_num_blocks, block_elements, heads, dim)
k_c = k.view(batch, kv_num_blocks, block_elements, heads, dim)
v_c = v.view(batch, kv_num_blocks, block_elements, heads, dim)
q_c = (q_c.float().sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = (k_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = (v_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_c = ((q_c.float() * q_token_valid).sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
k_c = ((k_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
v_c = ((v_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
q_ch = q_c.permute(0, 2, 1, 3).contiguous()
k_ch = k_c.permute(0, 2, 1, 3).contiguous()
v_ch = v_c.permute(0, 2, 1, 3).contiguous()
@@ -195,15 +195,13 @@ def video_sparse_attn_bshd(
# Sparse branch (fused Triton topk mask + CuTe BSHD).
mask = fused_topk_mask(scores, topk)
attention = block_sparse_attn_128_bshd if block_elements == 128 else block_sparse_attn_256_bshd
out_s, _ = attention(q, k, v, mask, variable_block_sizes)
out_s, _ = block_sparse_attn_256_bshd(q, k, v, mask, variable_block_sizes)
# Out-of-place: ``out_s`` is the tensor FA4's autograd node saved for its
# backward, so mutating it in place invalidates the graph.
out_view = out_s.view(batch, q_num_blocks, block_elements, heads, dim)
out = out_s
out_view = out.view(batch, q_num_blocks, block_elements, heads, dim)
if compress_attn_weight is not None:
gate_view = compress_attn_weight.view(batch, q_num_blocks, block_elements, heads, dim)
out = out_view + out_c_blk.unsqueeze(2) * gate_view
out_view.add_(out_c_blk.unsqueeze(2) * gate_view)
else:
out = out_view + out_c_blk.unsqueeze(2)
return out.view(batch, q_seq_len, heads, dim)
out_view.add_(out_c_blk.unsqueeze(2))
return out
@@ -237,12 +237,7 @@ def _attn_bwd_dkdv(
# Load m before computing qk to reduce pipeline stall.
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
m = tl.load(M + offs_m)
# Recompute logits exactly as the forward does: raw bf16 operands into
# the dot, fp32 scale after accumulation. A bf16 pre-scaled K perturbs
# the recomputed logits relative to the saved M by an error
# proportional to |logit|, which exp2 amplifies into arbitrarily wrong
# probabilities at large activations.
qkT = tl.dot(k, qT) * (sm_scale * 1.4426950408889634)
qkT = tl.dot(k, qT)
pT = tl.math.exp2(qkT - m[None, :])
mask = tl.arange(0, BLOCK_N1) < block_size
pT = tl.where(mask[:, None], pT, 0.0)
@@ -273,7 +268,6 @@ def _attn_bwd_dq(
do,
m,
D,
sm_scale,
# shared by Q/K/V/DO.
q2k_index,
q2k_num,
@@ -321,7 +315,7 @@ def _attn_bwd_dq(
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
kT = tl.load(kT_ptrs + block_sparse_offset)
vT = tl.load(vT_ptrs + block_sparse_offset)
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
qk = tl.dot(q, kT)
p = tl.math.exp2(qk - m)
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
mask = offs_in_block < block_size
@@ -330,7 +324,8 @@ def _attn_bwd_dq(
dp = tl.dot(do, vT).to(tl.float32)
ds = p * (dp - Di[:, None])
ds = ds.to(tl.bfloat16)
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
# Compute dQ.
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
dq += tl.dot(ds, tl.trans(kT))
# Increment pointers.
return dq
@@ -458,7 +453,6 @@ def _attn_bwd(
do,
m,
D, #
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
@@ -476,7 +470,7 @@ def _attn_bwd(
)
# Write back dQ.
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq *= sm_scale
dq *= LN2
tl.store(dq_ptrs, dq)
@@ -597,7 +591,6 @@ def _attn_bwd_dq_kernel(
Q,
K,
V,
sm_scale,
DO, #
DQ,
M,
@@ -670,7 +663,6 @@ def _attn_bwd_dq_kernel(
do,
m,
D,
sm_scale,
q2k_index,
q2k_num,
max_kv_blks,
@@ -688,7 +680,7 @@ def _attn_bwd_dq_kernel(
)
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
dq_acc *= sm_scale
dq_acc *= LN2
tl.store(dq_ptrs, dq_acc)
@@ -756,11 +748,9 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
dv = torch.empty_like(v)
BATCH, N_HEAD = q.shape[:2]
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
# K stays raw: the backward kernels apply sm_scale in fp32 after the dot,
# matching the forward's rounding exactly. (A bf16 pre-scaled K perturbs
# the recomputed logits vs the saved M; exp2 turns that into unboundedly
# wrong probabilities at large activations.)
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
arg_k = k
arg_k = arg_k * (sm_scale * RCP_LN2)
PRE_BLOCK = 64
assert Tq % PRE_BLOCK == 0
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
@@ -823,7 +813,6 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
q,
arg_k,
v,
sm_scale,
do,
dq,
M,
@@ -14,9 +14,7 @@ import math
import torch
VSA_TILE_SIZE = (4, 4, 4)
# 128 is served by the sm_100a CUDA backend (block_sparse_attn_sm100a); 64 and 256 by
# Triton and the CuTe-DSL path. A volume here only needs a backend that accepts it.
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 128, 256)
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 256)
def _canonicalize_device(device: torch.device | str) -> torch.device:
@@ -1,146 +0,0 @@
"""VSA-128 CuTe/Triton forward and backward parity on Blackwell."""
from __future__ import annotations
import math
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 128
_BLOCK_SIZE_3D = (2, 8, 8)
def _select_backend(monkeypatch, backend: str) -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
if backend == "cute":
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
else:
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
def _dense_sparse_reference(q, k, v, block_map, variable_block_sizes):
token_mask = block_map.repeat_interleave(_BLOCK, dim=2).repeat_interleave(_BLOCK, dim=3)
kv_valid = torch.arange(_BLOCK, device=k.device) < variable_block_sizes[:, None]
token_mask = token_mask & kv_valid.reshape(1, 1, 1, -1)
logits = torch.matmul(q.float(), k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
probabilities = torch.softmax(logits.masked_fill(~token_mask, float("-inf")), dim=-1)
return torch.matmul(probabilities, v.float()).to(q.dtype)
def _check(tag: str, expected: torch.Tensor, actual: torch.Tensor, avg_tol: float, rel_tol: float) -> None:
assert torch.isfinite(actual).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(expected, actual)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < avg_tol
assert max_rel < rel_tol
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
def test_vsa128_explicit_routes_forward_backward(backend: str, monkeypatch) -> None:
"""Adjacent Q128 blocks must keep independent routes instead of merging."""
_select_backend(monkeypatch, backend)
torch.manual_seed(53)
shape = (1, 1, 3 * _BLOCK, 128)
base = [torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3)]
grad_output = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
variable_block_sizes = torch.tensor([128, 91, 37], device="cuda", dtype=torch.int32)
block_map = torch.eye(3, device="cuda", dtype=torch.bool).view(1, 1, 3, 3)
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
actual, _ = block_sparse_attn_128(*actual_inputs, block_map, variable_block_sizes)
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
expected = _dense_sparse_reference(*reference_inputs, block_map, variable_block_sizes)
(expected * grad_output).sum().backward()
print(f"[vsa128-explicit-{backend}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_inputs, strict=True):
_check(name, reference.grad, candidate.grad, 2e-2, 0.5)
def _zero_kv_tail(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
valid = torch.arange(_BLOCK, device=x.device) < variable_block_sizes[:, None]
valid = valid.view(1, 1, -1, _BLOCK, 1).expand_as(x.view(1, x.shape[1], -1, _BLOCK, x.shape[-1]))
return x * valid.reshape_as(x).to(x.dtype)
@pytest.mark.cuda
@pytest.mark.parametrize("backend", ["cute", "triton"])
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa128_wrapper_forward_backward(backend: str, layout: str, monkeypatch) -> None:
_select_backend(monkeypatch, backend)
torch.manual_seed(59)
batch, heads, dim = 1, 2, 128
q_blocks, kv_blocks, topk = 3, 4, 2
q_shape = (batch, heads, q_blocks * _BLOCK, dim)
kv_shape = (batch, heads, kv_blocks * _BLOCK, dim)
q_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
kv_sizes = torch.tensor([128, 91, 37, 128], device="cuda", dtype=torch.int32)
q_sizes = torch.full((q_blocks, ), _BLOCK, device="cuda", dtype=torch.int32)
k_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
v_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
gate_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16) * 0.1
grad_output = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
if layout == "bhsd":
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
actual_gate = gate_base.detach().clone().requires_grad_()
actual = video_sparse_attn(
*actual_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=actual_gate,
)
actual_grads = actual_inputs
else:
bshd_inputs = [tensor.transpose(1, 2).contiguous().detach().requires_grad_()
for tensor in (q_base, k_base, v_base)]
bshd_gate = gate_base.transpose(1, 2).contiguous().detach().requires_grad_()
actual = video_sparse_attn_bshd(
*bshd_inputs,
kv_sizes,
q_sizes,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=bshd_gate,
).transpose(1, 2)
actual_grads = bshd_inputs
(actual * grad_output).sum().backward()
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
reference_gate = gate_base.detach().clone().requires_grad_()
expected = _torch_vsa256_reference(
*reference_inputs,
q_sizes,
kv_sizes,
topk,
compress_attn_weight=reference_gate,
)
(expected * grad_output).sum().backward()
print(f"[vsa128-wrapper-{backend}-{layout}]")
_check("out", expected, actual, 1e-3, 0.2)
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_grads, strict=True):
candidate_grad = candidate.grad if layout == "bhsd" else candidate.grad.transpose(1, 2)
_check(name, reference.grad, candidate_grad, 2e-2, 0.5)
actual_gate_grad = actual_gate.grad if layout == "bhsd" else bshd_gate.grad.transpose(1, 2)
_check("dgate", reference_gate.grad, actual_gate_grad, 1e-3, 0.2)
@@ -1,224 +0,0 @@
"""VSA-256 FA4 CuTe forward/backward parity for BHSD and BSHD APIs.
Covers the shapes the CuTe backward actually sees in production: the gated
compression branch (`compress_attn_weight`), partially filled Q tiles,
and q_len != kv_len. Also pins the inference fast path, which must skip the
KV-owned backward metadata without changing the forward result.
"""
from __future__ import annotations
from typing import Tuple
import pytest
import torch
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
_BLOCK = 256
_BLOCK_SIZE_3D = (4, 8, 8) # prod == 256
# Measured on GB200 (sm_100) with bf16 inputs: grads land around 1e-4 avg_abs
# and <=0.11 max_rel across every case below, so these leave ~10x headroom
# without being loose enough to hide a real regression.
_OUT_TOL = (1e-3, 0.2)
_GRAD_TOL = (1e-3, 0.25)
@pytest.fixture(autouse=True)
def _require_cute_backend(monkeypatch):
pytest.importorskip(
"flash_attn.cute.block_sparsity",
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
)
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
def _zero_pad_tail(x: torch.Tensor, var: torch.Tensor) -> torch.Tensor:
"""Zero the padded tail of every 256-token tile of a [B, H, S, D] tensor.
VSA callers scatter into a zeroed tile buffer, so padded slots are zero;
both the kernel and the reference rely on that.
"""
bsz, heads, _, dim = x.shape
blocks = var.numel()
token_idx = torch.arange(_BLOCK, device=x.device, dtype=torch.int32)
valid = (token_idx.view(1, -1) < var.view(-1, 1)).view(1, 1, blocks, _BLOCK, 1)
valid = valid.expand(bsz, heads, blocks, _BLOCK, dim).reshape_as(x)
return x * valid.to(x.dtype)
def _make_inputs(
q_blocks: int,
kv_blocks: int,
kv_var: torch.Tensor,
q_var: torch.Tensor,
heads: int = 2,
dim: int = 128,
seed: int = 42,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
torch.manual_seed(seed)
device = torch.device("cuda")
dtype = torch.bfloat16
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=dtype)
k = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
v = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
return _zero_pad_tail(q, q_var), _zero_pad_tail(k, kv_var), _zero_pad_tail(v, kv_var), grad_out
def _check(tag: str, ref: torch.Tensor, got: torch.Tensor, tol: Tuple[float, float]) -> None:
assert torch.isfinite(got).all().item(), f"{tag}: non-finite values"
avg_abs, max_rel = _metrics(ref, got)
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
assert avg_abs < tol[0], f"{tag}: avg_abs {avg_abs:.3e} >= {tol[0]:.3e}"
assert max_rel < tol[1], f"{tag}: max_rel {max_rel:.3e} >= {tol[1]:.3e}"
def _run_bhsd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = video_sparse_attn(qg, kg, vg, kv_var, q_var, topk, block_size=_BLOCK_SIZE_3D, compress_attn_weight=gate)
return out, (qg, kg, vg)
def _run_bshd(q, k, v, kv_var, q_var, topk, gate=None):
qg, kg, vg = (t.transpose(1, 2).contiguous().requires_grad_(True) for t in (q, k, v))
gate_bshd = None if gate is None else gate.transpose(1, 2).contiguous()
out = video_sparse_attn_bshd(qg,
kg,
vg,
kv_var,
q_var,
topk,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=gate_bshd)
return out.transpose(1, 2), (qg, kg, vg)
def _reference(q, k, v, q_var, kv_var, topk, gate=None):
qr, kr, vr = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out = _torch_vsa256_reference(qr, kr, vr, q_var, kv_var, topk, compress_attn_weight=gate)
return out, (qr, kr, vr)
def _compare(tag, layout, q, k, v, kv_var, q_var, topk, grad_out, gate=None):
runner = _run_bhsd if layout == "bhsd" else _run_bshd
out, (qg, kg, vg) = runner(q, k, v, kv_var, q_var, topk, gate=gate)
(out * grad_out).sum().backward()
grads = [g.grad if g.grad.dim() == 4 and layout == "bhsd" else g.grad for g in (qg, kg, vg)]
if layout == "bshd":
grads = [g.transpose(1, 2) for g in grads]
out_ref, refs = _reference(q, k, v, q_var, kv_var, topk, gate=gate)
(out_ref * grad_out).sum().backward()
print(f"[{tag}-{layout}]")
_check("out", out_ref, out, _OUT_TOL)
for name, ref, got in zip(("dq", "dk", "dv"), refs, grads):
_check(name, ref.grad, got, _GRAD_TOL)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_forward_backward_vs_torch_ref(layout: str) -> None:
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var)
_compare("vsa256-cute", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_with_compress_gate(layout: str) -> None:
"""The gated compression branch is what Wan and MiniMax-H3 actually run.
It is also the branch that composes the sparse output with the compression
output, so it is the one that breaks if that composition mutates FA4's
saved output in place.
"""
kv_var = torch.tensor([256, 200, 256, 91], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var, seed=7)
gate = torch.randn_like(q) * 0.1
_compare("vsa256-cute-gated", layout, q, k, v, kv_var, q_var, 2, grad_out, gate=gate)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_partial_q_blocks(layout: str) -> None:
"""Q tiles that are not full: only the compression divisor depends on it,
but it is the one axis the existing coverage held constant."""
kv_var = torch.tensor([256, 128, 256], dtype=torch.int32, device="cuda")
q_var = torch.tensor([256, 61, 199], dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(3, 3, kv_var, q_var, seed=11)
_compare("vsa256-cute-partial-q", layout, q, k, v, kv_var, q_var, 2, grad_out)
@pytest.mark.cuda
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
def test_vsa256_cute_backward_cross_q_kv(layout: str) -> None:
"""q_len != kv_len: forward has coverage, backward did not."""
kv_var = torch.tensor([256, 143, 256, 256, 88], dtype=torch.int32, device="cuda")
q_var = torch.full((2, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, grad_out = _make_inputs(2, 5, kv_var, q_var, seed=13)
_compare("vsa256-cute-cross", layout, q, k, v, kv_var, q_var, 3, grad_out)
@pytest.mark.cuda
def test_vsa256_cute_inference_matches_training_forward() -> None:
"""The KV-owned backward metadata is only built when something requires
grad. Skipping it must not perturb the forward result."""
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
q, k, v, _ = _make_inputs(3, 4, kv_var, q_var, seed=5)
with torch.no_grad():
out_infer = video_sparse_attn_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
kv_var,
q_var,
2,
block_size=_BLOCK_SIZE_3D,
compress_attn_weight=None,
)
out_train, _ = _run_bshd(q, k, v, kv_var, q_var, 2)
torch.testing.assert_close(out_infer, out_train.transpose(1, 2).detach(), rtol=0, atol=0)
@pytest.mark.cuda
def test_vsa256_cute_lse_is_bhs() -> None:
"""The aux return is [B, H, S] on both entrypoints, matching the Triton
path's contract."""
from fastvideo_kernel.block_sparse_attn_256 import (block_sparse_attn_256, block_sparse_attn_256_bshd)
device = torch.device("cuda")
heads, dim, q_blocks, kv_blocks = 2, 128, 3, 4
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
q = torch.randn(1, heads, sq, dim, device=device, dtype=torch.bfloat16)
k = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
v = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
vbs = torch.full((kv_blocks, ), _BLOCK, dtype=torch.int32, device=device)
mask = torch.zeros(1, heads, q_blocks, kv_blocks, dtype=torch.bool, device=device)
mask[..., :2] = True
_, lse_bhsd = block_sparse_attn_256(q, k, v, mask, vbs)
assert lse_bhsd.shape == (1, heads, sq), lse_bhsd.shape
_, lse_bshd = block_sparse_attn_256_bshd(
q.transpose(1, 2).contiguous(),
k.transpose(1, 2).contiguous(),
v.transpose(1, 2).contiguous(),
mask,
vbs,
)
assert lse_bshd.shape == (1, heads, sq), lse_bshd.shape
@@ -22,7 +22,6 @@ def _torch_vsa256_reference(
q_var: torch.Tensor,
kv_var: torch.Tensor,
topk_logical: int,
compress_attn_weight: torch.Tensor | None = None,
) -> torch.Tensor:
bsz, heads, _sq, dim = q.shape
q_blocks = q_var.numel()
@@ -56,8 +55,6 @@ def _torch_vsa256_reference(
logits = logits.masked_fill(~token_mask, float("-inf"))
prob = torch.softmax(logits, dim=-1)
out_s = torch.matmul(prob, vf).to(q.dtype)
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
return out_c + out_s
@@ -1,104 +0,0 @@
"""Regression: Triton block-sparse backward gradient parity at realistic activation scale.
The backward used to fold ``sm_scale / ln(2)`` into K in bf16 before the
exp2-based logit recompute. The bf16 rounding error on the pre-scaled K grows
proportionally to |logit| and exp2 amplifies it into exponentially wrong
probabilities, so dQ/dK/dV were correct at unit scale (every pre-existing test)
but off by orders of magnitude at real activation magnitudes.
This test sweeps the input scale and checks the Triton kernel's gradients
against an fp32 masked-dense SDPA reference. The unit-scale case is the
control (it passed even with the broken kernel); the large-scale cases are
the regression.
"""
import pytest
import torch
from fastvideo_kernel.block_sparse_attn import _map_to_index, block_sparse_attn_triton
from .utils import generate_block_sparse_mask_for_function
BLOCK = 64
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Same convention as test_vsa_varlen.py: every tensor here comes from the
global torch RNG and the checks use tight thresholds, so an unseeded run
would shift inputs whenever an earlier test file draws a different number
of randoms.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _dense_reference(q, k, v, block_mask):
"""fp32 masked-dense SDPA over the token-expanded block mask.
q/k/v: [B, H, S, D]; block_mask: [B, H, S // BLOCK, S // BLOCK] bool.
"""
qf, kf, vf = q.float(), k.float(), v.float()
token_mask = block_mask.repeat_interleave(BLOCK, dim=-2).repeat_interleave(BLOCK, dim=-1)
logits = torch.matmul(qf, kf.transpose(-2, -1)) * (q.shape[-1]**-0.5)
logits = logits.masked_fill(~token_mask, float("-inf"))
return torch.matmul(logits.softmax(dim=-1), vf)
@pytest.mark.cuda
@pytest.mark.parametrize("scale", [1.0, 4.0, 16.0])
def test_triton_backward_grad_parity_across_input_scales(scale: float) -> None:
"""Kernel dQ/dK/dV must stay within a few percent of the fp32 reference
regardless of input magnitude.
With the bf16 K pre-scaling bug, scale<=4.0 passes at this geometry while
scale=16.0 fails (measured on GB200: dq relative L2 error 5.9e-1 vs 6.9e-3
fixed); at larger geometries and real activation magnitudes the broken
kernel is off by orders of magnitude. The passing unit-scale case is
exactly how the bug survived the original test suite.
"""
if not torch.cuda.is_available():
pytest.skip("CUDA is required")
device = torch.device("cuda")
dtype = torch.bfloat16
batch, heads, dim = 1, 4, 128
num_blocks = 8
seq = num_blocks * BLOCK
q = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
k = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
v = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype)
grad_out = torch.randn_like(q)
block_mask = generate_block_sparse_mask_for_function(heads, num_blocks, num_blocks, k=3,
device=device).unsqueeze(0)
q2k_idx, q2k_num = _map_to_index(block_mask)
variable_block_sizes = torch.full((num_blocks, ), BLOCK, dtype=torch.int32, device=device)
q_ker, k_ker, v_ker = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out_ker, _ = block_sparse_attn_triton(q_ker, k_ker, v_ker, q2k_idx, q2k_num, variable_block_sizes)
(out_ker.float() * grad_out.float()).sum().backward()
q_ref, k_ref, v_ref = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
out_ref = _dense_reference(q_ref, k_ref, v_ref, block_mask)
(out_ref * grad_out.float()).sum().backward()
# Forward is exact at any scale; this pins the harness itself.
fwd_rel = ((out_ker.float() - out_ref).norm() / out_ref.norm()).item()
assert fwd_rel < 2e-2, f"scale={scale}: forward rel err {fwd_rel:.3e}"
for name, g_ker, g_ref in (
("dq", q_ker.grad, q_ref.grad),
("dk", k_ker.grad, k_ref.grad),
("dv", v_ker.grad, v_ref.grad),
):
assert torch.isfinite(g_ker).all().item(), f"scale={scale}: non-finite {name}"
ref_norm = g_ref.float().norm()
rel = ((g_ker.float() - g_ref.float()).norm() / ref_norm.clamp_min(1e-12)).item()
ratio = (g_ker.float().norm() / ref_norm.clamp_min(1e-12)).item()
print(f"scale={scale} {name}: rel_l2={rel:.4e} norm_ratio={ratio:.4f}")
assert rel < 5e-2, f"scale={scale}: {name} rel l2 err {rel:.3e} >= 5e-2"
assert 0.98 < ratio < 1.02, f"scale={scale}: {name} grad-norm ratio {ratio:.4f}"
-14
View File
@@ -23,20 +23,6 @@ from fastvideo_kernel.block_sparse_attn import (
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
@pytest.fixture(autouse=True)
def _seed_rng():
"""Pin the RNG so these cases do not depend on what ran before them.
Every tensor and every variable block size here comes from the global
torch RNG, and the gradient checks use a tight max_rel threshold. Without
a seed the inputs shift whenever an earlier test file draws a different
number of randoms, which surfaces as an unrelated-looking failure in
whichever case happens to land on unlucky data.
"""
torch.manual_seed(0)
torch.cuda.manual_seed_all(0)
def _reference_per_sequence(
q_list,
k_list,
+1 -6
View File
@@ -19,12 +19,7 @@ from fastvideo.attention.backends.abstract import (
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Every worker records the loaded FlashAttention implementation so a
# distributed profiling log contains one backend receipt per rank.
logger.info("Worker %s Using FlashAttention-%s backend",
os.environ.get("RANK", "0"),
fa_version,
local_main_process_only=False)
logger.info("Using FlashAttention-%s backend", fa_version)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
@@ -5,17 +5,13 @@ H3 runs one joint bidirectional attention over
``[text | condition keyframes | audio | generated video]``, so this
backend differs from the Wan-tuned ``video_sparse_attn``:
- Tiles are ``[segment-pure prefix chunks] + [3D video tiles]``; prefix
tiles never straddle segment boundaries. The tile size is selectable at
metadata build time: 256 tokens ``(4,8,8)`` (default) or 64 tokens
``(4,4,4)`` (see ``VSA_H3_TILE_SHAPES``).
- Tiles are ``[segment-pure prefix chunks] + [3D (4,8,8) video tiles]``;
prefix tiles never straddle segment boundaries.
- Selection is pure Python on pooled tile scores; the block-sparse kernel
consumes an explicit bool mask, so no kernel changes are needed.
- The compression branch is gated by ``to_gate_compress``, which the base
H3 checkpoint does not carry: the loader zero-initializes it, so
untrained inference is exactly pure sparse and finetuning can learn the
gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship
trained gates, which load and activate the branch.
- The compression branch is gated by ``to_gate_compress``, which the H3
checkpoint does not carry: the loader zero-initializes it, so untrained
inference is exactly pure sparse and finetuning can learn the gate.
- Non-video *queries* are always dense. Non-video *keys* are either
always-selected for every query ("exempt", default) or compete in
top-k under a FLOP-matched budget ("compete") — the ablation axis,
@@ -24,46 +20,22 @@ backend differs from the Wan-tuned ``video_sparse_attn``:
(``vsa_dense_first_n_steps``, ``vsa_dense_layers``) let mixed schedules
run the diffuse steps/layers dense while pushing the rest harder.
At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path
Targets sm10.x through the FA4 CuTe 256-tile path
(``FASTVIDEO_VSA_CUTEDSL=1``); the Triton 256→64 expansion is the
fallback and keeps identical mask semantics. At tile 64 the block map is
already at the kernels' native 64-token granularity, so both forward and
backward run the Triton block-sparse kernels directly (no expansion,
``FASTVIDEO_VSA_CUTEDSL`` does not apply). A third, opt-in route exists
for the tile-64 FORWARD only: ``FASTVIDEO_VSA_SM100A=1`` sends no-grad
forwards through the sm_100a CUDA block-sparse kernel
(``fastvideo_kernel.block_sparse_attn_sm100a``, upstream PR #1719 plus
our per-q-tile ``q2k_num`` fix) when the extension is built, the device
is sm_100, and the geometry qualifies; grad-tracking forwards and every
backward stay on Triton unchanged. If the env is set but a precondition
fails, the route logs one warning and falls back.
fallback and keeps identical mask semantics.
"""
import functools
import math
import os
from dataclasses import dataclass
from typing import Any
import torch
try:
from fastvideo_kernel.block_sparse_attn import block_sparse_attn as block_sparse_attn_64_bhsd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
from fastvideo_kernel.triton_kernels.index import map_to_index
except ImportError:
block_sparse_attn_64_bhsd = None
block_sparse_attn_256_bshd = None
map_to_index = None
try:
# Optional: only present in fastvideo_kernel builds that carry the sm_100a
# CUDA block-sparse forward (upstream PR #1719). The module itself imports
# fine without the compiled symbols (`_HAS_VSA_SM100A` is then False and
# `is_supported` says no), so this only guards *module* availability.
from fastvideo_kernel import block_sparse_attn_sm100a as _sm100a
except ImportError:
_sm100a = None
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder, layer_idx_from_prefix)
@@ -71,115 +43,51 @@ from fastvideo.attention.backends.video_sparse_attn import (compute_topk, constr
get_non_pad_index, get_tile_partition_indices,
scatter_into_tile_buf)
from fastvideo.attention.backends.video_sparse_attn_h3_probe import probe_enabled, record_probe
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Opt-in switch for the sm_100a CUDA forward on the tile-64 no-grad path.
VSA_SM100A_ENV = "FASTVIDEO_VSA_SM100A"
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x (default)
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
_TILE_ELEMS = math.prod(VSA_H3_TILE_SIZE)
# Selectable tile geometries, keyed by element count (= the build-time
# ``tile_size``). 64 runs the native 64-token Triton block-sparse kernels for
# forward AND backward — the block map is already at kernel granularity, so no
# 256->64 mask expansion is involved and FASTVIDEO_VSA_CUTEDSL does not apply.
VSA_H3_TILE_SHAPES: dict[int, tuple[int, int, int]] = {
_TILE_ELEMS: VSA_H3_TILE_SIZE,
64: (4, 4, 4),
}
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
def token_tile_and_valid(variable_block_sizes: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""Per padded-token tile id and pad-validity mask.
The single encoding of the padding contract, shared by the probe and the
test oracle so they cannot drift from the backend's tile geometry.
``tile_elems`` must match the metadata the sizes came from
(``MiniMaxH3VSAMetadata.tile_elems``).
"""
device = variable_block_sizes.device
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(_TILE_ELEMS)
token_valid = (torch.arange(_TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
return token_tile, token_valid
def _validate_h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
variable_block_sizes: torch.Tensor,
untile_combined_index: torch.Tensor,
tile_elems: int = _TILE_ELEMS,
) -> None:
"""Fail synchronously on out-of-bounds tile geometry.
Invariants the block-sparse kernel trusts without checking:
every tile's valid size is in (0, tile_elems]; the sizes sum to the
packed sequence length; and ``untile_combined_index`` maps each packed
row to exactly one non-pad slot of the padded tile buffer. A violation
would surface only as an async device fault at some later kernel or
collective (e.g. an FSDP all-gather), which is unattributable — so raise
here, once per cached geometry, with the numbers in hand.
"""
total = sum(prefix_segments) + math.prod(dit_seq_shape)
n_pad = variable_block_sizes.numel() * tile_elems
sizes_min = int(variable_block_sizes.min())
sizes_max = int(variable_block_sizes.max())
sizes_sum = int(variable_block_sizes.sum())
if sizes_min < 1 or sizes_max > tile_elems or sizes_sum != total:
raise ValueError(f"VSA-H3 tile sizes out of bounds for prefix={prefix_segments}, video={dit_seq_shape}, "
f"tile_elems={tile_elems}: min={sizes_min}, max={sizes_max}, sum={sizes_sum}, "
f"expected sum={total}.")
if untile_combined_index.numel() != total:
raise ValueError(f"VSA-H3 untile index has {untile_combined_index.numel()} entries for a packed "
f"sequence of {total} rows (prefix={prefix_segments}, video={dit_seq_shape}).")
idx_min = int(untile_combined_index.min())
idx_max = int(untile_combined_index.max())
if idx_min < 0 or idx_max >= n_pad:
# Range first: the pad-slot gather below would itself index out of
# bounds (the very async fault this guard exists to preempt).
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: range "
f"[{idx_min}, {idx_max}] vs padded length {n_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
in_tile_offset = untile_combined_index % tile_elems
maps_into_pad = bool((in_tile_offset >= variable_block_sizes[untile_combined_index // tile_elems]).any())
if maps_into_pad or int(torch.unique(untile_combined_index).numel()) != total:
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: "
f"pad-slot hit={maps_into_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
@functools.lru_cache(maxsize=10)
def _h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
device: torch.device,
tile_shape: tuple[int, int, int] = VSA_H3_TILE_SIZE,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
"""Tile the packed sequence: segment-pure prefix chunks, then video tiles.
Returns (tile_partition_indices, variable_block_sizes,
untile_combined_index, num_prefix_tiles, num_video_tiles).
"""
tile_elems = math.prod(tile_shape)
prefix_len = sum(prefix_segments)
prefix_sizes: list[int] = []
for segment in prefix_segments:
full, rem = divmod(segment, tile_elems)
prefix_sizes.extend([tile_elems] * full)
full, rem = divmod(segment, _TILE_ELEMS)
prefix_sizes.extend([_TILE_ELEMS] * full)
if rem:
prefix_sizes.append(rem)
num_prefix_tiles = len(prefix_sizes)
ts_t, ts_h, ts_w = tile_shape
ts_t, ts_h, ts_w = VSA_H3_TILE_SIZE
t, h, w = dit_seq_shape
num_tiles = (math.ceil(t / ts_t), math.ceil(h / ts_h), math.ceil(w / ts_w))
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, VSA_H3_TILE_SIZE)
num_video_tiles = int(video_sizes.numel())
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
tile_partition_indices = torch.cat([
torch.arange(prefix_len, device=device, dtype=torch.long),
video_indices,
@@ -192,11 +100,9 @@ def _h3_tile_geometry(
# get_non_pad_index is lru-cached on tensor identity; variable_block_sizes
# is itself cached by this function, so the identity stays stable.
non_pad_index = get_non_pad_index(variable_block_sizes, tile_elems)
non_pad_index = get_non_pad_index(variable_block_sizes, _TILE_ELEMS)
untile_combined_index = non_pad_index[torch.argsort(tile_partition_indices)]
# One-time (lru-cached) synchronous bounds check; see _validate_h3_tile_geometry.
_validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
return (tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles)
@@ -233,9 +139,6 @@ class MiniMaxH3VSAMetadata(AttentionMetadata):
exempt: bool
variable_block_sizes: torch.Tensor
untile_combined_index: torch.Tensor
# tokens per tile (256 or 64); selects the tile geometry AND the kernel
# route in forward() (256 -> VSA-256 CuTe/Triton, 64 -> native Triton)
tile_elems: int = _TILE_ELEMS
# layers forced dense regardless of sparsity (probe-guided opt-outs)
dense_layers: tuple[int, ...] = ()
# Single-slot holder for the padded tile buffer, owned by the BUILDER so
@@ -255,28 +158,24 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
prefix_segments: tuple[int, ...],
device: torch.device,
exempt: bool = True,
dense_layers: tuple[int, ...] = (),
tile_size: int = _TILE_ELEMS,
**kwargs: dict[str, Any],
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
prefix_segments: tuple[int, ...],
device: torch.device,
exempt: bool = True,
dense_layers: tuple[int, ...] = (),
**kwargs: dict[str, Any],
) -> MiniMaxH3VSAMetadata:
tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
if tile_shape is None:
raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
raw_latent_shape[2] // patch_size[2])
prefix_segments = tuple(int(s) for s in prefix_segments if s > 0)
total_seq_length = sum(prefix_segments) + math.prod(dit_seq_shape)
(_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles,
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
return MiniMaxH3VSAMetadata(
current_timestep=current_timestep,
@@ -287,14 +186,13 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
exempt=exempt,
variable_block_sizes=variable_block_sizes,
untile_combined_index=untile_combined_index,
tile_elems=int(tile_size),
dense_layers=tuple(int(layer) for layer in dense_layers),
tile_buf_holder=self._tile_buf_holder,
)
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems: int = _TILE_ELEMS) -> torch.Tensor:
"""fp32 mean over each tile_elems-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
written), so a plain sum with fp32 accumulation needs no validity mask
@@ -302,8 +200,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems:
the masked mean exactly.
"""
batch, seq_len, heads, dim = x.shape
n_tiles = seq_len // tile_elems
pooled = x.view(batch, n_tiles, tile_elems, heads, dim).sum(dim=2, dtype=torch.float32)
n_tiles = seq_len // _TILE_ELEMS
pooled = x.view(batch, n_tiles, _TILE_ELEMS, heads, dim).sum(dim=2, dtype=torch.float32)
pooled = pooled / variable_block_sizes.view(1, -1, 1, 1)
return pooled.permute(0, 2, 1, 3)
@@ -334,24 +232,6 @@ def _build_block_mask(
return mask
def _sm100a_unavailable_reason(sm100a_mod: Any, query_bhsd: torch.Tensor, variable_block_sizes: torch.Tensor,
grad_mode: bool) -> str | None:
"""Why the opt-in sm_100a forward route cannot run here, or None if it can.
Pure decision logic, split out so the routing is unit-testable without a
GPU or the compiled extension (tests substitute ``sm100a_mod``). Order
matters only for the message: the cheapest, most actionable reason first.
"""
if sm100a_mod is None:
return "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
if grad_mode:
return "inputs require grad and the sm_100a kernel is forward-only; grad paths keep Triton"
if not sm100a_mod.is_supported(query_bhsd, variable_block_sizes):
return ("block_sparse_attn_sm100a.is_supported returned False (needs an sm_100 device, a built "
"extension, bf16, head_dim 128, an even tile count, and integer tile sizes)")
return None
class MiniMaxH3VSAImpl(AttentionImpl):
def __init__(
@@ -379,7 +259,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
"routed to the VSA-H3 backend; exclude it from the supported backends.")
n_tiles = attn_metadata.variable_block_sizes.numel()
target_shape = (x.shape[0], n_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])
target_shape = (x.shape[0], n_tiles * _TILE_ELEMS, x.shape[-2], x.shape[-1])
# single scatter: untile_combined_index maps original row i to its
# padded slot, so this is exactly the inverse of postprocess_output
@@ -401,11 +281,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
gate_compress: torch.Tensor | None,
attn_metadata: MiniMaxH3VSAMetadata,
) -> torch.Tensor:
tile_elems = attn_metadata.tile_elems
if tile_elems == 64:
if block_sparse_attn_64_bhsd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn is not installed")
elif block_sparse_attn_256_bshd is None:
if block_sparse_attn_256_bshd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
@@ -415,8 +291,8 @@ class MiniMaxH3VSAImpl(AttentionImpl):
scores = None
if layer_sparsity > 0.0 or gate_compress is not None or probe_dir is not None:
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (query.shape[-1]**0.5)
if probe_dir is not None:
record_probe(probe_dir, self.layer_idx, query, key, scores, attn_metadata)
@@ -433,75 +309,18 @@ class MiniMaxH3VSAImpl(AttentionImpl):
attn_metadata.exempt,
)
if tile_elems == 64:
# Native 64-token path: the block map is already at the kernels'
# granularity. Both 64-token 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()
# Opt-in sm_100a CUDA forward (upstream PR #1719 + per-q-tile
# q2k_num fix). Forward-only: grad-tracking calls stay on Triton
# so autograd keeps the Triton fwd+bwd pairing untouched. The
# kernel does return an LSE in Triton's M format, so a future
# fwd/bwd pairing is possible, but it is not built here.
use_sm100a = False
if os.environ.get(VSA_SM100A_ENV, "0") == "1":
grad_mode = torch.is_grad_enabled() and (query.requires_grad or key.requires_grad
or value.requires_grad)
reason = _sm100a_unavailable_reason(_sm100a, q_bhsd, attn_metadata.variable_block_sizes, grad_mode)
if reason is None and map_to_index is None:
reason = "fastvideo_kernel.triton_kernels.index (map_to_index) is not importable"
if reason is None:
use_sm100a = True
elif not torch.compiler.is_compiling():
logger.warning_once(f"{VSA_SM100A_ENV}=1 but falling back to the Triton-64 kernels: {reason}")
if use_sm100a:
# The sm_100a entry is index-native; compact the bool map the
# same way the Triton bool entry does internally. Per-row
# counts are NON-uniform here (prefix query tiles are dense,
# video tiles run prefix+top-k) -- legal for the fixed kernel,
# silently wrong on the pre-fix upstream one.
q2k_idx, q2k_num = map_to_index(mask)
out_bhsd, _ = _sm100a.block_sparse_attn_sm100a(
q_bhsd,
k_bhsd,
v_bhsd,
q2k_idx,
q2k_num,
attn_metadata.variable_block_sizes.to(torch.int32),
need_lse=False,
)
else:
out_bhsd, _ = block_sparse_attn_64_bhsd(
q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
)
out = out_bhsd.transpose(1, 2).contiguous()
else:
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
if gate_compress is not None:
# Wan-style compression branch: dense attention over pooled tiles,
# broadcast to each tile's rows, scaled by the learned gate
# (zero-initialized for H3 => branch contributes nothing until
# finetuned; the model layer skips it entirely for all-zero gates).
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
batch, seq_len, heads, dim = out.shape
batch, _, heads, dim = out.shape
n_tiles = attn_metadata.variable_block_sizes.numel()
# Out-of-place: on the CuTe backend ``out`` is the tensor FA4's
# autograd node saved for its backward, so an in-place add here
# bumps its version counter and backward dies with "one of the
# variables needed for gradient computation has been modified".
out_tiled = out.view(batch, n_tiles, tile_elems, heads, dim)
gate_tiled = gate_compress.view(batch, n_tiles, tile_elems, heads, dim)
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
out.view(batch, n_tiles, _TILE_ELEMS, heads,
dim).addcmul_(out_c.unsqueeze(2), gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim))
return out
@@ -60,7 +60,7 @@ def record_probe(
gen = torch.Generator(device="cpu").manual_seed(step * 1000 + layer)
# sample among video rows in the PADDED/tiled domain that are non-pad
from fastvideo.attention.backends.video_sparse_attn_h3 import token_tile_and_valid
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes)
video_rows = torch.nonzero((token_tile >= P) & token_valid, as_tuple=False).flatten()
idx = video_rows[torch.randint(0, video_rows.numel(), (_TRUE_ROWS, ), generator=gen).to(query.device)]
-1
View File
@@ -1 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
-375
View File
@@ -1,375 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Benchmark generate-fewer-frames plus MLX RIFE interpolation.
This script keeps the video diffusion path delegated to
``fastvideo.benchmarks.mlx_fastwan_bench._generate_cell``. It only orchestrates
two frame-count cells and the postprocess interpolation step.
"""
from __future__ import annotations
import argparse
import json
import shutil
import time
from pathlib import Path
from types import SimpleNamespace
import imageio.v2 as imageio
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import encode_prompt, make_rotary_embeddings
from fastvideo.benchmarks.mlx_fastwan_bench import _generate_cell, _ms_ssim
from fastvideo.mlx_runtime.rife_interp import RIFEBackendError, interpolate, load_model
from fastvideo.mlx_runtime.memory import add_memory_limit_args, apply_memory_limits
FOX_PROMPT = "A fox runs through a misty pine forest, leaves kicking up behind it."
DEFAULT_MODEL_ROOT = Path("/Users/aryank/models/qad_int8_v2")
def _parse_timesteps(raw: str) -> list[int]:
timesteps = [int(part.strip()) for part in raw.split(",") if part.strip()]
if not timesteps:
raise SystemExit("No DMD timesteps parsed from --dmd-denoising-steps")
return timesteps
def _read_video(path: Path) -> list[np.ndarray]:
if not path.is_file():
raise FileNotFoundError(f"Video does not exist: {path}")
frames = [np.asarray(frame[:, :, :3], dtype=np.uint8) for frame in imageio.mimread(path)]
if not frames:
raise RuntimeError(f"No frames decoded from {path}")
return frames
def _write_video(path: Path, frames: list[np.ndarray], fps: int) -> None:
if not frames:
raise ValueError("Cannot write an empty frame list")
path.parent.mkdir(parents=True, exist_ok=True)
with imageio.get_writer(str(path), fps=fps, macro_block_size=1, codec="libx264", quality=8) as writer:
for frame in frames:
writer.append_data(np.asarray(frame, dtype=np.uint8))
def _copy_video(src: Path, dst: Path) -> Path:
dst.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(src, dst)
return dst
def _make_generation_inputs(args, num_frames: int):
import mlx.core as mx
import torch
config_path = args.model_root / "transformer" / "config.json"
checkpoint_path = args.model_root / "transformer" / "diffusion_pytorch_model.safetensors"
if not config_path.is_file():
raise SystemExit(f"Missing DiT config: {config_path}")
if not checkpoint_path.is_file():
raise SystemExit(f"Missing DiT checkpoint: {checkpoint_path}")
config = json.loads(config_path.read_text())
latent_frames = (num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
freqs_cis = make_rotary_embeddings(
config,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
)
generator = torch.Generator(device="cpu").manual_seed(args.seed)
latents_seed = torch.randn(
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
generator=generator,
dtype=torch.float32,
).numpy()
timesteps = _parse_timesteps(args.dmd_denoising_steps)
renoise_by_step = [
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
for _ in range(max(0,
len(timesteps) - 1))
]
prompt_embeds = encode_prompt(
model_root=args.model_root,
prompt=args.prompt,
max_sequence_length=args.max_sequence_length,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
)
return {
"checkpoint_path": checkpoint_path,
"config_path": config_path,
"encoder_hidden_states": mx.array(prompt_embeds.numpy()),
"freqs_cis": freqs_cis,
"timesteps": timesteps,
"latents_seed": latents_seed,
"renoise_by_step": renoise_by_step,
"latent_frames": latent_frames,
}
def _generate(args, num_frames: int):
cell_args = SimpleNamespace(
model_root=args.model_root,
height=args.height,
width=args.width,
num_frames=num_frames,
fps=args.fps,
flow_shift=args.flow_shift,
torch_device=args.torch_device,
torch_dtype=args.torch_dtype,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
mlx_checkpoint_cache=args.mlx_checkpoint_cache,
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
mlx_disable_cache=args.mlx_disable_cache,
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
benchmark_preset="metalfx-rife",
current_prompt_id=f"fox-{num_frames}f",
current_prompt=args.prompt,
output_dir=args.output_dir,
)
inputs = _make_generation_inputs(args, num_frames)
print(f"=== generate: {num_frames} frames ({inputs['latent_frames']} latent frames) ===", flush=True)
return _generate_cell(
args=cell_args,
mode=args.mode,
decoder=args.decoder,
checkpoint_path=inputs["checkpoint_path"],
config_path=inputs["config_path"],
encoder_hidden_states=inputs["encoder_hidden_states"],
freqs_cis=inputs["freqs_cis"],
timesteps=inputs["timesteps"],
latents_seed=inputs["latents_seed"],
renoise_by_step=inputs["renoise_by_step"],
)
def _relative_to_output(path: Path, output_dir: Path) -> str:
try:
return str(path.relative_to(output_dir))
except ValueError:
return str(path)
def _write_report(args, result: dict) -> Path:
report_path = args.output_dir / "metalfx_rife_report.md"
rows = result["speed_rows"]
table_lines = [
"| path | denoise_s | decode_s | gen_total_s | rife_s | net_s | speedup_vs_81 |",
"| --- | ---: | ---: | ---: | ---: | ---: | ---: |",
]
for row in rows:
table_lines.append(
"| {path} | {denoise_s:.3f} | {decode_s:.3f} | {gen_total_s:.3f} | {rife_s:.3f} | {net_s:.3f} | {speedup_vs_81:.3f}x |"
.format(**row))
text = f"""# MetalFX/RIFE Generate-Fewer-Frames Benchmark Run
Prompt: {args.prompt}
Resolution: {args.height}x{args.width}, fps={args.fps}, mode={args.mode}, decoder={args.decoder}
| metric | value |
| --- | ---: |
| reconstruction_ms_ssim | {result['reconstruction_ms_ssim']:.6f} |
| reference_frames | {result['reference_frames']} |
| reduced_frames | {result['reduced_frames']} |
| interpolated_frames | {result['interpolated_frames']} |
{chr(10).join(table_lines)}
Videos:
- reference: `{result['videos']['reference']}`
- reference drop-41 RIFE reconstruction: `{result['videos']['drop41_rife81']}`
- generated 41 RIFE to 81: `{result['videos']['generated41_rife81']}`
- generated 41 direct: `{result['videos']['generated41']}`
Raw metrics are in `{_relative_to_output(args.output_dir / 'metrics.json', args.output_dir)}`.
"""
report_path.write_text(text)
return report_path
def main() -> None:
parser = argparse.ArgumentParser(
description="Evaluate 41-frame generation plus MLX RIFE interpolation vs 81-frame generation.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--prompt", default=FOX_PROMPT)
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=832)
parser.add_argument("--reference-frames", type=int, default=81)
parser.add_argument("--reduced-frames", type=int, default=41)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--mode", default="int8", choices=("fp16", "bf16", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"))
parser.add_argument("--decoder", default="taehv", choices=("taehv", "wan-vae"))
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
parser.add_argument("--flow-shift", type=float, default=8.0)
parser.add_argument("--max-sequence-length", type=int, default=512)
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument("--output-dir", type=Path, default=Path("bench/metalfx_rife"))
parser.add_argument("--mlx-checkpoint-cache", type=Path, default=None)
parser.add_argument("--compile", action="store_true", help="Enable FASTVIDEO_MLX_COMPILE=1 for DiT denoise.")
parser.add_argument("--rife-scale", type=float, default=1.0)
parser.add_argument("--taehv-source-path", type=Path, default=None)
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
parser.add_argument("--taehv-parallel", action="store_true")
add_memory_limit_args(parser)
args = parser.parse_args()
if args.reference_frames != (args.reduced_frames - 1) * 2 + 1:
raise SystemExit(
"--reference-frames must equal (--reduced-frames - 1) * 2 + 1 for the default every-other-frame test")
if args.compile:
import os
os.environ["FASTVIDEO_MLX_COMPILE"] = "1"
args.model_root = args.model_root.expanduser().resolve()
args.output_dir = args.output_dir.expanduser().resolve()
args.output_dir.mkdir(parents=True, exist_ok=True)
if args.mlx_checkpoint_cache is None:
args.mlx_checkpoint_cache = args.output_dir / "mlx_checkpoint_cache"
else:
args.mlx_checkpoint_cache = args.mlx_checkpoint_cache.expanduser().resolve()
import mlx.core as mx
import torch
runtime_limits = apply_memory_limits(
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
mlx_disable_cache=args.mlx_disable_cache,
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
mx_module=mx,
).as_metrics()
mx.random.seed(args.seed)
torch.manual_seed(args.seed)
reference_cell = _generate(args, args.reference_frames)
reduced_cell = _generate(args, args.reduced_frames)
reference_video = _copy_video(reference_cell.video_path, args.output_dir / "fox_reference_81.mp4")
generated41_video = _copy_video(reduced_cell.video_path, args.output_dir / "fox_generated_41.mp4")
reference_frames = _read_video(reference_video)
dropped_reference_frames = reference_frames[::2]
if len(dropped_reference_frames) != args.reduced_frames:
raise RuntimeError(f"Expected {args.reduced_frames} dropped frames, got {len(dropped_reference_frames)}")
try:
rife_model = load_model("4.25")
except RIFEBackendError:
raise
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
raise RIFEBackendError(f"Unexpected RIFE load failure: {exc}") from exc
print("=== RIFE reconstruction: drop every other reference frame, 41 -> 81 ===", flush=True)
start = time.perf_counter()
reconstructed_frames = interpolate(dropped_reference_frames, factor=2, model=rife_model, scale=args.rife_scale)
recon_rife_s = time.perf_counter() - start
reconstructed_video = args.output_dir / "fox_reference_drop41_rife81.mp4"
_write_video(reconstructed_video, reconstructed_frames, args.fps)
reconstruction_ms_ssim = _ms_ssim(reference_video, reconstructed_video, required=True)
if reconstruction_ms_ssim is None:
raise RuntimeError("MS-SSIM returned None for reconstruction comparison")
print("=== RIFE actual reduced path: generated 41 -> 81 ===", flush=True)
generated41_frames = _read_video(generated41_video)
start = time.perf_counter()
generated41_rife_frames = interpolate(generated41_frames, factor=2, model=rife_model, scale=args.rife_scale)
generated41_rife_s = time.perf_counter() - start
generated41_rife_video = args.output_dir / "fox_generated41_rife81.mp4"
_write_video(generated41_rife_video, generated41_rife_frames, args.fps)
direct_denoise_s = float(reference_cell.metrics["denoise_s"])
reduced_denoise_s = float(reduced_cell.metrics["denoise_s"])
direct_gen_total_s = float(reference_cell.metrics["denoise_s"]) + float(reference_cell.metrics["decode_s"])
reduced_gen_total_s = float(reduced_cell.metrics["denoise_s"]) + float(reduced_cell.metrics["decode_s"])
net_denoise_rife_s = reduced_denoise_s + generated41_rife_s
net_gen_rife_s = reduced_gen_total_s + generated41_rife_s
result = {
"prompt":
args.prompt,
"model_root":
str(args.model_root),
"rife_impl":
"rife-mlx vendored at fastvideo/third_party/rife_mlx, weights mlx-community/RIFE-4.25",
"runtime_limits":
runtime_limits,
"reference_frames":
len(reference_frames),
"reduced_frames":
len(generated41_frames),
"interpolated_frames":
len(generated41_rife_frames),
"reconstruction_ms_ssim":
reconstruction_ms_ssim,
"reconstruction_rife_s":
recon_rife_s,
"generated41_rife_s":
generated41_rife_s,
"reference_metrics":
reference_cell.metrics,
"reduced_metrics":
reduced_cell.metrics,
"speed_rows": [
{
"path": "generate_81",
"denoise_s": direct_denoise_s,
"decode_s": float(reference_cell.metrics["decode_s"]),
"gen_total_s": direct_gen_total_s,
"rife_s": 0.0,
"net_s": direct_gen_total_s,
"speedup_vs_81": 1.0,
},
{
"path": "generate_41_plus_rife81_denoise_only",
"denoise_s": reduced_denoise_s,
"decode_s": 0.0,
"gen_total_s": reduced_denoise_s,
"rife_s": generated41_rife_s,
"net_s": net_denoise_rife_s,
"speedup_vs_81": direct_denoise_s / net_denoise_rife_s,
},
{
"path": "generate_41_plus_rife81_decode_included",
"denoise_s": reduced_denoise_s,
"decode_s": float(reduced_cell.metrics["decode_s"]),
"gen_total_s": reduced_gen_total_s,
"rife_s": generated41_rife_s,
"net_s": net_gen_rife_s,
"speedup_vs_81": direct_gen_total_s / net_gen_rife_s,
},
],
"videos": {
"reference": _relative_to_output(reference_video, args.output_dir),
"drop41_rife81": _relative_to_output(reconstructed_video, args.output_dir),
"generated41": _relative_to_output(generated41_video, args.output_dir),
"generated41_rife81": _relative_to_output(generated41_rife_video, args.output_dir),
},
}
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(result, indent=2))
report_path = _write_report(args, result)
print(json.dumps(result["speed_rows"], indent=2))
print(f"reconstruction_ms_ssim={reconstruction_ms_ssim:.6f}")
print(f"wrote {metrics_path}")
print(f"wrote {report_path}")
if __name__ == "__main__":
main()
-826
View File
@@ -1,826 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Prove-out benchmark for the MLX FastWan runtime (Apple Silicon).
Sweeps ``{dtype/quant} x {decoder}``, generates a clip per cell, and records the
latency breakdown, peak unified memory, and MS-SSIM (optionally LPIPS) against a
reference video. It emits a JSON blob and a markdown table -- the artifact that
turns "int8 + TAEHV looks good" into defensible numbers, and (via
``--assert-min-ssim``) a regression gate for the ``mx.compile`` work.
Design notes:
- Generation reuses the hybrid POC helpers in
``examples/inference/basic/mlx_wan_prompt_to_video.py`` (torch-MPS UMT5 encode
and Wan-VAE/TAEHV decode) plus the on-device MLX DMD sampler
(``fastvideo/mlx_runtime/sampling.py``); the denoise loop never leaves the
device.
- Quality reuses the tested MS-SSIM primitive
``fastvideo/tests/utils.py::compute_video_ssim_torchvision``.
- Reference: by default each cell is scored against the highest-fidelity cell
in the sweep (``fp16`` + ``wan-vae``), which needs no CUDA box and answers
"how much does int8/int4/TAEHV degrade vs the best local config". Pass
``--reference PATH`` to score against an external clip instead (e.g. the
torch-MPS or CUDA FastVideo output of the same model) for a "vs. the original
model" column.
Run on an Apple Silicon Mac (needs ``mlx`` + a torch build with MPS):
python fastvideo/benchmarks/mlx_fastwan_bench.py \
--modes fp16,bf16,int8,int4 --decoders taehv,wan-vae
"""
from __future__ import annotations
import argparse
import html
import json
import os
import time
from dataclasses import dataclass, field
from pathlib import Path
import numpy as np
from examples.inference.basic.mlx_wan_prompt_to_video import (
DEFAULT_MODEL_ROOT,
decode_latents_to_video,
encode_prompt,
make_rotary_embeddings,
)
from fastvideo.mlx_runtime.memory import add_memory_limit_args, apply_memory_limits, cleanup_mlx
# The highest-fidelity cell; used as the default SSIM reference when no external
# reference video is supplied.
REFERENCE_MODE = "fp16"
REFERENCE_DECODER = "wan-vae"
ALLOWED_MODES = ("fp16", "bf16", "int8", "int4", "mxfp8", "mxfp4", "nvfp4")
ALLOWED_DECODERS = ("taehv", "wan-vae")
@dataclass(frozen=True)
class PromptCase:
id: str
prompt: str
@dataclass(frozen=True)
class BenchmarkPreset:
height: int
width: int
num_frames: int
modes: str
decoders: str
mlx_memory_limit_gib: float | None = None
mlx_disable_cache: bool = False
torch_mps_high_watermark_ratio: float | None = None
torch_mps_low_watermark_ratio: float | None = None
PROMPT_SETS = {
"motion7": (
PromptCase("beach-sunset", "A slow cinematic sunset over ocean waves at a quiet beach."),
PromptCase("fox-forest", "A fox runs through a misty pine forest, leaves kicking up behind it."),
PromptCase("raccoon-sunflowers", "A raccoon walks through a sunflower field as petals move in the wind."),
PromptCase("surfing-cat", "A cat wearing sunglasses surfs across a bright blue ocean wave."),
PromptCase("burning-clock", "A vintage table clock burns on a wooden desk, flames flickering realistically."),
PromptCase("forest-walk", "Video game style footage of a man walking through a dense forest path."),
PromptCase("sea-dock-yachts", "Several yachts are parked at a sea dock while water ripples around them."),
),
}
BENCHMARK_PRESETS = {
"default":
BenchmarkPreset(
height=480,
width=832,
num_frames=81,
modes="fp16,bf16,int8,int4",
decoders="taehv,wan-vae",
),
"mac-16gb":
BenchmarkPreset(
height=448,
width=832,
num_frames=61,
modes="int8",
decoders="taehv",
mlx_memory_limit_gib=16.0,
mlx_disable_cache=True,
torch_mps_high_watermark_ratio=0.57,
torch_mps_low_watermark_ratio=0.0,
),
"mac-32gb":
BenchmarkPreset(
height=480,
width=832,
num_frames=81,
modes="int8,fp16",
decoders="taehv",
),
"mac-64gb":
BenchmarkPreset(
height=480,
width=832,
num_frames=81,
modes="int8,fp16",
decoders="taehv,wan-vae",
),
}
@dataclass
class Cell:
prompt_id: str
prompt: str
mode: str
decoder: str
video_path: Path
latents: np.ndarray
metrics: dict[str, float | int | str | bool | None] = field(default_factory=dict)
def _mode_to_dtype_quant(mode: str) -> tuple[str, str | None]:
"""Map a sweep mode to (MLX compute dtype, quantization spec).
Quantized modes keep fp16 activations and quantize only the DiT linear
weights (matching ``mlx_dit_from_diffusers_safetensors``).
"""
if mode == "bf16":
return "bf16", None
if mode == "fp16":
return "fp16", None
# int8/int4/mxfp*/nvfp4 -> fp16 activations + quantized weights.
return "fp16", mode
def _mx_dtype(mx, base: str):
return {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[base]
def _parse_list(raw: str, allowed: tuple[str, ...], label: str) -> list[str]:
items = [x.strip() for x in raw.split(",") if x.strip()]
unknown = sorted(set(items) - set(allowed))
if unknown:
raise ValueError(f"Unsupported {label}: {unknown} (allowed: {list(allowed)})")
return items
def _safe_slug(value: str, *, fallback: str) -> str:
slug = "".join(ch.lower() if ch.isalnum() else "-" for ch in value.strip())
slug = "-".join(part for part in slug.split("-") if part)
return slug[:64] or fallback
def _load_prompt_cases(prompt: str, prompt_file: Path | None, prompt_set: str = "single") -> list[PromptCase]:
"""Load one prompt, a built-in prompt set, or a text/jsonl prompt file."""
if prompt_file is None:
if prompt_set == "single":
return [PromptCase(id="prompt-001", prompt=prompt)]
if prompt_set not in PROMPT_SETS:
raise ValueError(f"Unsupported prompt set: {prompt_set} (allowed: {sorted(PROMPT_SETS) + ['single']})")
return list(PROMPT_SETS[prompt_set])
cases: list[PromptCase] = []
for line_index, raw_line in enumerate(prompt_file.read_text().splitlines(), start=1):
line = raw_line.strip()
if not line or line.startswith("#"):
continue
prompt_id = f"prompt-{len(cases) + 1:03d}"
prompt_text = line
if prompt_file.suffix.lower() == ".jsonl":
item = json.loads(line)
prompt_text = str(item.get("prompt") or item.get("text") or item.get("caption") or "").strip()
if not prompt_text:
raise ValueError(f"{prompt_file}:{line_index} has no prompt/text/caption field")
prompt_id = str(item.get("id") or item.get("name") or prompt_id)
cases.append(PromptCase(id=_safe_slug(prompt_id, fallback=f"prompt-{len(cases) + 1:03d}"), prompt=prompt_text))
if not cases:
raise ValueError(f"No prompts found in {prompt_file}")
return cases
def denoise_dmd_on_device(
*,
mx,
dit,
latents,
encoder_hidden_states,
freqs_cis,
timesteps: list[int],
renoise_by_step: list[np.ndarray],
schedule,
dmd_step,
mx_dtype,
) -> tuple[np.ndarray, list[float]]:
"""Run the FastWan DMD loop entirely on the MLX device.
Mirrors the loop in ``mlx_wan_prompt_to_video.py`` (fp32 affine math, MLX RNG
re-noise) so the benchmark measures exactly the shipped path.
Returns the final latents plus per-step wall times. The first step carries
one-time costs (mx.compile tracing, kernel warm-up), so first-vs-steady
step timing is how the benchmark separates cold-start from steady-state
denoise throughput.
All host-side tensors (timesteps, re-noise draws) are uploaded before the
loop starts, so the per-step body performs no bulk host->device transfers
and step timings measure device work rather than staging copies.
"""
timesteps_mx = [mx.array([float(timestep)]).astype(mx.float32) for timestep in timesteps]
renoise_mx = [mx.array(renoise).astype(mx.float32) for renoise in renoise_by_step]
if timesteps_mx or renoise_mx:
mx.eval(*timesteps_mx, *renoise_mx)
step_times: list[float] = []
for step_index, timestep in enumerate(timesteps):
step_start = time.perf_counter()
noise_input_latent = latents
noise_pred = dit(latents.astype(mx_dtype), encoder_hidden_states, timesteps_mx[step_index], freqs_cis)
noise_input_f32 = noise_input_latent.astype(mx.float32)
pred_noise_f32 = noise_pred.astype(mx.float32)
if step_index < len(timesteps) - 1:
next_ts: float | None = float(timesteps[step_index + 1])
renoise = renoise_mx[step_index]
else:
next_ts, renoise = None, None
latents = dmd_step(
latents=noise_input_f32,
noise_input_latent=noise_input_f32,
pred_noise=pred_noise_f32,
schedule=schedule,
timestep=float(timestep),
next_timestep=next_ts,
noise=renoise,
).astype(mx_dtype)
mx.eval(latents)
step_times.append(time.perf_counter() - step_start)
return np.array(latents.astype(mx.float32)), step_times
def _peak_memory_bytes(mx) -> int:
try:
return int(mx.get_peak_memory())
except Exception: # noqa: BLE001 - best-effort telemetry only.
return 0
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
mse = float(np.mean(np.square(diff)))
signal = float(np.mean(np.square(baseline.astype(np.float32))))
return {
"latent_mse_vs_ref_mode": mse,
"latent_snr_db_vs_ref_mode": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
}
def _ms_ssim(reference_video: Path, candidate_video: Path, *, required: bool = False) -> float | None:
"""Mean MS-SSIM between two mp4s, via the repo's tested helper."""
if not reference_video.exists() or not candidate_video.exists():
return None
try:
from fastvideo.tests.utils import compute_video_ssim_torchvision
except ImportError as exc:
message = ("MS-SSIM is unavailable because `pytorch-msssim` is not installed. "
"Install FastVideo with the test extra, e.g. `uv pip install -e '.[mlx,test]'`, "
"or run without an SSIM assertion.")
if required:
raise RuntimeError(message) from exc
print(f"{message} Skipping MS-SSIM.")
return None
try:
ssim_values = compute_video_ssim_torchvision(str(reference_video), str(candidate_video), use_ms_ssim=True)
except ImportError as exc:
message = ("MS-SSIM is unavailable because `pytorch-msssim` is not installed. "
"Install FastVideo with the test extra, e.g. `uv pip install -e '.[mlx,test]'`, "
"or run without an SSIM assertion.")
if required:
raise RuntimeError(message) from exc
print(f"{message} Skipping MS-SSIM.")
return None
return float(ssim_values[0])
def _markdown_table(rows: list[dict]) -> str:
columns = [
("prompt_id", "prompt"),
("mode", "mode"),
("decoder", "decoder"),
("status", "status"),
("denoise_s", "denoise s"),
("decode_s", "decode s"),
("total_s", "total s"),
("peak_gib", "peak GiB"),
("ms_ssim_vs_ref", "MS-SSIM"),
("lpips_vs_ref", "LPIPS"),
]
header = "| " + " | ".join(label for _, label in columns) + " |"
sep = "| " + " | ".join("---" for _ in columns) + " |"
lines = [header, sep]
for row in rows:
cells = []
for key, _ in columns:
value = row.get(key)
if isinstance(value, float):
cells.append(f"{value:.3f}")
elif value is None:
cells.append("-")
else:
cells.append(str(value))
lines.append("| " + " | ".join(cells) + " |")
return "\n".join(lines)
def _format_metric(value) -> str:
if isinstance(value, float):
return f"{value:.3f}"
if value is None:
return "-"
return str(value)
def _html_grid(rows: list[dict]) -> str:
groups: dict[str, list[dict]] = {}
for row in rows:
groups.setdefault(str(row.get("prompt_id", "prompt")), []).append(row)
sections = []
for prompt_id, group_rows in groups.items():
prompt = next((str(row.get("prompt", "")) for row in group_rows if row.get("prompt")), "")
cards = []
for row in group_rows:
title = f"{row.get('mode', '-')} / {row.get('decoder', '-')}"
status = row.get("status", "-")
video_path = row.get("video_path")
if video_path:
media = f'<video src="{html.escape(str(video_path))}" muted loop controls playsinline></video>'
else:
media = f'<div class="missing">No video<br>{html.escape(str(row.get("error", "")))}</div>'
metrics = (f"status={status} · total={_format_metric(row.get('total_s'))}s · "
f"denoise={_format_metric(row.get('denoise_s'))}s · "
f"decode={_format_metric(row.get('decode_s'))}s · "
f"peak={_format_metric(row.get('peak_gib'))}GiB")
cards.append("<article>"
f"<h3>{html.escape(title)}</h3>"
f"{media}"
f"<p>{html.escape(metrics)}</p>"
"</article>")
sections.append("<section>"
f"<h2>{html.escape(prompt_id)}</h2>"
f"<p class=\"prompt\">{html.escape(prompt)}</p>"
f"<div class=\"grid\">{''.join(cards)}</div>"
"</section>")
return """<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>FastVideo MLX benchmark grid</title>
<style>
body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif; margin: 24px; background: #111; color: #eee; }
button { margin-right: 8px; padding: 8px 12px; border-radius: 8px; border: 1px solid #555; background: #222; color: #eee; }
section { margin-top: 28px; }
.prompt { color: #bbb; max-width: 900px; }
.grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(280px, 1fr)); gap: 16px; }
article { background: #1b1b1b; border: 1px solid #333; border-radius: 12px; padding: 12px; }
h1, h2, h3 { margin: 0 0 10px; }
video { width: 100%; border-radius: 8px; background: #000; }
article p { color: #bbb; font-size: 13px; line-height: 1.4; }
.missing { min-height: 160px; display: grid; place-items: center; text-align: center; color: #f5b5b5; background: #2a1515; border-radius: 8px; padding: 12px; }
</style>
</head>
<body>
<h1>FastVideo MLX benchmark grid</h1>
<p>Use the controls below to start/stop every clip together for side-by-side inspection.</p>
<button onclick="for (const v of document.querySelectorAll('video')) { v.currentTime = 0; v.play(); }">Restart + play all</button>
<button onclick="for (const v of document.querySelectorAll('video')) v.pause();">Pause all</button>
""" + "\n".join(sections) + """
</body>
</html>
"""
def _write_html_grid(rows: list[dict], output_dir: Path) -> Path:
html_path = output_dir / "index.html"
html_path.write_text(_html_grid(rows))
return html_path
def _generate_cell(
*,
args,
mode: str,
decoder: str,
checkpoint_path: Path,
config_path: Path,
encoder_hidden_states,
freqs_cis,
timesteps: list[int],
latents_seed: np.ndarray,
renoise_by_step: list[np.ndarray],
) -> Cell:
import mlx.core as mx
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
base_dtype, quantization = _mode_to_dtype_quant(mode)
mx_dtype = _mx_dtype(mx, base_dtype)
mx.clear_cache()
mx.reset_peak_memory()
load_start = time.perf_counter()
load_source = "diffusers"
if args.mlx_checkpoint_cache is not None:
from fastvideo.mlx_runtime.checkpoint import (
load_mlx_dit_checkpoint,
save_mlx_dit_checkpoint,
)
mode_ckpt_dir = args.mlx_checkpoint_cache / mode
if (mode_ckpt_dir / "mlx_dit.json").exists():
dit = load_mlx_dit_checkpoint(mode_ckpt_dir)
load_source = "mlx_checkpoint"
else:
dit = mlx_dit_from_diffusers_safetensors(
checkpoint_path,
config_path,
dtype=base_dtype,
quantization=quantization,
)
save_mlx_dit_checkpoint(dit, mode_ckpt_dir)
load_source = "diffusers_then_saved"
else:
dit = mlx_dit_from_diffusers_safetensors(
checkpoint_path,
config_path,
dtype=base_dtype,
quantization=quantization,
)
load_s = time.perf_counter() - load_start
load_peak = _peak_memory_bytes(mx)
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
latents = mx.array(latents_seed).astype(mx_dtype)
mx.reset_peak_memory()
denoise_start = time.perf_counter()
latents_np, step_times = denoise_dmd_on_device(
mx=mx,
dit=dit,
latents=latents,
encoder_hidden_states=encoder_hidden_states.astype(mx_dtype),
freqs_cis=freqs_cis,
timesteps=timesteps,
renoise_by_step=renoise_by_step,
schedule=schedule,
dmd_step=dmd_step,
mx_dtype=mx_dtype,
)
denoise_s = time.perf_counter() - denoise_start
denoise_peak = _peak_memory_bytes(mx)
del dit, latents
cleanup_mlx(mx)
video_path = (args.output_dir / f"{args.current_prompt_id}" /
f"video_{mode}_{decoder}_{args.height}x{args.width}x{args.num_frames}.mp4")
decode_start = time.perf_counter()
decode_latents_to_video(
model_root=args.model_root,
latents_np=latents_np,
output_path=video_path,
fps=args.fps,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
backend=decoder,
taehv_source_path=args.taehv_source_path,
taehv_checkpoint_path=args.taehv_checkpoint_path,
taehv_parallel=args.taehv_parallel,
)
decode_s = time.perf_counter() - decode_start
metrics: dict[str, float | int | str | bool | None] = {
"prompt_id": args.current_prompt_id,
"prompt": args.current_prompt,
"benchmark_preset": args.benchmark_preset,
"mode": mode,
"decoder": decoder,
"status": "ok",
"video_path": str(video_path.relative_to(args.output_dir)),
"load_s": load_s,
"load_source": load_source,
"denoise_s": denoise_s,
# The first step carries one-time costs (mx.compile tracing, kernel
# warm-up); steady-state throughput is the median of the rest.
"denoise_first_step_s": step_times[0] if step_times else None,
"denoise_steady_step_s": (float(np.median(step_times[1:])) if len(step_times) > 1 else None),
"decode_s": decode_s,
"total_s": load_s + denoise_s + decode_s,
"load_peak_gib": load_peak / (1024**3),
"peak_gib": max(load_peak, denoise_peak) / (1024**3),
"quantization": quantization or "none",
"compute_dtype": base_dtype,
"compile": os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1",
"fast_norm": os.environ.get("FASTVIDEO_MLX_FAST_NORM", "0") == "1",
"mlx_memory_limit_gib": args.mlx_memory_limit_gib,
"mlx_cache_limit_gib": args.mlx_cache_limit_gib,
"mlx_disable_cache": args.mlx_disable_cache,
"mlx_wired_limit_gib": args.mlx_wired_limit_gib,
"torch_mps_high_watermark_ratio": args.torch_mps_high_watermark_ratio,
"torch_mps_low_watermark_ratio": args.torch_mps_low_watermark_ratio,
}
return Cell(
prompt_id=args.current_prompt_id,
prompt=args.current_prompt,
mode=mode,
decoder=decoder,
video_path=video_path,
latents=latents_np,
metrics=metrics,
)
def main() -> None:
preset_parser = argparse.ArgumentParser(add_help=False)
preset_parser.add_argument("--benchmark-preset", choices=tuple(BENCHMARK_PRESETS), default="default")
preset_args, _ = preset_parser.parse_known_args()
preset = BENCHMARK_PRESETS[preset_args.benchmark_preset]
parser = argparse.ArgumentParser(description="MLX FastWan prove-out benchmark (latency + quality).")
parser.add_argument("--benchmark-preset",
choices=tuple(BENCHMARK_PRESETS),
default=preset_args.benchmark_preset,
help="Memory-tier benchmark defaults. Explicit CLI flags override preset values.")
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
parser.add_argument("--prompt", default="A paper boat sails through a shallow stream in a mossy forest.")
parser.add_argument(
"--prompt-file",
type=Path,
default=None,
help=
"Optional prompt set. Plain text uses one prompt per non-empty line; .jsonl accepts prompt/text/caption plus optional id/name.",
)
parser.add_argument(
"--prompt-set",
choices=("single", *PROMPT_SETS.keys()),
default="single",
help="Built-in standard prompt set. Ignored when --prompt-file is supplied.",
)
parser.add_argument("--height", type=int, default=preset.height)
parser.add_argument("--width", type=int, default=preset.width)
parser.add_argument("--num-frames", type=int, default=preset.num_frames)
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
parser.add_argument("--flow-shift", type=float, default=8.0)
parser.add_argument("--max-sequence-length", type=int, default=512)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--fps", type=int, default=16)
parser.add_argument("--modes", default=preset.modes)
parser.add_argument("--decoders", default=preset.decoders)
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_fastwan_bench"))
parser.add_argument("--torch-device", default="auto")
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
parser.add_argument(
"--reference",
type=Path,
default=None,
help="External reference mp4 to score every cell against. Defaults to the fp16+wan-vae cell.",
)
parser.add_argument("--assert-min-ssim",
type=float,
default=None,
help="Fail if any cell's MS-SSIM vs the reference falls below this value.")
parser.add_argument("--compile",
action="store_true",
help="Enable mx.compile on the DiT forward (sets FASTVIDEO_MLX_COMPILE=1).")
parser.add_argument("--lpips", action="store_true", help="Also compute LPIPS (needs the `lpips` package).")
parser.add_argument("--taehv-source-path", type=Path, default=None)
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
parser.add_argument("--taehv-parallel", action="store_true")
parser.add_argument(
"--mlx-checkpoint-cache",
type=Path,
default=None,
help="Directory of per-mode pre-quantized MLX checkpoints. The first cell of a mode "
"converts from Diffusers weights and saves here (load_source=diffusers_then_saved); "
"later cells and later runs reload without requantizing (load_source=mlx_checkpoint), "
"which is also how the checkpoint load-time win is measured.",
)
add_memory_limit_args(
parser,
mlx_memory_limit_gib=preset.mlx_memory_limit_gib,
mlx_disable_cache=preset.mlx_disable_cache,
torch_mps_high_watermark_ratio=preset.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=preset.torch_mps_low_watermark_ratio,
)
args = parser.parse_args()
if args.compile:
os.environ["FASTVIDEO_MLX_COMPILE"] = "1"
import mlx.core as mx
runtime_limits = apply_memory_limits(
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
mlx_disable_cache=args.mlx_disable_cache,
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
mx_module=mx,
).as_metrics()
import torch
mx.random.seed(args.seed)
torch.manual_seed(args.seed)
args.output_dir.mkdir(parents=True, exist_ok=True)
modes = _parse_list(args.modes, ALLOWED_MODES, "modes")
decoders = _parse_list(args.decoders, ALLOWED_DECODERS, "decoders")
config_path = args.model_root / "transformer/config.json"
checkpoint_path = args.model_root / "transformer/diffusion_pytorch_model.safetensors"
config = json.loads(config_path.read_text())
latent_frames = (args.num_frames - 1) // 4 + 1
latent_height = args.height // 8
latent_width = args.width // 8
freqs_cis = make_rotary_embeddings(
config,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
)
generator = torch.Generator(device="cpu").manual_seed(args.seed)
latents_seed = torch.randn(
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
generator=generator,
dtype=torch.float32,
).numpy()
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
# Keep DMD stochasticity identical across benchmark cells. Without this,
# FP16/INT8/decoder comparisons can accidentally measure different re-noise
# samples instead of only quantization or decoder differences.
renoise_by_step = [
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
for _ in range(max(0,
len(timesteps) - 1))
]
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
prompt_cases = _load_prompt_cases(args.prompt, args.prompt_file, args.prompt_set)
cells: list[Cell] = []
unsupported_rows: list[dict] = []
for prompt_case in prompt_cases:
args.current_prompt_id = prompt_case.id
args.current_prompt = prompt_case.prompt
prompt_embeds = encode_prompt(
model_root=args.model_root,
prompt=prompt_case.prompt,
max_sequence_length=args.max_sequence_length,
device_arg=args.torch_device,
dtype_arg=args.torch_dtype,
)
encoder_hidden_states = mx.array(prompt_embeds.numpy())
for mode in modes:
for decoder in decoders:
print(f"=== cell: prompt={prompt_case.id} mode={mode} decoder={decoder} ===")
try:
cells.append(
_generate_cell(
args=args,
mode=mode,
decoder=decoder,
checkpoint_path=checkpoint_path,
config_path=config_path,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis,
timesteps=timesteps,
latents_seed=latents_seed,
renoise_by_step=renoise_by_step,
))
except UnsupportedMLXQuantizationError as exc:
# Record the cell as unsupported and keep sweeping: a partial
# report on this MLX build beats crashing the whole run.
print(f"skipping cell (unsupported by this MLX build): {exc}")
unsupported_rows.append({
"prompt_id": prompt_case.id,
"prompt": prompt_case.prompt,
"mode": mode,
"decoder": decoder,
"status": "unsupported_by_mlx",
"error": str(exc),
})
if not cells:
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(unsupported_rows, indent=2))
raise SystemExit(f"No benchmark cell could run: every requested mode is unsupported by this MLX build. "
f"Wrote {metrics_path}.")
# Resolve one internal reference per prompt. A single external reference, if
# supplied, is used for every prompt and only video metrics are computed.
reference_by_prompt: dict[str, tuple[Path, np.ndarray | None]] = {}
if args.reference is not None:
for prompt_case in prompt_cases:
reference_by_prompt[prompt_case.id] = (args.reference, None)
else:
for prompt_case in prompt_cases:
prompt_cells = [c for c in cells if c.prompt_id == prompt_case.id]
if not prompt_cells:
continue
ref_cell = next(
(c for c in prompt_cells if c.mode == REFERENCE_MODE and c.decoder == REFERENCE_DECODER),
prompt_cells[0],
)
reference_by_prompt[prompt_case.id] = (ref_cell.video_path, ref_cell.latents)
print(f"Using internal reference cell for {prompt_case.id}: "
f"mode={ref_cell.mode} decoder={ref_cell.decoder}")
lpips_fn = _load_lpips() if args.lpips else None
rows: list[dict] = []
failures: list[str] = []
for cell in cells:
reference_video, reference_latents = reference_by_prompt[cell.prompt_id]
ms_ssim = _ms_ssim(Path(reference_video), cell.video_path, required=args.assert_min_ssim is not None)
cell.metrics["ms_ssim_vs_ref"] = ms_ssim
cell.metrics.update(runtime_limits)
if reference_latents is not None:
cell.metrics.update(_latent_delta_metrics(cell.latents, reference_latents))
cell.metrics["lpips_vs_ref"] = (_lpips_between(lpips_fn, Path(reference_video), cell.video_path)
if lpips_fn else None)
if args.assert_min_ssim is not None and ms_ssim is not None and ms_ssim < args.assert_min_ssim:
failures.append(
f"{cell.prompt_id}/{cell.mode}/{cell.decoder}: MS-SSIM {ms_ssim:.4f} < {args.assert_min_ssim}")
rows.append(dict(cell.metrics))
print(json.dumps(cell.metrics, indent=2))
rows.extend(unsupported_rows)
metrics_path = args.output_dir / "metrics.json"
metrics_path.write_text(json.dumps(rows, indent=2))
table_path = args.output_dir / "metrics.md"
table = _markdown_table(rows)
table_path.write_text(table + "\n")
html_path = _write_html_grid(rows, args.output_dir)
print("\n" + table)
print(f"\nWrote {metrics_path}, {table_path}, and {html_path}")
if failures:
raise SystemExit("SSIM regression gate failed:\n " + "\n ".join(failures))
def _load_lpips() -> object | None:
"""Return an LPIPS model, or ``None`` if the optional dep is unavailable."""
try:
import lpips # noqa: PLC0415 - optional dependency.
except ImportError:
print("LPIPS requested but the `lpips` package is not installed; skipping (install `.[eval]`).")
return None
return lpips.LPIPS(net="alex")
def _lpips_between(lpips_fn, reference_video: Path, candidate_video: Path) -> float | None:
if lpips_fn is None or not reference_video.exists() or not candidate_video.exists():
return None
import torch
ref = _read_video_frames(reference_video)
cand = _read_video_frames(candidate_video)
if ref is None or cand is None or ref.shape != cand.shape:
return None
# LPIPS expects NCHW in [-1, 1].
ref_t = torch.from_numpy(ref).permute(0, 3, 1, 2).float() / 127.5 - 1.0
cand_t = torch.from_numpy(cand).permute(0, 3, 1, 2).float() / 127.5 - 1.0
with torch.no_grad():
scores = lpips_fn(ref_t, cand_t)
return float(scores.mean().item())
def _read_video_frames(path: Path) -> np.ndarray | None:
try:
import cv2
except ImportError:
return None
cap = cv2.VideoCapture(str(path))
frames = []
try:
while True:
ok, frame_bgr = cap.read()
if not ok:
break
frames.append(cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB))
finally:
cap.release()
if not frames:
return None
return np.stack(frames, axis=0)
if __name__ == "__main__":
main()
+7
View File
@@ -55,6 +55,13 @@ class LTX2VideoArchConfig(DiTArchConfig):
# LTX-2.3 gated extensions. All default OFF == LTX-2.0 behavior.
cross_attention_adaln: bool = False
caption_proj_before_connector: bool = False
# FFN bias, per stream. Some checkpoints ship the video FFN without bias
# (``ff_bias: false`` in their metadata), which drops the 96
# ``transformer_blocks.*.ff.net.{0.proj,2}.bias`` tensors; the audio FFN is
# configured independently and commonly keeps its biases. Both default True,
# which is the existing behavior.
ff_bias: bool = True
audio_ff_bias: bool = True
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(default_factory=lambda: [20, 2048, 2048])
@@ -62,8 +62,6 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
hidden_size: int = 5120
intermediate_size: int = 25600
num_hidden_layers: int = 64
output_hidden_state_index: int = 50
num_hidden_layers_override: int | None = 50
num_attention_heads: int = 64
num_key_value_heads: int = 8
head_dim: int = 128
@@ -109,7 +107,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
vision_initializer_range: float = 0.02
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
output_hidden_states: bool = False
output_hidden_states: bool = True
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
_fsdp_shard_conditions: list = field(default_factory=lambda: [
_is_language_transformer_layer,
@@ -120,17 +118,6 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
])
def __post_init__(self) -> None:
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
if self.num_hidden_layers_override is not None:
if self.num_hidden_layers_override <= 0:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
if self.num_hidden_layers_override < self.output_hidden_state_index:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
f"hidden_states[{self.output_hidden_state_index}], got "
f"{self.num_hidden_layers_override}.")
rope_scaling = dict(self.rope_scaling or {})
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
if not self.mrope_interleaved:
+18 -43
View File
@@ -808,19 +808,14 @@ class VideoGenerator:
latent_batch_size = _infer_latent_batch_size(batch)
is_latent_output = fastvideo_args.output_type == "latent"
needs_frame_output = batch.return_frames or (batch.save_video and not is_latent_output)
# A populated ``samples`` has exactly one consumer — the result
# dict (``"samples": samples if batch.return_frames else None``).
# Post-decode frame building reads ``output_batch.output``
# directly (the GPU ``vid_u8`` path), not ``samples``. So when
# ``return_frames=False`` the pinned fp32 alloc + D->H copy are
# dead weight — the CLI generate flow (``save_video=True``,
# ``return_frames=False``) hits this on every call.
# ``output_type == "latent"`` keeps its existing branch (shape
# mismatch falls through to ``.cpu()`` below) for callers that
# *do* ask for the latent samples via ``return_frames=True``.
needs_samples_buffer = batch.return_frames or needs_frame_output
# When ``output_type == "latent"`` the forward output has latent
# shape (e.g. ``[B, C_latent, T_latent, H_latent, W_latent]``)
# rather than the pre-allocation's pixel shape. Skip the pinned
# ~50 MB buffer entirely. Also skip it for metadata-only calls;
# neither the result nor save path will consume the decoded tensor.
# ``skip_pixel_prealloc`` also gates the slow-path warning.
needs_samples_out = batch.return_frames
skip_pixel_prealloc = is_latent_output or not needs_samples_out
skip_pixel_prealloc = is_latent_output or not needs_samples_buffer
if skip_pixel_prealloc:
samples = torch.empty(0, device='cpu')
else:
@@ -840,11 +835,9 @@ class VideoGenerator:
"This usually means the executor/pipeline failed earlier.")
audio_only = bool(output_batch.extra.get("audio_only"))
if not needs_samples_out:
# Nothing downstream reads ``samples`` (the result dict
# returns None when ``return_frames=False``); keep the empty
# placeholder allocated above and skip the fp32 D->H copy
# entirely.
if not needs_samples_buffer or (audio_only and not batch.return_frames):
# Metadata-only/audio-only request: keep the empty placeholder and
# avoid the decoded tensor D->H copy.
pass
elif audio_only:
# Audio-only return-frames requests expose the small placeholder
@@ -876,13 +869,8 @@ class VideoGenerator:
# `GenerationResult.size` describes the produced media, not only the
# base-stage request. Refiner pipelines can change the final pixel
# dimensions, so derive this result metadata from the decoded output.
# Read the geometry from `output_batch.output` (a shape-only access,
# no D->H copy): when `return_frames=False` the `samples` mirror
# stays an empty placeholder and no longer carries the decoded
# shape. Metadata-only calls keep the request fallback and never
# inspect the (possibly dropped) worker output.
output_size = _resolve_output_size(
output_batch.output if needs_frame_output else samples,
samples,
(target_height, target_width, batch.num_frames),
pixel_output=not is_latent_output and not audio_only,
)
@@ -894,26 +882,13 @@ class VideoGenerator:
elif not needs_frame_output:
frames = None
else:
# Quantize on the source device (typically CUDA) BEFORE the
# device->host copy. `samples` above is just the pinned-CPU
# mirror of `output_batch.output` (`samples.copy_(output)` or
# `output.cpu()`) with no intervening preprocessing, so reading
# `output_batch.output` here is the same data. The old path
# paid a full fp32 video D->H copy (which scales with
# resolution x frames x batch) and then a single-threaded
# per-frame CPU *255/cast loop. Casting to uint8 on-device
# makes the transfer 4x smaller, ships it in a single copy,
# and moves the elementwise work onto the GPU. clamp_() also
# fixes a latent overflow bug: VAE output slightly outside
# [0, 1] wrapped mod 256 in the old unclamped cast.
# (Equivalence is SSIM-gated, not bit-exact: float->uint8
# differs <=1 LSB CPU vs GPU.)
src = output_batch.output
vid_u8 = (src * 255).clamp_(0, 255).to(torch.uint8)
vid_u8 = rearrange(vid_u8, "b c t h w -> t b c h w").cpu()
frames = [
torchvision.utils.make_grid(x, nrow=6).permute(1, 2, 0).squeeze(-1).contiguous().numpy() for x in vid_u8
]
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.permute(1, 2, 0).squeeze(-1)
x = (x * 255).to(torch.uint8)
frames.append(x.contiguous().cpu().numpy())
postprocess_time = time.perf_counter() - postprocess_start
logger.info("PostDecodeFrameProcessStage completed in %.3f s", postprocess_time)
if logging_info is not None:
-29
View File
@@ -21,17 +21,12 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_MINIMAX_H3_FUSIONS: str = ""
FASTVIDEO_VAE_PARALLEL_DECODE: bool = False
FASTVIDEO_VAE_PARALLEL_ENCODE: bool = False
FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY: str | None = None
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_NVTX_PROFILE: bool = False
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
@@ -222,34 +217,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# If set (=1), MiniMax-H3 VAE decode (and, with the ENCODE variant,
# reference-video encode) round-robins its temporal chunks across the
# sequence-parallel ranks instead of running serially on the output rank.
# Folded into FastVideoArgs.vae_parallel_decode / vae_parallel_encode at
# construction (parse-once). The STRATEGY variant picks the chunk
# transport collective: "gather" (default) or "all_gather".
"FASTVIDEO_VAE_PARALLEL_DECODE":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE", "0") != "0",
"FASTVIDEO_VAE_PARALLEL_ENCODE":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_ENCODE", "0") != "0",
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY":
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY", None),
# Opt-in MiniMax-H3 inference-only Triton fusions adapted from the
# NVlabs/Sana Sol-Engine implementation. Accepts `all`, `1`, or a
# comma-separated subset of `modulate,qknorm_rope,swiglu`. An empty value
# (the default), `0`, or `none` keeps the eager implementation.
"FASTVIDEO_MINIMAX_H3_FUSIONS":
lambda: os.getenv("FASTVIDEO_MINIMAX_H3_FUSIONS", ""),
# Use dedicated multiprocess context for workers.
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
# Emit lightweight NVTX ranges for external profilers such as Nsight Systems.
"FASTVIDEO_NVTX_PROFILE":
lambda: os.getenv("FASTVIDEO_NVTX_PROFILE", "0") != "0",
# Enables torch profiler if set. Path to the directory where torch profiler
# traces are saved. Note that it must be an absolute path.
"FASTVIDEO_TORCH_PROFILER_DIR":
-51
View File
@@ -146,19 +146,6 @@ class FastVideoArgs:
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
# Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the
# video VAE's temporal chunks (decode) and clips (reference encode) are
# round-robined across the sequence-parallel ranks and reassembled
# bit-exactly on the group's first rank instead of running serially on
# one rank while the others idle. ``__post_init__`` folds the
# FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars
# into these fields (parse-once, like attention_backend), and
# FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport
# collective ("gather" or "all_gather").
vae_parallel_decode: bool = False
vae_parallel_encode: bool = False
vae_parallel_decode_strategy: str | None = None
# Compilation
# ``enable_torch_compile`` covers the DiT path (transformer,
# transformer_2, and the LTX-2 stage-2 transformer_refine).
@@ -182,7 +169,6 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
VSA_tile_size: int = 256 # VSA-H3 tile size (256 or 64); 64 = native Triton path
# V-MoBA parameters
moba_config_path: str | None = None
@@ -300,27 +286,8 @@ class FastVideoArgs:
env_backend = envs.FASTVIDEO_ATTENTION_BACKEND
if env_backend is not None and backend_name_to_enum(env_backend) is not None:
self.attention_backend = env_backend
self._fold_vae_parallel_env()
self.check_fastvideo_args()
def _fold_vae_parallel_env(self) -> None:
"""Parse-once adapters for the sequence-parallel VAE env vars."""
import fastvideo.envs as envs
# Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES /
# DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args
# never imports model modules; a unit test pins the two in sync).
strategies = ("gather", "all_gather")
if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE:
self.vae_parallel_decode = True
if not self.vae_parallel_encode and envs.FASTVIDEO_VAE_PARALLEL_ENCODE:
self.vae_parallel_encode = True
if self.vae_parallel_decode_strategy is None:
self.vae_parallel_decode_strategy = envs.FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY or "gather"
if self.vae_parallel_decode_strategy not in strategies:
raise ValueError(f"vae_parallel_decode_strategy must be one of {strategies}, "
f"got {self.vae_parallel_decode_strategy!r}.")
def _apply_transformer_quant(self) -> None:
"""Pin the typed ``transformer_quant`` instance onto ``dit_config``.
@@ -664,18 +631,6 @@ class FastVideoArgs:
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
"Should be enabled in almost all cases",
)
parser.add_argument(
"--vae-parallel-decode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 VAE decode chunks across the SP ranks "
"and reassemble bit-exactly on the output rank (default: serial decode on the output rank)",
)
parser.add_argument(
"--vae-parallel-encode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 reference-video VAE encode clips across "
"the SP ranks; every rank keeps the identical full encoding (default: serial encode on every rank)",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -689,12 +644,6 @@ class FastVideoArgs:
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
parser.add_argument(
"--VSA-tile-size",
type=int,
default=FastVideoArgs.VSA_tile_size,
help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path",
)
# Master port for distributed training/inference
parser.add_argument(
@@ -82,6 +82,79 @@ def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool:
return prefix in _LTX2_NVFP4_LINEAR_PREFIXES
def audit_nvfp4_coverage(
model,
predicate=None,
target_classes=None,
skip_classes=(),
) -> dict[str, object]:
"""Report which declared target classes actually attached, and raise if any did not.
The target set is written as module-prefix suffixes, but a tensor carries
three different names on its way here -- the checkpoint key, the
``named_modules()`` path, and the ``prefix=`` string a layer is constructed
with -- and only the last one reaches the predicate. A rename on any of the
others leaves this set syntactically valid and silently matching nothing,
and the attached COUNT alone cannot show it: a set that quietly stops
covering a whole class still reports a large, healthy-looking number.
So: a declared class matching zero modules is an error, not a warning. This
guard did NOT fire when it was written -- the shipped set is correct today.
It exists for the next rename, which is the only kind of failure that can
reach production looking like success.
Returns the receipt: attached count and params, per-class attached counts,
and the classes deliberately NOT targeted, so a reader can see the
exclusions were chosen rather than lost.
"""
from fastvideo.layers.linear import LinearBase
if predicate is None:
predicate = is_ltx2_nvfp4_linear_prefix
if target_classes is None:
target_classes = _LTX2_NVFP4_BLOCK_LINEAR_SUFFIXES
linears = [(getattr(m, "prefix", None) or n, n, m)
for n, m in model.named_modules() if isinstance(m, LinearBase)]
def _params(mod):
w = getattr(mod, "weight", None)
return int(w.shape[0] * w.shape[1]) if w is not None and w.dim() == 2 else 0
per_class: dict[str, int] = {}
for suffix in target_classes:
per_class[suffix] = sum(1 for pfx, _, _ in linears
if pfx.endswith("." + suffix) and predicate(pfx))
empty = sorted(k for k, v in per_class.items() if v == 0)
if empty:
raise ValueError(
"NVFP4 target classes matched ZERO modules: " + ", ".join(empty) +
". The model's actual Linear prefixes look like: " +
", ".join(sorted({p for p, _, _ in linears})[:5]) +
". Either the module naming changed or the target set is stale -- "
"refusing to train a model that silently quantizes less than declared.")
attached = [(p, n, m) for p, n, m in linears if predicate(p)]
skipped: dict[str, int] = {}
for pfx, name, _ in linears:
if predicate(pfx):
continue
tail = name.rsplit(".", 1)[-1]
for known in skip_classes:
if known in name:
skipped[known] = skipped.get(known, 0) + 1
break
else:
skipped["other:" + tail] = skipped.get("other:" + tail, 0) + 1
return {
"linears_total": len(linears),
"attached": len(attached),
"attached_params_M": round(sum(_params(m) for _, _, m in attached) / 1e6, 1),
"attached_per_class": per_class,
"skipped_by_rule": dict(sorted(skipped.items())),
}
def _is_ltx2_refine_only_prefix(prefix: str) -> bool:
return any(prefix.endswith(suffix) for suffix in _LTX2_REFINE_ONLY_SUFFIXES)
@@ -552,4 +625,5 @@ __all__ = [
"NVFP4QuantizeMethod",
"convert_model_to_nvfp4",
"is_ltx2_nvfp4_linear_prefix",
"audit_nvfp4_coverage",
]
+1 -4
View File
@@ -114,10 +114,7 @@ def _info(logger: Logger,
is_local_main_process = local_rank == 0
if (main_process_only and is_main_process) or (local_main_process_only and is_local_main_process):
# Honor an explicit stacklevel (info_once routes through here with
# stacklevel already set) instead of passing the keyword twice.
stacklevel = kwargs.pop("stacklevel", 2)
logger.log(logging.INFO, msg, *args, stacklevel=stacklevel, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
global _warned_local_main_process, _warned_main_process
-115
View File
@@ -1,115 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Experimental Apple MLX runtime helpers.
This package is intentionally small for now. It exists to grow the Apple-native
FastWan path in measurable steps: shape planning, primitive benchmarks, then
Wan block parity, then full DiT/runtime support.
"""
from fastvideo.mlx_runtime.fastwan import (
FastWanShape,
MLXQuantizationSpec,
MLXWanDiT,
MLXWanTransformerBlock,
UnsupportedMLXQuantizationError,
ensure_quantization_supported,
fastwan_shape,
fastwan_shape_from_config,
mlx_dit_from_diffusers_safetensors,
mlx_block_weights_from_torch,
mlx_block_weights_from_diffusers_safetensors,
quantization_support_error,
)
from fastvideo.mlx_runtime.checkpoint import (
load_mlx_dit_checkpoint,
save_mlx_dit_checkpoint,
)
from fastvideo.mlx_runtime.memory import (
AppliedMemoryLimits,
add_memory_limit_args,
apply_memory_limits,
gib_to_bytes,
)
from fastvideo.mlx_runtime.refine import (
DEFAULT_REFINE_SIGMA,
RefinePlan,
TwoPassResult,
default_refine_timesteps,
plan_refine_resolutions,
prepare_refine_latents,
refine_sigma_from_schedule,
run_dmd_loop,
run_two_pass_dmd,
upsample_latents_spatial,
)
from fastvideo.mlx_runtime.frame_upsample import (
DEFAULT_PIXEL_UPSAMPLE_MODE,
PIXEL_UPSAMPLE_MODES,
unsharp,
upsample_frame,
upsample_frames,
)
from fastvideo.mlx_runtime.fast_spatial import (
DEFAULT_FAST_SPATIAL_SHARPEN,
FastSpatialPlan,
apply_fast_spatial_upsample,
plan_fast_spatial,
resolve_spatial_mode,
)
from fastvideo.mlx_runtime.prompt_enhance import (
DEFAULT_ENHANCE_SYSTEM_PROMPT,
DEFAULT_MLX_LM_MODEL,
EnhanceResult,
enhance_prompt,
enhance_prompt_template,
enhance_result_as_metrics,
load_or_enhance_prompt,
)
__all__ = [
"AppliedMemoryLimits",
"DEFAULT_ENHANCE_SYSTEM_PROMPT",
"DEFAULT_MLX_LM_MODEL",
"DEFAULT_FAST_SPATIAL_SHARPEN",
"DEFAULT_PIXEL_UPSAMPLE_MODE",
"DEFAULT_REFINE_SIGMA",
"EnhanceResult",
"FastSpatialPlan",
"FastWanShape",
"MLXQuantizationSpec",
"MLXWanDiT",
"MLXWanTransformerBlock",
"RefinePlan",
"TwoPassResult",
"UnsupportedMLXQuantizationError",
"add_memory_limit_args",
"apply_fast_spatial_upsample",
"apply_memory_limits",
"enhance_prompt",
"enhance_prompt_template",
"enhance_result_as_metrics",
"ensure_quantization_supported",
"fastwan_shape",
"fastwan_shape_from_config",
"gib_to_bytes",
"load_mlx_dit_checkpoint",
"load_or_enhance_prompt",
"mlx_dit_from_diffusers_safetensors",
"mlx_block_weights_from_diffusers_safetensors",
"mlx_block_weights_from_torch",
"PIXEL_UPSAMPLE_MODES",
"default_refine_timesteps",
"plan_fast_spatial",
"plan_refine_resolutions",
"prepare_refine_latents",
"quantization_support_error",
"refine_sigma_from_schedule",
"resolve_spatial_mode",
"run_dmd_loop",
"run_two_pass_dmd",
"save_mlx_dit_checkpoint",
"unsharp",
"upsample_frame",
"upsample_frames",
"upsample_latents_spatial",
]
-273
View File
@@ -1,273 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Pre-quantized MLX checkpoint save/load for the FastWan DiT.
Loading the Diffusers fp32/fp16 checkpoint and quantizing at startup costs
both download size and load time on every run. This module persists an already
cast (and optionally already quantized) ``MLXWanDiT`` so 16 GB users download
and load roughly half the bytes and skip requantization entirely:
dit = mlx_dit_from_diffusers_safetensors(ckpt, cfg, quantization="int8")
save_mlx_dit_checkpoint(dit, "FastWan2.1-T2V-1.3B-mlx-int8")
...
dit = load_mlx_dit_checkpoint("FastWan2.1-T2V-1.3B-mlx-int8")
Format (one directory):
- ``mlx_dit.safetensors`` — every array, saved with ``mx.save_safetensors``.
Plain weights keep their key; a quantized weight ``K`` is stored as the
packed ``K`` plus ``K.scales`` (and ``K.biases`` for affine modes).
- ``mlx_dit.json`` — format version, the model config, the quantization spec,
and which keys are quantized, so the loader can rebuild ``QuantizedMatrix``
objects without guessing.
"""
from __future__ import annotations
import json
import shutil
import tempfile
from pathlib import Path
from typing import Any
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.fastwan import (
MLXQuantizationSpec,
MLXWanDiT,
MLXWanTransformerBlock,
QuantizedMatrix,
ensure_quantization_supported,
)
logger = init_logger(__name__)
FORMAT_VERSION = 1
WEIGHTS_FILENAME = "mlx_dit.safetensors"
MANIFEST_FILENAME = "mlx_dit.json"
_BLOCK_PREFIX = "blocks"
_DTYPE_TO_NAME = {"float16": "fp16", "bfloat16": "bf16", "float32": "fp32"}
def _dtype_name(dtype) -> str:
"""Return the manifest name for a supported MLX data type.
Parameters:
dtype: The MLX data type to convert.
Returns:
str: The manifest name corresponding to the data type.
Raises:
ValueError: If the data type is not supported for checkpointing.
"""
import mlx.core as mx
for raw, name in _DTYPE_TO_NAME.items():
if dtype == getattr(mx, raw):
return name
raise ValueError(f"Unsupported MLX dtype for checkpointing: {dtype}")
def _name_to_dtype(name: str):
"""Convert a manifest dtype name to its corresponding MLX dtype.
Parameters:
name (str): Manifest name, such as ``"fp16"``, ``"bf16"``, or ``"fp32"``.
Returns:
The corresponding MLX dtype.
"""
import mlx.core as mx
return {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[name]
def _flatten_weights(dit: MLXWanDiT) -> dict[str, Any]:
"""Combine model and transformer-block weights into a single flattened mapping.
Parameters:
dit (MLXWanDiT): Model whose weights should be flattened.
Returns:
dict[str, Any]: Mapping containing top-level weights and indexed transformer-block weights.
"""
flat: dict[str, Any] = dict(dit.weights)
for index, block in enumerate(dit.blocks):
for name, value in block.weights.items():
flat[f"{_BLOCK_PREFIX}.{index}.{name}"] = value
return flat
def save_mlx_dit_checkpoint(dit: MLXWanDiT, checkpoint_dir: str | Path) -> Path:
"""Save a plain or quantized MLX Wan DiT checkpoint to a directory.
Parameters:
dit (MLXWanDiT): Model whose weights and configuration will be saved.
checkpoint_dir (str | Path): Destination directory for the checkpoint.
Returns:
Path: Path to the checkpoint directory.
"""
import mlx.core as mx
checkpoint_dir = Path(checkpoint_dir)
arrays: dict[str, Any] = {}
quantized: dict[str, dict[str, Any]] = {}
spec: MLXQuantizationSpec | None = None
for key, value in _flatten_weights(dit).items():
if isinstance(value, QuantizedMatrix):
if spec is not None and value.spec != spec:
raise ValueError(f"Mixed quantization specs in one checkpoint ({spec} vs {value.spec} at '{key}') "
"are not supported.")
spec = value.spec
arrays[key] = value.weight
arrays[f"{key}.scales"] = value.scales
if value.biases is not None:
arrays[f"{key}.biases"] = value.biases
quantized[key] = {
"dequantized_dtype": _dtype_name(value.dequantized_dtype),
"has_biases": value.biases is not None,
}
else:
arrays[key] = value
manifest = {
"format_version": FORMAT_VERSION,
"config": dit.config,
"num_blocks": len(dit.blocks),
"quantization": None if spec is None else {
"mode": spec.mode,
"bits": spec.bits,
"group_size": spec.group_size,
},
"quantized_keys": quantized,
}
manifest_json = json.dumps(manifest, indent=2)
checkpoint_dir.parent.mkdir(parents=True, exist_ok=True)
staging_dir = Path(tempfile.mkdtemp(dir=checkpoint_dir.parent, prefix=f".{checkpoint_dir.name}.staging-"))
backup_root: Path | None = None
try:
staged_weights = staging_dir / WEIGHTS_FILENAME
staged_manifest = staging_dir / MANIFEST_FILENAME
mx.save_safetensors(str(staged_weights), arrays)
staged_manifest.write_text(manifest_json)
if checkpoint_dir.exists():
backup_root = Path(tempfile.mkdtemp(dir=checkpoint_dir.parent, prefix=f".{checkpoint_dir.name}.backup-"))
try:
checkpoint_dir.replace(backup_root / checkpoint_dir.name)
except Exception:
shutil.rmtree(backup_root, ignore_errors=True)
raise
try:
staging_dir.replace(checkpoint_dir)
except Exception:
if backup_root is not None:
(backup_root / checkpoint_dir.name).replace(checkpoint_dir)
shutil.rmtree(backup_root, ignore_errors=True)
raise
if backup_root is not None:
shutil.rmtree(backup_root, ignore_errors=True)
finally:
shutil.rmtree(staging_dir, ignore_errors=True)
logger.info("Saved MLX DiT checkpoint (%d arrays, quantization=%s) to %s", len(arrays),
spec.label if spec else "none", checkpoint_dir)
return checkpoint_dir
def load_mlx_dit_checkpoint(checkpoint_dir: str | Path, *, compile: bool = False) -> MLXWanDiT:
"""
Reconstruct an MLXWanDiT model from a versioned checkpoint.
Parameters:
checkpoint_dir (str | Path): Directory containing the checkpoint manifest and weights.
compile (bool): Whether to configure the reconstructed model for compilation.
Returns:
MLXWanDiT: The reconstructed model.
Raises:
FileNotFoundError: If the checkpoint manifest or weights file is missing.
ValueError: If the checkpoint format is unsupported or block weights are incomplete.
"""
import mlx.core as mx
checkpoint_dir = Path(checkpoint_dir)
manifest_path = checkpoint_dir / MANIFEST_FILENAME
weights_path = checkpoint_dir / WEIGHTS_FILENAME
if not manifest_path.exists() or not weights_path.exists():
raise FileNotFoundError(f"Not an MLX DiT checkpoint directory: {checkpoint_dir} "
f"(expected {MANIFEST_FILENAME} and {WEIGHTS_FILENAME}).")
manifest = json.loads(manifest_path.read_text())
version = manifest.get("format_version")
if version != FORMAT_VERSION:
raise ValueError(f"MLX DiT checkpoint {checkpoint_dir} has format_version={version}; "
f"this FastVideo build reads version {FORMAT_VERSION}. Re-export the checkpoint.")
spec = None
if manifest["quantization"] is not None:
spec = MLXQuantizationSpec(**manifest["quantization"])
# The packed layout of mx.quantize output is mode-specific, so a build
# that cannot run the mode cannot use these arrays at all.
ensure_quantization_supported(spec)
arrays = mx.load(str(weights_path))
quantized_keys: dict[str, dict[str, Any]] = manifest["quantized_keys"]
def rebuild(key: str):
"""
Reconstructs a weight array or quantized matrix from checkpoint data.
Parameters:
key (str): The weight key to rebuild.
Returns:
The stored array for an unquantized weight or a reconstructed quantized matrix.
"""
if key not in quantized_keys:
return arrays[key]
info = quantized_keys[key]
assert spec is not None, f"Quantized key '{key}' in a checkpoint without a quantization spec"
return QuantizedMatrix(
weight=arrays[key],
scales=arrays[f"{key}.scales"],
biases=arrays[f"{key}.biases"] if info["has_biases"] else None,
spec=spec,
dequantized_dtype=_name_to_dtype(info["dequantized_dtype"]),
)
config = manifest["config"]
block_keys: dict[int, list[str]] = {}
top_level_keys: list[str] = []
for key in arrays:
if key.endswith(".scales") or key.endswith(".biases"):
continue
if key.startswith(f"{_BLOCK_PREFIX}."):
index_str, _, _ = key[len(_BLOCK_PREFIX) + 1:].partition(".")
block_keys.setdefault(int(index_str), []).append(key)
else:
top_level_keys.append(key)
weights = {key: rebuild(key) for key in top_level_keys}
num_blocks = int(manifest["num_blocks"])
if sorted(block_keys) != list(range(num_blocks)):
raise ValueError(f"MLX DiT checkpoint {checkpoint_dir} is missing block weights: "
f"manifest says {num_blocks} blocks, found indices {sorted(block_keys)}.")
inner_dim = int(config["num_attention_heads"]) * int(config["attention_head_dim"])
blocks = []
for index in range(num_blocks):
prefix = f"{_BLOCK_PREFIX}.{index}."
block_weights = {key[len(prefix):]: rebuild(key) for key in block_keys[index]}
blocks.append(
MLXWanTransformerBlock(
block_weights,
dim=inner_dim,
ffn_dim=int(config["ffn_dim"]),
num_heads=int(config["num_attention_heads"]),
eps=float(config["eps"]),
))
return MLXWanDiT(weights, blocks, config, compile=compile)
-228
View File
@@ -1,228 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Spatial fast mode for the MLX Wan runtime (RIFE's spatial twin).
RIFE ``--fast`` cuts *frames* (temporal). This module cuts *pixels*
(spatial): denoise at ``target // scale``, decode at that size, then
resample the decoded frames up to the target. No second denoise pass —
that is ``--refine`` (quality). The two compose:
* ``--fast-spatial`` alone → speed (≈ scale² fewer tokens)
* ``--refine`` alone → quality two-pass (H3 / LTX-2)
* ``--fast`` + ``--refine`` → fewer frames at base res, full-res refine
* ``--fast`` + ``--fast-spatial`` → fewer frames *and* fewer pixels
The upsample runs in **pixel** space, after the VAE decode. It used to run
in latent space (bilinear over the latent H/W plane, sharing the refine
hand-off primitive) and that is what made spatial fast mode incoherent: an
interpolated Wan latent is off the decoder's manifold, so decode returned
the right silhouette under a smeared veil. ``--refine`` can get away with
the latent-space upsample because a second DMD pass re-denoises the result;
spatial fast mode hands the latent straight to the decoder, so it cannot.
See :mod:`fastvideo.mlx_runtime.frame_upsample` for the full rationale.
MetalFX is intentionally not used: it needs game-engine motion vectors
and depth that diffusion output lacks.
"""
from __future__ import annotations
from collections.abc import Iterable
from dataclasses import dataclass
import numpy as np
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.frame_upsample import (
DEFAULT_PIXEL_UPSAMPLE_MODE,
PIXEL_UPSAMPLE_MODES,
upsample_frames,
)
from fastvideo.mlx_runtime.refine import RefinePlan, plan_refine_resolutions
logger = init_logger(__name__)
# Resampling from a smaller decode loses high-frequency detail the same way
# RIFE's flow warp does, so spatial fast mode borrows ``--fast``'s remedy: a
# light unsharp pass. 0.4 recovers perceived crispness on Wan2.1 output at 2x
# without the halos that show up by ~0.8.
DEFAULT_FAST_SPATIAL_SHARPEN = 0.4
@dataclass(frozen=True)
class FastSpatialPlan:
"""Resolved geometry for a spatial-fast (upsample-only) run."""
plan: RefinePlan
upsample_mode: str
sharpen: float = DEFAULT_FAST_SPATIAL_SHARPEN
@property
def enabled(self) -> bool:
"""
Determine whether spatial scaling is enabled.
Returns:
`true` if the spatial scale is greater than one, `false` otherwise.
"""
return self.plan.spatial_scale > 1
@property
def scale(self) -> int:
"""Provides the configured spatial scaling factor.
Returns:
int: The spatial scaling factor.
"""
return self.plan.spatial_scale
@property
def target_height(self) -> int:
"""
Return the target output height for the spatial plan.
Returns:
int: Target output height in pixels.
"""
return self.plan.target_height
@property
def target_width(self) -> int:
"""Return the target image width in pixels.
Returns:
int: The target image width.
"""
return self.plan.target_width
@property
def stage1_height(self) -> int:
"""
Provide the stage-one latent height used for reduced-resolution processing.
Returns:
int: The stage-one latent height.
"""
return self.plan.stage1_height
@property
def stage1_width(self) -> int:
"""Get the stage-one latent width.
Returns:
int: The stage-one latent width.
"""
return self.plan.stage1_width
def plan_fast_spatial(
*,
height: int,
width: int,
num_frames: int,
spatial_scale: int = 2,
vae_spatial_compression: int = 8,
vae_temporal_compression: int = 4,
patch_size: tuple[int, int, int] = (1, 2, 2),
upsample_mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
sharpen: float = DEFAULT_FAST_SPATIAL_SHARPEN,
enabled: bool = True,
) -> FastSpatialPlan:
"""
Build a plan for reduced-resolution denoising followed by pixel-space upsampling.
Parameters:
upsample_mode (str): Pixel interpolation kernel, one of
:data:`~fastvideo.mlx_runtime.frame_upsample.PIXEL_UPSAMPLE_MODES`.
sharpen (float): Unsharp strength applied after the resize.
Returns:
FastSpatialPlan: The validated spatial-fast processing plan.
Raises:
ValueError: If the upsample mode is unsupported or ``sharpen`` is negative.
"""
if upsample_mode not in PIXEL_UPSAMPLE_MODES:
raise ValueError(f"Unsupported upsample mode: {upsample_mode!r} "
f"(expected one of {', '.join(PIXEL_UPSAMPLE_MODES)})")
if sharpen < 0.0:
raise ValueError(f"sharpen must be >= 0, got {sharpen}")
plan = plan_refine_resolutions(
height=height,
width=width,
num_frames=num_frames,
spatial_scale=spatial_scale,
vae_spatial_compression=vae_spatial_compression,
vae_temporal_compression=vae_temporal_compression,
patch_size=patch_size,
enabled=enabled,
mode_label="fast-spatial",
)
if plan.spatial_scale > 1:
logger.info(
"[MLX fast-spatial] denoise+decode %dx%d → upsample %dx to %dx%d (%s, sharpen=%.2f)",
plan.stage1_width,
plan.stage1_height,
plan.spatial_scale,
plan.target_width,
plan.target_height,
upsample_mode,
sharpen,
)
return FastSpatialPlan(plan=plan, upsample_mode=upsample_mode, sharpen=sharpen)
def apply_fast_spatial_upsample(
frames: Iterable[np.ndarray],
spatial: FastSpatialPlan,
) -> list[np.ndarray]:
"""Resample decoded stage-1 frames up to the target resolution.
This runs on decoded RGB frames, *not* on latents: see the module
docstring for why the latent-space version produced a blurred veil.
Parameters:
frames (Iterable[np.ndarray]): Decoded HxWx3 uint8 RGB frames, produced
by decoding at the stage-one resolution.
spatial (FastSpatialPlan): Plan defining the target size, interpolation
kernel, and unsharp strength.
Returns:
list[np.ndarray]: Frames at the target resolution. When spatial scaling
is disabled the frames are returned unchanged, as a list.
"""
if not spatial.enabled:
return list(frames)
return upsample_frames(
frames,
width=spatial.target_width,
height=spatial.target_height,
mode=spatial.upsample_mode,
sharpen=spatial.sharpen,
)
def resolve_spatial_mode(
*,
refine: bool,
fast_spatial: bool,
) -> str:
"""Select the active spatial processing mode, with refinement taking precedence.
Returns:
str: ``"refine"`` when refinement is enabled, ``"fast_spatial"`` when
spatial-fast processing is enabled, or ``"off"`` otherwise.
"""
if refine:
return "refine"
if fast_spatial:
return "fast_spatial"
return "off"
__all__ = [
"DEFAULT_FAST_SPATIAL_SHARPEN",
"FastSpatialPlan",
"apply_fast_spatial_upsample",
"plan_fast_spatial",
"resolve_spatial_mode",
]
-980
View File
@@ -1,980 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# mypy: disable-error-code=no-untyped-call
"""FastWan-oriented helpers for the experimental MLX runtime path."""
from __future__ import annotations
import json
import math
import os
import statistics
import time
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any
from collections.abc import Callable
from fastvideo.logger import init_logger
if TYPE_CHECKING:
import mlx.core as mx
import torch
logger = init_logger(__name__)
@dataclass(frozen=True)
class FastWanShape:
height: int
width: int
num_frames: int
latent_frames: int
latent_height: int
latent_width: int
patch_frames: int
patch_height: int
patch_width: int
tokens: int
hidden_size: int
num_heads: int
head_dim: int
class UnsupportedMLXQuantizationError(ValueError):
"""A quantization mode the installed MLX build cannot execute.
Raised by :func:`ensure_quantization_supported` before any model weights
are loaded, so callers (CLI flags, benchmark sweeps) can fail fast with an
actionable message -- or skip the mode -- instead of crashing deep inside
``mx.quantize`` mid-load.
"""
@dataclass(frozen=True)
class MLXQuantizationSpec:
"""MLX quantized-matmul configuration for DiT linear weights."""
mode: str
bits: int | None = None
group_size: int | None = None
@classmethod
def from_name(cls, name: str | None) -> MLXQuantizationSpec | None:
if name is None or name in {"", "none", "fp16", "fp32"}:
return None
if name == "int8":
return cls(mode="affine", bits=8, group_size=64)
if name == "int4":
return cls(mode="affine", bits=4, group_size=64)
if name == "mxfp8":
return cls(mode="mxfp8")
if name == "mxfp4":
return cls(mode="mxfp4")
if name == "nvfp4":
return cls(mode="nvfp4")
raise ValueError(f"Unsupported MLX quantization mode: {name}")
@property
def label(self) -> str:
if self.mode == "affine":
return f"int{self.bits}"
return self.mode
@dataclass(frozen=True)
class QuantizedMatrix:
weight: mx.array
scales: mx.array
biases: mx.array | None
spec: MLXQuantizationSpec
dequantized_dtype: mx.Dtype
def fastwan_shape(
*,
height: int,
width: int,
num_frames: int,
vae_temporal_compression: int = 4,
vae_spatial_compression: int = 8,
patch_size: tuple[int, int, int] = (1, 2, 2),
num_heads: int = 12,
head_dim: int = 128,
) -> FastWanShape:
"""Return the approximate DiT token shape for Wan/FastWan T2V inference."""
latent_frames = (num_frames - 1) // vae_temporal_compression + 1
latent_height = height // vae_spatial_compression
latent_width = width // vae_spatial_compression
patch_frames = latent_frames // patch_size[0]
patch_height = latent_height // patch_size[1]
patch_width = latent_width // patch_size[2]
tokens = patch_frames * patch_height * patch_width
return FastWanShape(
height=height,
width=width,
num_frames=num_frames,
latent_frames=latent_frames,
latent_height=latent_height,
latent_width=latent_width,
patch_frames=patch_frames,
patch_height=patch_height,
patch_width=patch_width,
tokens=tokens,
hidden_size=num_heads * head_dim,
num_heads=num_heads,
head_dim=head_dim,
)
def fastwan_shape_from_config(
config_path: str | Path,
*,
height: int,
width: int,
num_frames: int,
) -> FastWanShape:
config = json.loads(Path(config_path).read_text())
return fastwan_shape(
height=height,
width=width,
num_frames=num_frames,
patch_size=tuple(config["patch_size"]),
num_heads=int(config["num_attention_heads"]),
head_dim=int(config["attention_head_dim"]),
)
def replace_tokens(shape: FastWanShape, tokens: int) -> FastWanShape:
return FastWanShape(**{**shape.__dict__, "tokens": tokens})
def median_ms(samples: list[float]) -> float:
return statistics.median(samples) * 1000.0
def benchmark_mlx_attention(shape: FastWanShape, warmup: int, iters: int) -> float:
import mlx.core as mx
q = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
k = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
v = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
scale = shape.head_dim**-0.5
for _ in range(warmup):
y = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
mx.eval(y)
samples = []
for _ in range(iters):
start = time.perf_counter()
y = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
mx.eval(y)
samples.append(time.perf_counter() - start)
return median_ms(samples)
def benchmark_mlx_linear(shape: FastWanShape, warmup: int, iters: int) -> float:
import mlx.core as mx
x = mx.random.normal((shape.tokens, shape.hidden_size), dtype=mx.float16)
w = mx.random.normal((shape.hidden_size, shape.hidden_size), dtype=mx.float16)
b = mx.zeros((shape.hidden_size, ), dtype=mx.float16)
for _ in range(warmup):
y = x @ w + b
mx.eval(y)
samples = []
for _ in range(iters):
start = time.perf_counter()
y = x @ w + b
mx.eval(y)
samples.append(time.perf_counter() - start)
return median_ms(samples)
def benchmark_torch_mps_attention(shape: FastWanShape, warmup: int, iters: int) -> float | None:
try:
import torch
import torch.nn.functional as F
except ImportError:
return None
if not torch.backends.mps.is_available():
return None
device = torch.device("mps")
q = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
k = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
v = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
for _ in range(warmup):
y = F.scaled_dot_product_attention(q, k, v)
torch.mps.synchronize()
_ = y
samples = []
for _ in range(iters):
start = time.perf_counter()
y = F.scaled_dot_product_attention(q, k, v)
torch.mps.synchronize()
_ = y
samples.append(time.perf_counter() - start)
return median_ms(samples)
def torch_to_mx(tensor) -> mx.array:
import mlx.core as mx
return mx.array(tensor.detach().cpu().float().numpy())
def weight_dtype(weight):
if isinstance(weight, QuantizedMatrix):
return weight.dequantized_dtype
return weight.dtype
_QUANT_SUPPORT_CACHE: dict[tuple[str, int | None, int | None], str | None] = {}
def quantization_support_error(spec: MLXQuantizationSpec) -> str | None:
"""Probe whether the installed MLX build supports ``spec``.
Runs a tiny ``mx.quantize`` + ``mx.quantized_matmul`` with exactly the
arguments :func:`quantize_matrix` / :func:`linear` use, so the result
reflects the real runtime path. The affine (int8/int4) modes are stable
across MLX releases, but the ``mxfp8``/``mxfp4``/``nvfp4`` mode strings
require newer MLX builds and raise otherwise. Returns ``None`` when the
mode works, else the underlying error message. Cached per spec.
"""
key = (spec.mode, spec.bits, spec.group_size)
if key not in _QUANT_SUPPORT_CACHE:
import mlx.core as mx
try:
probe_dim = max(spec.group_size or 0, 64)
weight = mx.zeros((probe_dim, probe_dim), dtype=mx.float16)
quantized = quantize_matrix(weight, spec)
y = linear(mx.zeros((1, probe_dim), dtype=mx.float16), quantized)
mx.eval(y)
_QUANT_SUPPORT_CACHE[key] = None
except Exception as exc: # noqa: BLE001 - MLX raises varied error types per backend/version.
_QUANT_SUPPORT_CACHE[key] = f"{type(exc).__name__}: {exc}"
return _QUANT_SUPPORT_CACHE[key]
def ensure_quantization_supported(spec: MLXQuantizationSpec | None) -> None:
"""Raise :class:`UnsupportedMLXQuantizationError` if ``spec`` cannot run here."""
if spec is None:
return
error = quantization_support_error(spec)
if error is None:
return
import mlx.core as mx
mlx_version = getattr(mx, "__version__", "unknown")
raise UnsupportedMLXQuantizationError(f"MLX quantization mode '{spec.label}' is not supported by the installed mlx "
f"({mlx_version}): {error}. Upgrade mlx or pick a supported mode "
f"(int8 is currently the most reliable quality/memory target).")
def quantize_matrix(weight, spec: MLXQuantizationSpec | None):
if spec is None:
return weight
import mlx.core as mx
if len(weight.shape) < 2:
return weight
q = mx.quantize(weight, group_size=spec.group_size, bits=spec.bits, mode=spec.mode)
biases = q[2] if len(q) == 3 else None
eval_args = [q[0], q[1]]
if biases is not None:
eval_args.append(biases)
mx.eval(*eval_args)
return QuantizedMatrix(
weight=q[0],
scales=q[1],
biases=biases,
spec=spec,
dequantized_dtype=weight.dtype,
)
def linear(x, weight, bias=None):
import mlx.core as mx
if isinstance(weight, QuantizedMatrix):
y = mx.quantized_matmul(
x,
weight.weight,
weight.scales,
weight.biases,
transpose=True,
group_size=weight.spec.group_size,
bits=weight.spec.bits,
mode=weight.spec.mode,
).astype(x.dtype)
else:
y = x @ weight.T
if bias is not None:
y = y + bias
return y
def _use_fast_norm() -> bool:
"""Opt-in to MLX's fused ``mx.fast`` normalization kernels.
Off by default so the numerically-explicit reference path stays the
baseline. Set ``FASTVIDEO_MLX_FAST_NORM=1`` to route LayerNorm/RMSNorm
through single fused Metal kernels (fewer intermediates, less memory
traffic) and benchmark the speedup.
"""
import os
return os.environ.get("FASTVIDEO_MLX_FAST_NORM", "0") == "1"
def layer_norm(x, weight=None, bias=None, eps: float = 1e-6):
import mlx.core as mx
if _use_fast_norm():
# Compute in fp32 (matching the reference below) so downstream dtype
# and precision are identical across call sites.
w = weight.astype(mx.float32) if weight is not None else None
b = bias.astype(mx.float32) if bias is not None else None
return mx.fast.layer_norm(x.astype(mx.float32), w, b, eps)
x_float = x.astype(mx.float32)
mean = mx.mean(x_float, axis=-1, keepdims=True)
var = mx.mean(mx.square(x_float - mean), axis=-1, keepdims=True)
y = (x_float - mean) * mx.rsqrt(var + eps)
if weight is not None:
y = y * weight
if bias is not None:
y = y + bias
return y
def rms_norm(x, weight, eps: float = 1e-6):
import mlx.core as mx
if _use_fast_norm():
return mx.fast.rms_norm(x, weight, eps)
orig_dtype = x.dtype
x_float = x.astype(mx.float32)
variance = mx.mean(mx.square(x_float), axis=-1, keepdims=True)
y = x_float * mx.rsqrt(variance + eps)
return y.astype(orig_dtype) * weight
def apply_rotary_emb(x, cos, sin, *, is_neox_style: bool = False):
"""Apply FastVideo's rotary convention to MLX tensors.
Args:
x: [batch, seq, heads, head_dim]
cos/sin: [seq, head_dim] for Wan's full-dimension rotate-pair style,
or [seq, head_dim // 2] for traditional RoPE.
"""
import mlx.core as mx
head_size = x.shape[-1]
rope_dim = cos.shape[-1]
cos = cos[None, :, None, :]
sin = sin[None, :, None, :]
x_float = x.astype(mx.float32)
if rope_dim == head_size:
x_pairs = x_float.reshape(*x.shape[:-1], -1, 2)
x_real = x_pairs[..., 0]
x_imag = x_pairs[..., 1]
x_rotated = mx.stack([-x_imag, x_real], axis=-1).reshape(*x.shape)
return (x_float * cos + x_rotated * sin).astype(x.dtype)
if is_neox_style:
x1, x2 = mx.split(x_float, 2, axis=-1)
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
return mx.concatenate([o1, o2], axis=-1).astype(x.dtype)
x1 = x_float[..., ::2]
x2 = x_float[..., 1::2]
o1 = x1 * cos - x2 * sin
o2 = x2 * cos + x1 * sin
return mx.stack([o1, o2], axis=-1).reshape(*x.shape).astype(x.dtype)
_WINDOWED_ATTENTION_WARNED = False
def _warn_windowed_attention_once(window: int) -> None:
"""Warn that FASTVIDEO_MLX_WINDOW degrades output on a dense-trained DiT.
Sliding-window self-attention is fast (6.6x at a +-3-frame window on 1.3B)
but these checkpoints were trained with dense attention, and restricting it
at inference produces heavy colour-block noise: structural agreement with
the dense baseline drops to 0.25 at +-3 frames and 0.03 at +-5. Sparsity of
this kind is a training-time method. Kept as a research knob, but it should
never be on by accident.
"""
global _WINDOWED_ATTENTION_WARNED
if _WINDOWED_ATTENTION_WARNED:
return
_WINDOWED_ATTENTION_WARNED = True
logger.warning(
"FASTVIDEO_MLX_WINDOW=%d enables sliding-window self-attention. These "
"checkpoints are trained dense; expect severely degraded output. This is "
"a research knob, not a speed setting — use --fast-spatial for real "
"denoise savings.",
window,
)
def gelu_tanh(x):
"""tanh-approximate GELU, as used by Wan's FFN.
``mlx.nn.gelu_approx`` is the same tanh approximation behind a fused
kernel. On the 1.3B FFN shape (32760x8960) it is bit-identical to the
expanded expression below and 3.3x faster — 28.9ms -> 8.7ms per layer,
which is 0.6s per denoise step across 30 layers.
"""
import mlx.nn as nn
return nn.gelu_approx(x)
def silu(x):
import mlx.core as mx
return x * mx.sigmoid(x)
def timestep_embedding(t, dim: int, max_period: int = 10000):
import mlx.core as mx
half = dim // 2
freqs = mx.exp(-math.log(max_period) * mx.arange(0, half, dtype=mx.float32) / half)
args = t[:, None].astype(mx.float32) * freqs[None]
embedding = mx.concatenate([mx.cos(args), mx.sin(args)], axis=-1)
if dim % 2:
embedding = mx.concatenate([embedding, mx.zeros_like(embedding[:, :1])], axis=-1)
return embedding
def scale_residual(residual, x, gate):
return residual + x * gate
def scale_residual_layer_norm_scale_shift(residual, x, gate, shift, scale, weight=None, bias=None, eps: float = 1e-6):
if isinstance(gate, int):
assert gate == 1
residual_output = residual + x
else:
residual_output = residual + x * gate
normalized = layer_norm(residual_output, weight=weight, bias=bias, eps=eps)
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
class MLXWanT2VCrossAttention:
def __init__(self, weights: dict[str, mx.array], *, dim: int, num_heads: int, eps: float = 1e-6) -> None:
self.weights = weights
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
def __call__(self, x, context):
import mlx.core as mx
batch = x.shape[0]
q = linear(x, self.weights["attn2.to_q.weight"], self.weights.get("attn2.to_q.bias"))
q = rms_norm(q, self.weights["attn2.norm_q.weight"], eps=self.eps).reshape(batch, -1, self.num_heads,
self.head_dim)
if context.shape[1] == 0:
attended = mx.zeros_like(q)
else:
k = linear(context, self.weights["attn2.to_k.weight"], self.weights.get("attn2.to_k.bias"))
k = rms_norm(k, self.weights["attn2.norm_k.weight"],
eps=self.eps).reshape(batch, -1, self.num_heads, self.head_dim)
v = linear(context, self.weights["attn2.to_v.weight"],
self.weights.get("attn2.to_v.bias")).reshape(batch, -1, self.num_heads, self.head_dim)
attended = mx.fast.scaled_dot_product_attention(
q.transpose(0, 2, 1, 3),
k.transpose(0, 2, 1, 3),
v.transpose(0, 2, 1, 3),
scale=self.head_dim**-0.5,
).transpose(0, 2, 1, 3)
attended = attended.reshape(batch, -1, self.dim)
return linear(attended, self.weights["attn2.to_out.weight"], self.weights.get("attn2.to_out.bias"))
class MLXWanTransformerBlock:
"""Dense T2V Wan transformer block for the experimental MLX runtime.
This mirrors the non-VSA PyTorch block for single-process dense attention.
Rotary embeddings and sequence-parallel paths are intentionally left out of
this first parity target.
"""
def __init__(self, weights: dict[str, mx.array], *, dim: int, ffn_dim: int, num_heads: int, eps: float = 1e-6):
self.weights = weights
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
self.attn2 = MLXWanT2VCrossAttention(weights, dim=dim, num_heads=num_heads, eps=eps)
def __call__(self, hidden_states, encoder_hidden_states, temb, freqs_cis=None):
import mlx.core as mx
orig_dtype = hidden_states.dtype
e = self.weights["scale_shift_table"] + temb.astype(mx.float32)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = mx.split(e, 6, axis=1)
norm_hidden_states = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm_hidden_states = (norm_hidden_states * (1.0 + scale_msa) + shift_msa).astype(orig_dtype)
query = linear(norm_hidden_states, self.weights["to_q.weight"], self.weights.get("to_q.bias"))
key = linear(norm_hidden_states, self.weights["to_k.weight"], self.weights.get("to_k.bias"))
value = linear(norm_hidden_states, self.weights["to_v.weight"], self.weights.get("to_v.bias"))
query = rms_norm(query, self.weights["norm_q.weight"],
eps=self.eps).reshape(hidden_states.shape[0], -1, self.num_heads, self.head_dim)
key = rms_norm(key, self.weights["norm_k.weight"], eps=self.eps).reshape(hidden_states.shape[0], -1,
self.num_heads, self.head_dim)
value = value.reshape(hidden_states.shape[0], -1, self.num_heads, self.head_dim)
if freqs_cis is not None:
cos, sin = freqs_cis
query = apply_rotary_emb(query, cos, sin, is_neox_style=False)
key = apply_rotary_emb(key, cos, sin, is_neox_style=False)
# Self-attention only. FASTVIDEO_MLX_WINDOW=0/unset → full SDPA (byte-identical
# to the historical path). When >0, use chunked symmetric sliding-window
# attention (see windowed_attention.py). Cross-attn (attn2) stays full.
# Optional FASTVIDEO_MLX_WINDOW_SINK (default 0) adds global sink tokens.
q_bh = query.transpose(0, 2, 1, 3) # (B, H, S, D)
k_bh = key.transpose(0, 2, 1, 3)
v_bh = value.transpose(0, 2, 1, 3)
scale = self.head_dim**-0.5
window = int(os.environ.get("FASTVIDEO_MLX_WINDOW", "0") or "0")
if window > 0:
from fastvideo.mlx_runtime.windowed_attention import windowed_attention
_warn_windowed_attention_once(window)
sink = int(os.environ.get("FASTVIDEO_MLX_WINDOW_SINK", "0") or "0")
attn_output = windowed_attention(q_bh, k_bh, v_bh, window=window, sink=sink, scale=scale)
else:
attn_output = mx.fast.scaled_dot_product_attention(q_bh, k_bh, v_bh, scale=scale)
attn_output = attn_output.transpose(0, 2, 1, 3)
attn_output = attn_output.reshape(hidden_states.shape[0], -1, self.dim)
attn_output = linear(attn_output, self.weights["to_out.weight"], self.weights.get("to_out.bias"))
norm_hidden_states, hidden_states = scale_residual_layer_norm_scale_shift(
hidden_states,
attn_output,
gate_msa,
0.0,
0.0,
weight=self.weights["self_attn_residual_norm.norm.weight"],
bias=self.weights["self_attn_residual_norm.norm.bias"],
eps=self.eps,
)
norm_hidden_states = norm_hidden_states.astype(orig_dtype)
hidden_states = hidden_states.astype(orig_dtype)
attn_output = self.attn2(norm_hidden_states, encoder_hidden_states)
norm_hidden_states, hidden_states = scale_residual_layer_norm_scale_shift(
hidden_states,
attn_output,
1,
c_shift_msa,
c_scale_msa,
eps=self.eps,
)
norm_hidden_states = norm_hidden_states.astype(orig_dtype)
hidden_states = hidden_states.astype(orig_dtype)
ff_output = linear(norm_hidden_states, self.weights["ffn.fc_in.weight"], self.weights.get("ffn.fc_in.bias"))
ff_output = gelu_tanh(ff_output)
ff_output = linear(ff_output, self.weights["ffn.fc_out.weight"], self.weights.get("ffn.fc_out.bias"))
hidden_states = scale_residual(hidden_states, ff_output, c_gate_msa)
return hidden_states.astype(orig_dtype)
def mlx_block_weights_from_torch(torch_block) -> dict[str, mx.array]:
return {name: torch_to_mx(value) for name, value in torch_block.state_dict().items()}
class MLXWanDiT:
"""Experimental FP16 Wan/FastWan DiT forward path in MLX."""
def __init__(
self,
weights: dict[str, mx.array],
blocks: list[MLXWanTransformerBlock],
config: dict,
*,
compile: bool = False,
) -> None:
import os
self.weights = weights
self.blocks = blocks
self.config = config
self.num_heads = int(config["num_attention_heads"])
self.head_dim = int(config["attention_head_dim"])
self.hidden_size = self.num_heads * self.head_dim
self.ffn_dim = int(config["ffn_dim"])
self.in_channels = int(config["in_channels"])
self.out_channels = int(config["out_channels"])
self.patch_size = tuple(config["patch_size"])
self.freq_dim = int(config["freq_dim"])
# Opt-in graph fusion. With fixed weights and static shapes, the whole
# denoise-step forward is a pure function of (latents, timestep) -- a
# good mx.compile target. Off by default so the eager path stays the
# baseline; enable via constructor or FASTVIDEO_MLX_COMPILE=1 and verify
# with the benchmark's SSIM ~= 1.0 check.
self._enable_compile = compile or os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1"
self._compiled_forward: Callable[..., Any] | None = None
self._compiled_signature: tuple | None = None
def patch_embed(self, hidden_states):
batch, channels, frames, height, width = hidden_states.shape
pt, ph, pw = self.patch_size
patch_dim = channels * pt * ph * pw
x = hidden_states.reshape(batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw)
x = x.transpose(0, 2, 4, 6, 1, 3, 5, 7).reshape(batch, -1, patch_dim)
return linear(x, self.weights["patch_embedding.weight"], self.weights.get("patch_embedding.bias"))
def condition(self, timestep, encoder_hidden_states):
t_freq = timestep_embedding(timestep, self.freq_dim).astype(
weight_dtype(self.weights["condition_embedder.time_embedder.linear_1.weight"]))
temb = linear(
t_freq,
self.weights["condition_embedder.time_embedder.linear_1.weight"],
self.weights["condition_embedder.time_embedder.linear_1.bias"],
)
temb = silu(temb)
temb = linear(
temb,
self.weights["condition_embedder.time_embedder.linear_2.weight"],
self.weights["condition_embedder.time_embedder.linear_2.bias"],
)
timestep_proj = silu(temb)
timestep_proj = linear(
timestep_proj,
self.weights["condition_embedder.time_proj.weight"],
self.weights["condition_embedder.time_proj.bias"],
).reshape(timestep.shape[0], 6, self.hidden_size)
encoder_hidden_states = linear(
encoder_hidden_states,
self.weights["condition_embedder.text_embedder.linear_1.weight"],
self.weights["condition_embedder.text_embedder.linear_1.bias"],
)
encoder_hidden_states = gelu_tanh(encoder_hidden_states)
encoder_hidden_states = linear(
encoder_hidden_states,
self.weights["condition_embedder.text_embedder.linear_2.weight"],
self.weights["condition_embedder.text_embedder.linear_2.bias"],
)
return temb, timestep_proj, encoder_hidden_states
def output(self, hidden_states, temb, *, batch: int, frames: int, height: int, width: int):
pt, ph, pw = self.patch_size
post_patch_frames = frames // pt
post_patch_height = height // ph
post_patch_width = width // pw
shift, scale = mx_split_two(self.weights["scale_shift_table"] + temb[:, None, :], axis=1)
hidden_states = layer_norm(hidden_states, eps=float(self.config["eps"])) * (1.0 + scale) + shift
hidden_states = hidden_states.astype(weight_dtype(self.weights["proj_out.weight"]))
hidden_states = linear(hidden_states, self.weights["proj_out.weight"], self.weights["proj_out.bias"])
hidden_states = hidden_states.reshape(
batch,
post_patch_frames,
post_patch_height,
post_patch_width,
pt,
ph,
pw,
self.out_channels,
)
hidden_states = hidden_states.transpose(0, 7, 1, 4, 2, 5, 3, 6)
return hidden_states.reshape(batch, self.out_channels, frames, height, width)
def _forward(self, hidden_states, encoder_hidden_states, timestep, cos, sin):
"""Pure forward used both eagerly and as the mx.compile target.
``cos``/``sin`` are passed as separate array args (rather than a tuple)
so the function traces cleanly under mx.compile.
"""
batch, _, frames, height, width = hidden_states.shape
freqs_cis = (cos, sin) if cos is not None else None
hidden_states = self.patch_embed(hidden_states)
temb, timestep_proj, encoder_hidden_states = self.condition(timestep, encoder_hidden_states)
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, freqs_cis=freqs_cis)
return self.output(hidden_states, temb, batch=batch, frames=frames, height=height, width=width)
def __call__(self, hidden_states, encoder_hidden_states, timestep, freqs_cis):
cos, sin = freqs_cis if freqs_cis is not None else (None, None)
if self._enable_compile and cos is not None:
import mlx.core as mx
# mx.compile keeps one traced graph per input signature, and each
# graph pins its own materialization of the quantized weights. The
# two-pass modes (--refine) call the DiT at a second resolution, so
# keeping both graphs alive doubles resident DiT memory: 14B refine
# peaked at 34.7 GiB instead of 20.8 GiB. Retire the previous graph
# when the signature changes; the retrace costs far less than a
# second copy of the weights.
signature = (hidden_states.shape, encoder_hidden_states.shape, timestep.shape)
if self._compiled_forward is not None and signature != self._compiled_signature:
self._compiled_forward = None
self._compiled_signature = None
mx.clear_cache()
if self._compiled_forward is None:
self._compiled_forward = mx.compile(self._forward)
self._compiled_signature = signature
compiled_forward = self._compiled_forward
try:
return compiled_forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
except Exception as exc: # noqa: BLE001 - some quant graphs may not trace; fall back to eager.
logger.warning("mx.compile forward failed (%s); falling back to eager execution.", exc)
self._enable_compile = False
self._compiled_forward = None
self._compiled_signature = None
return self._forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
def mx_split_two(x, *, axis: int):
import mlx.core as mx
left, right = mx.split(x, 2, axis=axis)
return left, right
def _load_safetensor_value(handle, name: str):
return handle.get_tensor(name)
def _load_mx_array_from_safetensor(handle, name: str, dtype):
"""Load a safetensors value and cast before creating the MLX array.
The FastWan Diffusers checkpoint is fp32. Creating an MLX array first and
then casting it to fp16 briefly materializes a large fp32 MLX allocation.
Casting the CPU tensor before crossing into MLX keeps the transient GPU-side
footprint lower.
"""
import mlx.core as mx
import torch
tensor = handle.get_tensor(name)
if dtype == mx.float16:
tensor = tensor.to(torch.float16)
elif dtype == mx.float32:
tensor = tensor.to(torch.float32)
elif dtype == mx.bfloat16:
# NumPy has no bfloat16, so bridge through fp32 and cast on-device below.
tensor = tensor.to(torch.float32)
array = mx.array(tensor.numpy())
del tensor
if dtype is not None and array.dtype != dtype:
array = array.astype(dtype)
mx.eval(array)
return array
def _eval_loaded_weight(value) -> None:
import mlx.core as mx
if isinstance(value, QuantizedMatrix):
eval_args = [value.weight, value.scales]
if value.biases is not None:
eval_args.append(value.biases)
mx.eval(*eval_args)
else:
mx.eval(value)
# Diffusers-to-FastVideo key mapping for WanTransformerBlock weights.
# Shared by both MLX and torch block loaders to keep mappings synchronized.
_WAN_BLOCK_KEY_MAP = {
"scale_shift_table": "scale_shift_table",
"attn1.to_q.weight": "to_q.weight",
"attn1.to_q.bias": "to_q.bias",
"attn1.to_k.weight": "to_k.weight",
"attn1.to_k.bias": "to_k.bias",
"attn1.to_v.weight": "to_v.weight",
"attn1.to_v.bias": "to_v.bias",
"attn1.to_out.0.weight": "to_out.weight",
"attn1.to_out.0.bias": "to_out.bias",
"attn1.norm_q.weight": "norm_q.weight",
"attn1.norm_k.weight": "norm_k.weight",
"attn2.to_q.weight": "attn2.to_q.weight",
"attn2.to_q.bias": "attn2.to_q.bias",
"attn2.to_k.weight": "attn2.to_k.weight",
"attn2.to_k.bias": "attn2.to_k.bias",
"attn2.to_v.weight": "attn2.to_v.weight",
"attn2.to_v.bias": "attn2.to_v.bias",
"attn2.to_out.0.weight": "attn2.to_out.weight",
"attn2.to_out.0.bias": "attn2.to_out.bias",
"attn2.norm_q.weight": "attn2.norm_q.weight",
"attn2.norm_k.weight": "attn2.norm_k.weight",
"ffn.net.0.proj.weight": "ffn.fc_in.weight",
"ffn.net.0.proj.bias": "ffn.fc_in.bias",
"ffn.net.2.weight": "ffn.fc_out.weight",
"ffn.net.2.bias": "ffn.fc_out.bias",
"norm2.weight": "self_attn_residual_norm.norm.weight",
"norm2.bias": "self_attn_residual_norm.norm.bias",
}
def mlx_block_weights_from_diffusers_safetensors(
checkpoint_path: str | Path,
*,
block_index: int = 0,
quantization: str | MLXQuantizationSpec | None = None,
dtype=None,
) -> dict[str, mx.array]:
"""Load one Diffusers-format Wan block into the MLX dense-block key layout."""
from safetensors import safe_open
prefix = f"blocks.{block_index}."
key_map = _WAN_BLOCK_KEY_MAP
spec = MLXQuantizationSpec.from_name(quantization) if (quantization is None
or isinstance(quantization, str)) else quantization
ensure_quantization_supported(spec)
matrix_targets = {target for target in key_map.values() if target.endswith(".weight") and "norm" not in target}
weights = {}
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for source_name, target_name in key_map.items():
full = prefix + source_name
if full not in available:
# Biases are optional: e.g. Wan2.1-14B has bias-free attention/FFN.
# The block forward already fetches biases via ``.get(...)``.
if source_name.endswith(".bias"):
continue
raise KeyError(f"missing required block weight: {full}")
array = _load_mx_array_from_safetensor(handle, full, dtype)
loaded = quantize_matrix(array, spec) if target_name in matrix_targets else array
_eval_loaded_weight(loaded)
weights[target_name] = loaded
del array
return weights
def mlx_dit_from_diffusers_safetensors(
checkpoint_path: str | Path,
config_path: str | Path,
*,
dtype: str = "fp16",
num_blocks: int | None = None,
quantization: str | MLXQuantizationSpec | None = None,
compile: bool = False,
) -> MLXWanDiT:
import mlx.core as mx
from safetensors import safe_open
config = json.loads(Path(config_path).read_text())
total_blocks = int(config["num_layers"])
if num_blocks is None:
num_blocks = total_blocks
cast_dtype = {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[dtype]
spec = MLXQuantizationSpec.from_name(quantization) if (quantization is None
or isinstance(quantization, str)) else quantization
ensure_quantization_supported(spec)
top_level_names = [
"patch_embedding.weight",
"patch_embedding.bias",
"condition_embedder.time_embedder.linear_1.weight",
"condition_embedder.time_embedder.linear_1.bias",
"condition_embedder.time_embedder.linear_2.weight",
"condition_embedder.time_embedder.linear_2.bias",
"condition_embedder.time_proj.weight",
"condition_embedder.time_proj.bias",
"condition_embedder.text_embedder.linear_1.weight",
"condition_embedder.text_embedder.linear_1.bias",
"condition_embedder.text_embedder.linear_2.weight",
"condition_embedder.text_embedder.linear_2.bias",
"scale_shift_table",
"proj_out.weight",
"proj_out.bias",
]
weights = {}
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for name in top_level_names:
if name not in available:
if name.endswith(".bias"):
continue
raise KeyError(f"missing required weight: {name}")
array = _load_mx_array_from_safetensor(handle, name, cast_dtype)
if name == "patch_embedding.weight":
array = array.reshape(int(config["num_attention_heads"]) * int(config["attention_head_dim"]), -1)
if name.endswith(".weight") and name not in {"scale_shift_table"}:
loaded = quantize_matrix(array, spec)
else:
loaded = array
_eval_loaded_weight(loaded)
weights[name] = loaded
del array
blocks = []
for block_index in range(num_blocks):
block_weights = mlx_block_weights_from_diffusers_safetensors(
checkpoint_path,
block_index=block_index,
quantization=spec,
dtype=cast_dtype,
)
block_weights = {
name: (value if isinstance(value, QuantizedMatrix) else value.astype(cast_dtype))
for name, value in block_weights.items()
}
for value in block_weights.values():
_eval_loaded_weight(value)
blocks.append(
MLXWanTransformerBlock(
block_weights,
dim=int(config["num_attention_heads"]) * int(config["attention_head_dim"]),
ffn_dim=int(config["ffn_dim"]),
num_heads=int(config["num_attention_heads"]),
eps=float(config["eps"]),
))
return MLXWanDiT(weights, blocks, config, compile=compile)
def torch_block_state_from_diffusers_safetensors(
checkpoint_path: str | Path,
*,
block_index: int = 0,
) -> dict[str, torch.Tensor]:
"""Load one Diffusers-format Wan block into FastVideo's dense block keys."""
from safetensors import safe_open
prefix = f"blocks.{block_index}."
key_map = _WAN_BLOCK_KEY_MAP
state = {}
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
available = set(handle.keys())
for source_name, target_name in key_map.items():
full = prefix + source_name
if full not in available:
# Biases are optional: e.g. Wan2.1-14B has bias-free attention/FFN.
if source_name.endswith(".bias"):
continue
raise KeyError(f"missing required block weight: {full}")
state[target_name] = handle.get_tensor(full).float()
return state
-154
View File
@@ -1,154 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Pixel-space spatial resampling for decoded MLX Wan frames.
Spatial fast mode denoises on a smaller latent grid and has to get back to
the requested output size. That resize belongs *here* — after the VAE
decode — and not in latent space.
A Wan latent cell is a learned code for an 8x8 (Wan2.1) or 16x16 (Wan2.2)
pixel block, not a low-pass sample of the image. Linearly blending two
adjacent codes does not produce the code of the blended blocks; it produces
a vector the decoder was never trained on. The decoder answers with smeared,
ringing texture laid over otherwise-correct structure — the silhouette
survives, the detail turns to haze. Measured on Wan2.1-1.3B at 480x832, a
2x bilinear latent upsample destroys 62% of the latent's high-frequency
energy while leaving its overall magnitude intact, which is exactly the
signature of that veil.
Resampling decoded RGB frames has none of that problem: an image *is* a
sampled 2-D signal, so Lanczos/cubic interpolation is the operation it was
defined for. The result is soft — it carries stage-1's real detail budget
and no more — but it is clean and coherent.
"""
from __future__ import annotations
from collections.abc import Iterable
import numpy as np
# Pixel-space interpolation kernels, best-quality first. ``lanczos`` is the
# default: it holds edges better than cubic at 2x with no visible ringing on
# decoder output, which is already band-limited.
PIXEL_UPSAMPLE_MODES = ("lanczos", "cubic", "bilinear", "nearest")
DEFAULT_PIXEL_UPSAMPLE_MODE = "lanczos"
def _interpolation_flag(mode: str) -> int:
"""
Map a pixel upsample mode name onto its OpenCV interpolation flag.
Parameters:
mode (str): One of :data:`PIXEL_UPSAMPLE_MODES`.
Returns:
int: The matching ``cv2.INTER_*`` flag.
Raises:
ValueError: If the mode is not a supported pixel upsample mode.
"""
import cv2
flags = {
"lanczos": cv2.INTER_LANCZOS4,
"cubic": cv2.INTER_CUBIC,
"bilinear": cv2.INTER_LINEAR,
"nearest": cv2.INTER_NEAREST,
}
try:
return flags[mode]
except KeyError:
raise ValueError(f"Unsupported pixel upsample mode: {mode!r} "
f"(expected one of {', '.join(PIXEL_UPSAMPLE_MODES)})") from None
def unsharp(frame: np.ndarray, amount: float) -> np.ndarray:
"""Light unsharp mask, used to counter resampling / optical-flow softening.
Parameters:
frame (np.ndarray): HxWx3 uint8 RGB frame.
amount (float): Strength; ``0`` returns the frame unchanged.
Returns:
np.ndarray: A new frame; the input is never modified in place.
"""
if amount <= 0.0:
return frame
import cv2
blur = cv2.GaussianBlur(frame, (0, 0), 1.0)
return cv2.addWeighted(frame, 1.0 + amount, blur, -amount, 0)
def upsample_frame(
frame: np.ndarray,
*,
width: int,
height: int,
mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
sharpen: float = 0.0,
) -> np.ndarray:
"""
Resample one decoded RGB frame to the target pixel size.
Parameters:
frame (np.ndarray): HxWx3 uint8 RGB frame.
width (int): Target width in pixels.
height (int): Target height in pixels.
mode (str): Interpolation kernel, one of :data:`PIXEL_UPSAMPLE_MODES`.
sharpen (float): Unsharp strength applied after the resize.
Returns:
np.ndarray: A new frame at ``height x width``; already-correct sizes
are still passed through ``sharpen``.
Raises:
ValueError: If the frame is not HxWx3, or the target size is not positive.
"""
import cv2
array = np.asarray(frame)
if array.ndim != 3 or array.shape[2] != 3:
raise ValueError(f"frame must have shape HxWx3, got {array.shape}")
if width <= 0 or height <= 0:
raise ValueError(f"target size must be positive, got {width}x{height}")
if array.dtype != np.uint8:
array = np.clip(array, 0, 255).astype(np.uint8)
if (array.shape[0], array.shape[1]) != (height, width):
array = cv2.resize(array, (width, height), interpolation=_interpolation_flag(mode))
return unsharp(array, sharpen)
def upsample_frames(
frames: Iterable[np.ndarray],
*,
width: int,
height: int,
mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
sharpen: float = 0.0,
) -> list[np.ndarray]:
"""
Resample every decoded frame to the target pixel size.
Parameters:
frames (Iterable[np.ndarray]): Decoded HxWx3 uint8 RGB frames.
width (int): Target width in pixels.
height (int): Target height in pixels.
mode (str): Interpolation kernel, one of :data:`PIXEL_UPSAMPLE_MODES`.
sharpen (float): Unsharp strength applied after each resize.
Returns:
list[np.ndarray]: New frames at the target size, in input order.
"""
return [upsample_frame(frame, width=width, height=height, mode=mode, sharpen=sharpen) for frame in frames]
__all__ = [
"DEFAULT_PIXEL_UPSAMPLE_MODE",
"PIXEL_UPSAMPLE_MODES",
"unsharp",
"upsample_frame",
"upsample_frames",
]
-243
View File
@@ -1,243 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Memory-tier helpers for Apple Silicon MLX/MPS experiments.
macOS does not expose a perfect "pretend this machine only has 16 GB unified
memory" switch. MLX can cap the allocator used by the Apple-native DiT path,
and PyTorch MPS exposes process-level watermark environment variables for the
hybrid prompt/decode stages. Applying both gives benchmark and generation
entrypoints a practical, explicit way to exercise memory-tier presets.
"""
from __future__ import annotations
import argparse
import gc
import os
from dataclasses import dataclass, field
from typing import Any
GIB = 1024**3
@dataclass(frozen=True)
class AppliedMemoryLimits:
"""Memory limits applied for one Apple Silicon benchmark/generation process."""
mlx_memory_limit_gib: float | None = None
mlx_cache_limit_gib: float | None = None
mlx_disable_cache: bool = False
mlx_wired_limit_gib: float | None = None
torch_mps_high_watermark_ratio: float | None = None
torch_mps_low_watermark_ratio: float | None = None
applied_bytes: dict[str, int] = field(default_factory=dict)
previous_bytes: dict[str, int] = field(default_factory=dict)
errors: dict[str, str] = field(default_factory=dict)
def as_metrics(self) -> dict[str, int | float | str | bool | None]:
"""Flatten the configured memory limits, applied values, previous values, and errors into a metrics dictionary.
Returns:
dict[str, int | float | str | bool | None]: Metrics keyed by limit names and their corresponding values.
"""
metrics: dict[str, int | float | str | bool | None] = {
"mlx_memory_limit_gib": self.mlx_memory_limit_gib,
"mlx_cache_limit_gib": self.mlx_cache_limit_gib,
"mlx_disable_cache": self.mlx_disable_cache,
"mlx_wired_limit_gib": self.mlx_wired_limit_gib,
"torch_mps_high_watermark_ratio": self.torch_mps_high_watermark_ratio,
"torch_mps_low_watermark_ratio": self.torch_mps_low_watermark_ratio,
}
for name, value in self.applied_bytes.items():
metrics[f"{name}_bytes"] = value
for name, value in self.previous_bytes.items():
metrics[f"previous_{name}_bytes"] = value
for name, error in self.errors.items():
metrics[f"{name}_error"] = error
return metrics
def gib_to_bytes(value: float | None) -> int | None:
"""
Convert a positive memory limit from GiB to bytes.
Parameters:
value (float | None): Memory limit in GiB, or `None` when unset.
Returns:
int | None: The memory limit in bytes, or `None` when no limit is provided.
Raises:
ValueError: If `value` is zero or negative.
"""
if value is None:
return None
if value <= 0:
raise ValueError(f"Memory limit must be positive GiB, got {value}")
return int(value * GIB)
def cleanup_mlx(mx_module: Any | None = None) -> None:
"""Collect unreachable MLX objects, then release their allocator cache."""
if mx_module is None:
import mlx.core as mx
mx_module = mx
gc.collect()
mx_module.clear_cache()
def cleanup_torch_mps(torch_module: Any | None = None) -> None:
"""Collect unreachable Torch objects, then release the MPS allocator cache."""
if torch_module is None:
import torch
torch_module = torch
gc.collect()
if torch_module.backends.mps.is_available():
torch_module.mps.empty_cache()
def _set_mps_env(name: str, value: float | None) -> float | None:
"""Set a PyTorch MPS watermark environment variable.
Parameters:
name (str): Name of the environment variable to set.
value (float | None): Watermark ratio, or `None` to leave the variable unchanged.
Returns:
float | None: The configured watermark ratio, or `None` when no value is provided.
Raises:
ValueError: If `value` is negative.
"""
if value is None:
return None
if value < 0:
raise ValueError(f"{name} must be non-negative, got {value}")
os.environ[name] = str(value)
return value
def apply_memory_limits(
*,
mlx_memory_limit_gib: float | None = None,
mlx_cache_limit_gib: float | None = None,
mlx_disable_cache: bool = False,
mlx_wired_limit_gib: float | None = None,
torch_mps_high_watermark_ratio: float | None = None,
torch_mps_low_watermark_ratio: float | None = None,
mx_module: Any | None = None,
) -> AppliedMemoryLimits:
"""Apply optional MLX allocator limits and PyTorch MPS watermarks.
PyTorch reads MPS watermark variables when the MPS backend initializes, so
call this before importing PyTorch. Specifying only a high watermark sets the
low watermark to ``0.0``. MLX limit-setting failures are recorded in the
result and do not prevent other limits from being applied.
Parameters:
mlx_memory_limit_gib (float | None): Maximum MLX memory in GiB.
mlx_cache_limit_gib (float | None): Maximum MLX cache size in GiB.
mlx_disable_cache (bool): Whether to disable the MLX cache.
mlx_wired_limit_gib (float | None): Maximum MLX wired memory in GiB.
torch_mps_high_watermark_ratio (float | None): PyTorch MPS high watermark
ratio.
torch_mps_low_watermark_ratio (float | None): PyTorch MPS low watermark
ratio.
Returns:
AppliedMemoryLimits: Configured values, applied and previous MLX byte
limits, MPS watermark values, and per-limit errors.
"""
if torch_mps_high_watermark_ratio is not None and torch_mps_low_watermark_ratio is None:
torch_mps_low_watermark_ratio = 0.0
high = _set_mps_env("PYTORCH_MPS_HIGH_WATERMARK_RATIO", torch_mps_high_watermark_ratio)
low = _set_mps_env("PYTORCH_MPS_LOW_WATERMARK_RATIO", torch_mps_low_watermark_ratio)
memory_bytes = gib_to_bytes(mlx_memory_limit_gib)
cache_bytes = 0 if mlx_disable_cache else gib_to_bytes(mlx_cache_limit_gib)
wired_bytes = gib_to_bytes(mlx_wired_limit_gib)
applied: dict[str, int] = {}
previous: dict[str, int] = {}
errors: dict[str, str] = {}
if memory_bytes is not None or cache_bytes is not None or wired_bytes is not None:
if mx_module is None:
import mlx.core as mx
mx_module = mx
# Apply each limit independently; record failures without stopping.
limits = [
("mlx_memory_limit", memory_bytes, mx_module.set_memory_limit),
("mlx_cache_limit", cache_bytes, mx_module.set_cache_limit),
("mlx_wired_limit", wired_bytes, mx_module.set_wired_limit),
]
for name, value, setter in limits:
if value is not None:
try:
previous[name] = int(setter(value))
applied[name] = value
except Exception as exc: # noqa: BLE001 - macOS/system-limit dependent.
errors[name] = f"{type(exc).__name__}: {exc}"
return AppliedMemoryLimits(
mlx_memory_limit_gib=mlx_memory_limit_gib,
mlx_cache_limit_gib=mlx_cache_limit_gib,
mlx_disable_cache=mlx_disable_cache,
mlx_wired_limit_gib=mlx_wired_limit_gib,
torch_mps_high_watermark_ratio=high,
torch_mps_low_watermark_ratio=low,
applied_bytes=applied,
previous_bytes=previous,
errors=errors,
)
def add_memory_limit_args(
parser: argparse.ArgumentParser,
*,
mlx_memory_limit_gib: float | None = None,
mlx_cache_limit_gib: float | None = None,
mlx_disable_cache: bool = False,
mlx_wired_limit_gib: float | None = None,
torch_mps_high_watermark_ratio: float | None = None,
torch_mps_low_watermark_ratio: float | None = None,
) -> None:
"""
Add configurable Apple Silicon memory-limit options to an argument parser.
Parameters:
parser (argparse.ArgumentParser): Parser to which the options are added.
mlx_memory_limit_gib (float | None): Default MLX memory limit in GiB.
mlx_cache_limit_gib (float | None): Default MLX cache limit in GiB.
mlx_disable_cache (bool): Whether the cache limit defaults to zero.
mlx_wired_limit_gib (float | None): Default MLX wired-memory limit in GiB.
torch_mps_high_watermark_ratio (float | None): Default PyTorch MPS high-watermark ratio.
torch_mps_low_watermark_ratio (float | None): Default PyTorch MPS low-watermark ratio.
"""
parser.add_argument("--mlx-memory-limit-gib",
type=float,
default=mlx_memory_limit_gib,
help="Set MLX memory limit in GiB for memory-tier testing (DiT path).")
parser.add_argument("--mlx-cache-limit-gib",
type=float,
default=mlx_cache_limit_gib,
help="Set MLX cache limit in GiB. Use --mlx-disable-cache to force 0.")
parser.add_argument("--mlx-disable-cache",
action="store_true",
default=mlx_disable_cache,
help="Set MLX cache limit to 0 for stricter memory-tier tests.")
parser.add_argument("--mlx-wired-limit-gib",
type=float,
default=mlx_wired_limit_gib,
help="Set MLX wired-memory limit in GiB where supported by macOS/MLX.")
parser.add_argument("--torch-mps-high-watermark-ratio",
type=float,
default=torch_mps_high_watermark_ratio,
help="Set PYTORCH_MPS_HIGH_WATERMARK_RATIO before importing torch.")
parser.add_argument("--torch-mps-low-watermark-ratio",
type=float,
default=torch_mps_low_watermark_ratio,
help="Set PYTORCH_MPS_LOW_WATERMARK_RATIO before importing torch.")
-129
View File
@@ -1,129 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Best-effort prompt-embedding cache shared by the MLX entrypoints."""
from __future__ import annotations
import hashlib
import io
import json
import logging
import tempfile
from pathlib import Path
import numpy as np
logger = logging.getLogger(__name__)
def fingerprint_digest(fingerprint: dict[str, object]) -> str:
payload = json.dumps(fingerprint, sort_keys=True, separators=(",", ":"))
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def text_encoder_fingerprint(model_root: Path) -> dict[str, object]:
"""Return a cheap identity for the tokenizer and text-encoder files."""
root = model_root.resolve()
components = [path for name in ("tokenizer", "text_encoder") if (path := root / name).is_dir()]
scan_roots = components or [root]
files: list[list[object]] = []
complete = True
try:
for scan_root in scan_roots:
for path in sorted(scan_root.rglob("*")):
try:
if not path.is_file():
continue
stat = path.stat()
files.append([
path.relative_to(root).as_posix(),
stat.st_size,
stat.st_mtime_ns,
stat.st_ctime_ns,
])
except OSError:
complete = False
except OSError:
complete = False
# ponytail: metadata avoids hashing multi-GB weights; use a model manifest
# if supported workflows ever preserve size, mtime, and ctime while mutating.
return {"root": str(root), "files": files, "complete": complete}
def prompt_cache_meta_path(cache_path: Path) -> Path:
return cache_path.with_suffix(cache_path.suffix + ".json")
def _fingerprint_is_complete(fingerprint: dict[str, object]) -> bool:
text_encoder = fingerprint.get("text_encoder")
return not isinstance(text_encoder, dict) or text_encoder.get("complete") is not False
def load_prompt_cache(
cache_path: Path | None,
fingerprint: dict[str, object],
) -> np.ndarray | None:
"""Load a matching cache entry, treating every cache failure as a miss."""
if cache_path is None or not _fingerprint_is_complete(fingerprint):
return None
try:
metadata = json.loads(prompt_cache_meta_path(cache_path).read_text())
if not isinstance(metadata, dict):
return None
if metadata.get("fingerprint_sha256") != fingerprint_digest(fingerprint):
return None
payload = cache_path.read_bytes()
if metadata.get("data_sha256") != hashlib.sha256(payload).hexdigest():
return None
array = np.load(io.BytesIO(payload), allow_pickle=False)
return array if isinstance(array, np.ndarray) else None
except (EOFError, OSError, UnicodeError, ValueError) as exc:
logger.info("Prompt cache read skipped for %s: %s", cache_path, exc)
return None
def _atomic_write(path: Path, payload: bytes) -> None:
temp_path: Path | None = None
try:
with tempfile.NamedTemporaryFile(
mode="wb",
dir=path.parent,
prefix=f".{path.name}.",
suffix=".tmp",
delete=False,
) as handle:
temp_path = Path(handle.name)
handle.write(payload)
temp_path.replace(path)
finally:
if temp_path is not None:
temp_path.unlink(missing_ok=True)
def save_prompt_cache(
cache_path: Path | None,
embeds: np.ndarray,
fingerprint: dict[str, object],
) -> bool:
"""Atomically publish an integrity-bound cache entry when possible."""
if cache_path is None or not _fingerprint_is_complete(fingerprint):
return False
try:
cache_path.parent.mkdir(parents=True, exist_ok=True)
buffer = io.BytesIO()
np.save(buffer, np.asarray(embeds), allow_pickle=False)
payload = buffer.getvalue()
metadata = (json.dumps(
{
"fingerprint_sha256": fingerprint_digest(fingerprint),
"data_sha256": hashlib.sha256(payload).hexdigest(),
"fingerprint": fingerprint,
},
indent=2) + "\n").encode("utf-8")
# Publish data first. Until metadata follows, old metadata's data digest
# makes the torn pair a harmless miss rather than a stale cache hit.
_atomic_write(cache_path, payload)
_atomic_write(prompt_cache_meta_path(cache_path), metadata)
return True
except (OSError, TypeError, ValueError) as exc:
logger.info("Prompt cache write skipped for %s: %s", cache_path, exc)
return False
-454
View File
@@ -1,454 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Local prompt enrichment for the MLX Wan runtime (H3 Context-IR-style).
Wan's training captions are long and cinematic; short user prompts leave
quality on the table. This module expands a raw prompt into Wan-style
shot language **on device** — no remote API, no training.
Backends (first match wins):
1. **mlx-lm** — optional local LLM (``--enhance-prompt-model``).
2. **template** — deterministic cinematic expansion (always available).
System-prompt contract matches the streaming server's enhancer defaults
in ``fastvideo/entrypoints/streaming/prompt/enhancer.py`` so remote and
local paths stay interchangeable.
"""
from __future__ import annotations
import hashlib
import json
import re
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Keep in lockstep with streaming PromptEnhancer defaults (enhance op).
DEFAULT_ENHANCE_SYSTEM_PROMPT = ("You are a prompt enhancer for cinematic video generation. Given "
"a user prompt, produce an enhanced prompt that is more vivid, "
"specific, and concrete. Keep the subject intact; add lighting, "
"camera, and motion detail. Reply with just the enhanced prompt.")
# Small default that fits 16 GB Macs alongside the 1.3B DiT when the user
# opts into mlx-lm. Override with --enhance-prompt-model.
DEFAULT_MLX_LM_MODEL = "mlx-community/Qwen2.5-0.5B-Instruct-4bit"
_CAMERA_CUES = (
"cinematic",
"camera",
"lens",
"shot",
"bokeh",
"dolly",
"tracking",
"close-up",
"wide shot",
"handheld",
"steadicam",
)
_LIGHT_CUES = (
"light",
"lighting",
"sun",
"golden hour",
"neon",
"rim light",
"softbox",
"overcast",
"moonlight",
"volumetric",
)
_MOTION_CUES = (
"moving",
"motion",
"walk",
"run",
"flies",
"flying",
"drifts",
"sails",
"flows",
"pan",
"tilt",
"zoom",
)
@dataclass(frozen=True)
class EnhanceResult:
"""Outcome of a prompt enrichment call."""
original: str
enhanced: str
backend: str
elapsed_s: float
model: str | None = None
@property
def changed(self) -> bool:
"""Indicates whether the enhanced prompt differs from the original after trimming surrounding whitespace.
Returns:
bool: `True` if the prompts differ, `False` otherwise.
"""
return self.enhanced.strip() != self.original.strip()
def _normalize_user_prompt(prompt: str) -> str:
"""
Normalize a user prompt for enhancement.
Parameters:
prompt (str): User-provided prompt text.
Returns:
str: The prompt with leading and trailing whitespace removed and internal whitespace collapsed.
Raises:
ValueError: If the prompt is empty after whitespace normalization.
"""
text = " ".join(prompt.strip().split())
if not text:
raise ValueError("prompt must be non-empty")
return text
def _already_rich(prompt: str) -> bool:
"""
Determine whether a prompt already contains substantial camera and lighting detail.
Returns:
bool: `true` if the prompt is at least 160 characters long and includes camera and lighting cues, `false` otherwise.
"""
lower = prompt.lower()
has_camera = any(c in lower for c in _CAMERA_CUES)
has_light = any(c in lower for c in _LIGHT_CUES)
return len(prompt) >= 160 and has_camera and has_light
def enhance_prompt_template(prompt: str) -> str:
"""
Expand a prompt with cinematic camera, lighting, motion, and visual-quality details.
Rich prompts are preserved, while thinner prompts receive deterministic enhancements
without changing their subject.
Returns:
str: The original or expanded prompt with normalized whitespace and punctuation.
"""
text = _normalize_user_prompt(prompt)
if _already_rich(text):
return text
lower = text.lower()
parts = [text.rstrip(".")]
if not any(c in lower for c in _CAMERA_CUES):
parts.append("shot on a 35mm anamorphic lens, gentle handheld micro-movement, "
"shallow depth of field")
if not any(c in lower for c in _LIGHT_CUES):
parts.append("natural cinematic lighting with soft volumetric haze and subtle "
"rim light separating subject from background")
if not any(c in lower for c in _MOTION_CUES):
parts.append("smooth continuous motion with grounded physics")
parts.append("highly detailed, coherent temporal continuity, film grain, "
"color graded like a contemporary drama")
enhanced = ", ".join(parts)
# Single trailing period; collapse duplicate whitespace.
enhanced = re.sub(r"\s+", " ", enhanced).strip()
if not enhanced.endswith("."):
enhanced += "."
return enhanced
def enhance_prompt_mlx_lm(
prompt: str,
*,
model: str = DEFAULT_MLX_LM_MODEL,
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
max_tokens: int = 128,
temp: float = 0.6,
) -> str:
"""
Enhance a user prompt with a locally hosted mlx-lm instruction model.
Parameters:
prompt (str): The prompt to enhance.
model (str): The mlx-lm model identifier or path.
system_prompt (str): Instructions that guide prompt enhancement.
max_tokens (int): Maximum number of tokens to generate.
temp (float): Sampling temperature for generation.
Returns:
str: The enhanced prompt.
Raises:
RuntimeError: If mlx-lm is unavailable or produces an empty result.
"""
try:
from mlx_lm import generate, load
except ImportError as exc: # pragma: no cover - optional dep
raise RuntimeError("mlx-lm is not installed. `uv pip install mlx-lm` or use "
"--enhance-prompt-backend template.") from exc
text = _normalize_user_prompt(prompt)
logger.info("[MLX enhance] loading %s", model)
mlx_model, tokenizer = load(model)
messages = [
{
"role": "system",
"content": system_prompt
},
{
"role": "user",
"content": text
},
]
if hasattr(tokenizer, "apply_chat_template"):
chat = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
)
else: # pragma: no cover - ancient tokenizers
chat = f"{system_prompt}\n\nUser: {text}\nAssistant:"
raw = generate(
mlx_model,
tokenizer,
prompt=chat,
max_tokens=max_tokens,
temp=temp,
verbose=False,
)
enhanced = _clean_llm_output(raw, original=text)
if not enhanced:
raise RuntimeError("mlx-lm returned an empty enhance result")
return enhanced
def _clean_llm_output(raw: str, *, original: str) -> str:
"""
Clean generated prompt text and fall back to the original when the result is too short.
Parameters:
raw (str): Raw text produced by the language model.
original (str): Original prompt used as the fallback value.
Returns:
str: Cleaned first paragraph of the generated text, or the original prompt when the generated text is too short.
"""
text = raw.strip()
# Drop common prefatory phrases.
for prefix in (
"enhanced prompt:",
"here's the enhanced prompt:",
"here is the enhanced prompt:",
"sure:",
"sure,",
):
if text.lower().startswith(prefix):
text = text[len(prefix):].strip()
# Keep first non-empty paragraph only.
para = text.split("\n\n")[0].strip()
para = " ".join(para.split())
if len(para) < max(12, len(original) // 4):
return original
return para
def enhance_prompt(
prompt: str,
*,
backend: str = "auto",
model: str | None = None,
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
max_tokens: int = 128,
) -> EnhanceResult:
"""Enhance a prompt using the selected backend, falling back to a deterministic template when configured for automatic selection.
Parameters:
prompt (str): The prompt to enhance.
backend (str): The enhancement backend: ``"auto"``, ``"mlx-lm"``, or ``"template"``.
model (str | None): The MLX language model to use.
system_prompt (str): Instructions provided to the MLX language model.
max_tokens (int): Maximum number of tokens generated by the MLX language model.
Returns:
EnhanceResult: The original and enhanced prompts, selected backend, timing information, and model metadata.
Raises:
ValueError: If the prompt is empty or the backend is unsupported.
Exception: If the explicitly selected ``"mlx-lm"`` backend fails.
"""
text = _normalize_user_prompt(prompt)
backend_norm = (backend or "auto").lower()
if backend_norm not in {"auto", "mlx-lm", "template"}:
raise ValueError(f"Unknown enhance backend: {backend}")
start = time.perf_counter()
used_model: str | None = None
if backend_norm in {"auto", "mlx-lm"}:
try:
used_model = model or DEFAULT_MLX_LM_MODEL
enhanced = enhance_prompt_mlx_lm(
text,
model=used_model,
system_prompt=system_prompt,
max_tokens=max_tokens,
)
return EnhanceResult(
original=text,
enhanced=enhanced,
backend="mlx-lm",
elapsed_s=time.perf_counter() - start,
model=used_model,
)
except Exception as exc:
if backend_norm == "mlx-lm":
raise
logger.info(
"[MLX enhance] mlx-lm unavailable (%s); using template backend",
exc,
)
enhanced = enhance_prompt_template(text)
return EnhanceResult(
original=text,
enhanced=enhanced,
backend="template",
elapsed_s=time.perf_counter() - start,
model=None,
)
def enhance_cache_path(
prompt: str,
*,
backend: str,
model: str | None,
cache_dir: Path | None = None,
) -> Path:
"""Content-addressed cache file for an enhanced prompt string."""
root = cache_dir or (Path.home() / ".cache" / "fastvideo" / "enhanced_prompts")
key = "\0".join([prompt, backend, model or ""])
digest = hashlib.sha256(key.encode("utf-8")).hexdigest()[:24]
return root / f"{digest}.json"
def load_or_enhance_prompt(
prompt: str,
*,
backend: str = "auto",
model: str | None = None,
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
max_tokens: int = 128,
cache: bool = True,
cache_dir: Path | None = None,
) -> EnhanceResult:
"""
Enhance a prompt, reusing a cached result when available.
Parameters:
prompt (str): The prompt to enhance.
backend (str): Enhancement backend to use.
model (str | None): Optional model identifier.
system_prompt (str): System prompt for model-based enhancement.
max_tokens (int): Maximum number of tokens generated by the model.
cache (bool): Whether to read and write the on-disk cache.
cache_dir (Path | None): Optional directory for cached results.
Returns:
EnhanceResult: The enhanced prompt and backend metadata. Cached results are marked with the ``"cache"`` backend.
"""
text = _normalize_user_prompt(prompt)
path = enhance_cache_path(text, backend=backend, model=model, cache_dir=cache_dir)
if cache and path.is_file():
try:
payload = json.loads(path.read_text())
return EnhanceResult(
original=str(payload.get("original", text)),
enhanced=str(payload["enhanced"]),
# Mark cache hits explicitly so metrics/logs can distinguish
# a free replay from a fresh template/mlx-lm call.
backend="cache",
elapsed_s=0.0,
model=payload.get("model"),
)
except (OSError, KeyError, json.JSONDecodeError):
pass
result = enhance_prompt(
text,
backend=backend,
model=model,
system_prompt=system_prompt,
max_tokens=max_tokens,
)
if cache:
try:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(
{
"original": result.original,
"enhanced": result.enhanced,
"backend": result.backend,
"model": result.model,
},
indent=2,
))
except OSError as exc: # pragma: no cover - cache is best-effort
logger.info("[MLX enhance] cache write skipped: %s", exc)
return result
def enhance_result_as_metrics(result: EnhanceResult | None) -> dict[str, Any]:
"""
Convert prompt enhancement results into metrics fields.
Parameters:
result (EnhanceResult | None): The enhancement result, or `None` when no enhancement was performed.
Returns:
dict[str, Any]: A metrics mapping containing enhancement status, backend metadata, timing, and original and enhanced prompts.
"""
if result is None:
return {
"enhance_prompt": False,
"enhance_backend": None,
"enhance_model": None,
"enhance_elapsed_s": None,
"prompt_original": None,
"prompt_enhanced": None,
}
return {
"enhance_prompt": True,
"enhance_backend": result.backend,
"enhance_model": result.model,
"enhance_elapsed_s": result.elapsed_s,
"prompt_original": result.original,
"prompt_enhanced": result.enhanced,
}
__all__ = [
"DEFAULT_ENHANCE_SYSTEM_PROMPT",
"DEFAULT_MLX_LM_MODEL",
"EnhanceResult",
"enhance_cache_path",
"enhance_prompt",
"enhance_prompt_mlx_lm",
"enhance_prompt_template",
"enhance_result_as_metrics",
"load_or_enhance_prompt",
]
-284
View File
@@ -1,284 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MLX block-scaled quantization backends (affine INT8, MXFP8/4, NVFP4).
Isolated experiment module: probes which ``mx.quantize`` modes the installed
MLX build supports and exposes a thin wrapper around native quantized matmul.
Depends only on ``mlx.core`` and the standard library — do not import the rest
of FastVideo from here.
"""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from typing import Final
from collections.abc import Mapping
import mlx.core as mx
# Probe matrix side length: must be divisible by every mode's group size
# (affine g64, mxfp* g32, nvfp4 g16).
_PROBE_DIM: Final[int] = 64
class QuantBackend(str, Enum):
"""Named MLX quantization backends evaluated for M5 Neural Accelerators."""
AFFINE_INT8_G64 = "affine_int8_g64"
MXFP8 = "mxfp8"
MXFP4 = "mxfp4"
NVFP4 = "nvfp4"
BACKENDS: Final[tuple[str, ...]] = tuple(b.value for b in QuantBackend)
# Backend name -> kwargs for mx.quantize / mx.quantized_matmul.
# Affine baseline matches FastVideo DiT load path (INT8, group size 64).
# MX/NV block-scaled modes use MLX defaults (see mx.quantize docs).
_BACKEND_KWARGS: Final[Mapping[str, Mapping[str, object]]] = {
QuantBackend.AFFINE_INT8_G64.value: {
"mode": "affine",
"bits": 8,
"group_size": 64,
},
QuantBackend.MXFP8.value: {
"mode": "mxfp8",
"bits": None,
"group_size": None,
},
QuantBackend.MXFP4.value: {
"mode": "mxfp4",
"bits": None,
"group_size": None,
},
QuantBackend.NVFP4.value: {
"mode": "nvfp4",
"bits": None,
"group_size": None,
},
}
_SUPPORT_CACHE: dict[str, bool] = {}
_SUPPORT_ERROR_CACHE: dict[str, str | None] = {}
_BYTES_CACHE: dict[str, float] = {}
@dataclass(frozen=True)
class QuantizedWeight:
"""Packed quantized weight plus scales/biases for one backend."""
weight: mx.array
scales: mx.array
biases: mx.array | None
backend: str
mode: str
bits: int | None
group_size: int | None
# Original (rows, cols) of the fp weight, used for bytes-per-element.
orig_shape: tuple[int, int]
def _normalize_backend(backend: str) -> str:
"""Normalize a quantization backend name and validate that it is supported.
Parameters:
backend (str): Backend name to normalize.
Returns:
str: The lowercase backend name without surrounding whitespace.
Raises:
ValueError: If the backend name is unknown.
"""
name = backend.strip().lower()
if name not in _BACKEND_KWARGS:
known = ", ".join(BACKENDS)
raise ValueError(f"Unknown quant backend {backend!r}. Expected one of: {known}")
return name
def _kwargs_for(backend: str) -> dict[str, object]:
"""Return the MLX quantization arguments configured for a backend.
Parameters:
backend (str): Backend name to resolve.
Returns:
dict[str, object]: Quantization arguments for the normalized backend.
"""
return dict(_BACKEND_KWARGS[_normalize_backend(backend)])
def support_error(backend: str) -> str | None:
"""
Check whether a quantization backend is supported by the current MLX runtime.
Parameters:
backend (str): Quantization backend name.
Returns:
str | None: An error description when the backend is unsupported, or `None` when supported.
"""
name = _normalize_backend(backend)
if name in _SUPPORT_ERROR_CACHE:
return _SUPPORT_ERROR_CACHE[name]
kwargs = _kwargs_for(name)
try:
w = mx.zeros((_PROBE_DIM, _PROBE_DIM), dtype=mx.float16)
quantized = mx.quantize(
w,
group_size=kwargs["group_size"], # type: ignore[arg-type]
bits=kwargs["bits"], # type: ignore[arg-type]
mode=str(kwargs["mode"]),
)
w_q = quantized[0]
scales = quantized[1]
biases = quantized[2] if len(quantized) == 3 else None
x = mx.zeros((1, _PROBE_DIM), dtype=mx.float16)
y = mx.quantized_matmul(
x,
w_q,
scales,
biases,
transpose=True,
group_size=kwargs["group_size"], # type: ignore[arg-type]
bits=kwargs["bits"], # type: ignore[arg-type]
mode=str(kwargs["mode"]),
)
mx.eval(y)
_SUPPORT_ERROR_CACHE[name] = None
_SUPPORT_CACHE[name] = True
except Exception as exc: # noqa: BLE001 - MLX raises varied types per mode/version.
msg = f"{type(exc).__name__}: {exc}"
_SUPPORT_ERROR_CACHE[name] = msg
_SUPPORT_CACHE[name] = False
return _SUPPORT_ERROR_CACHE[name]
def is_supported(backend: str) -> bool:
"""Return True if the installed MLX build can quantize/matmul with ``backend``."""
name = _normalize_backend(backend)
if name not in _SUPPORT_CACHE:
support_error(name)
return _SUPPORT_CACHE[name]
def quantize_weight(w: mx.array, backend: str) -> QuantizedWeight:
"""
Quantize a two-dimensional weight matrix using the specified native MLX backend.
Parameters:
w (mx.array): The two-dimensional weight matrix to quantize.
backend (str): The quantization backend to use.
Returns:
QuantizedWeight: The quantized weights and associated quantization metadata.
Raises:
ValueError: If the backend is unknown, the weight is not two-dimensional,
or its last dimension is not divisible by the backend's group size.
RuntimeError: If the backend is unsupported by the installed MLX build.
"""
name = _normalize_backend(backend)
err = support_error(name)
if err is not None:
mlx_version = getattr(mx, "__version__", "unknown")
raise RuntimeError(f"Quant backend {name!r} is not supported by installed mlx "
f"({mlx_version}): {err}")
if w.ndim != 2:
raise ValueError(f"quantize_weight expects a 2D weight, got shape {tuple(w.shape)}")
rows, cols = int(w.shape[0]), int(w.shape[1])
kwargs = _kwargs_for(name)
group_size = kwargs["group_size"]
# When group_size is None, MLX applies the mode default; only check when set.
if isinstance(group_size, int) and cols % group_size != 0:
raise ValueError(f"Weight last dim {cols} must be divisible by group_size={group_size} "
f"for backend {name!r}")
quantized = mx.quantize(
w,
group_size=kwargs["group_size"], # type: ignore[arg-type]
bits=kwargs["bits"], # type: ignore[arg-type]
mode=str(kwargs["mode"]),
)
w_q = quantized[0]
scales = quantized[1]
biases = quantized[2] if len(quantized) == 3 else None
eval_args = [w_q, scales] if biases is None else [w_q, scales, biases]
mx.eval(*eval_args)
return QuantizedWeight(
weight=w_q,
scales=scales,
biases=biases,
backend=name,
mode=str(kwargs["mode"]),
bits=kwargs["bits"] if isinstance(kwargs["bits"], int) else None,
group_size=group_size if isinstance(group_size, int) else None,
orig_shape=(rows, cols),
)
def quantized_matmul(x: mx.array, qw: QuantizedWeight) -> mx.array:
"""Compute ``x @ w.T`` in the quantized domain via ``mx.quantized_matmul``."""
return mx.quantized_matmul(
x,
qw.weight,
qw.scales,
qw.biases,
transpose=True,
group_size=qw.group_size,
bits=qw.bits,
mode=qw.mode,
)
def _artifact_nbytes(qw: QuantizedWeight) -> int:
"""
Calculate the total storage size of a quantized weight artifact in bytes.
Parameters:
qw (QuantizedWeight): Quantized weight artifact whose packed weights, scales, and optional biases are measured.
Returns:
int: Total number of bytes used by the artifact's stored arrays.
"""
total = int(qw.weight.nbytes) + int(qw.scales.nbytes)
if qw.biases is not None:
total += int(qw.biases.nbytes)
return total
def bytes_per_weight(backend: str) -> float:
"""
Measure the effective storage cost of a quantized weight.
Parameters:
backend (str): Quantization backend to measure.
Returns:
float: Stored bytes per original weight element, including packed weights,
scales, and optional biases.
Raises:
RuntimeError: If the backend is unsupported.
"""
name = _normalize_backend(backend)
if name in _BYTES_CACHE:
return _BYTES_CACHE[name]
err = support_error(name)
if err is not None:
mlx_version = getattr(mx, "__version__", "unknown")
raise RuntimeError(f"Cannot measure bytes_per_weight for unsupported backend {name!r} "
f"(mlx {mlx_version}): {err}")
probe = mx.zeros((_PROBE_DIM, _PROBE_DIM), dtype=mx.float16)
qw = quantize_weight(probe, name)
n_elem = qw.orig_shape[0] * qw.orig_shape[1]
value = _artifact_nbytes(qw) / float(n_elem)
_BYTES_CACHE[name] = value
return value
-690
View File
@@ -1,690 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Two-pass spatial refine for the MLX Wan runtime (H3 / LTX-2 pattern).
Biggest quality lever on Apple Silicon without a new model or training:
generate at base resolution, then run a second denoising pass with the
*same* DiT at a higher resolution.
This is the MLX-side port of the CUDA refine template in
``fastvideo/pipelines/basic/ltx2/stages/ltx2_refine.py`` and the H3
"base + regenerate" pattern documented in
``docs/design/mac_qad_two_product_strategy.md``:
1. :func:`plan_refine_resolutions` — split the request into stage-1
(base) and stage-2 (target) pixel sizes, validating VAE / patch
alignment the way :class:`LTX2RefineInitStage` does.
2. :func:`upsample_latents_spatial` — 2× (or N×) spatial upsample of
clean latents. Wan has no learned latent upsampler on Mac, so this
is bilinear over the H×W plane (temporal axis untouched) — same
role as LTX-2's ``upsample_video`` hand-off, without the learned
residual.
3. :func:`prepare_refine_latents` — upsample + re-noise the clean
stage-1 latents to the stage-2 sigma so the second denoise has
something to refine (mirrors :class:`LTX2UpsampleStage` +
``apply_ltx2_gaussian_noiser``).
4. :func:`run_two_pass_dmd` — orchestrate stage-1 denoise → refine
hand-off → stage-2 denoise with the same model / prompt embeds.
No LoRA swap, no dedicated SR weights, no new training — pure pipeline
work reusable by Wan2.1-14B and Wan2.2-5B on Apple Silicon.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from collections.abc import Callable, Sequence
import numpy as np
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, add_noise, dmd_step
if TYPE_CHECKING: # pragma: no cover - typing only
import mlx.core as mx
logger = init_logger(__name__)
# Default stage-2 noise level when the caller does not supply a schedule.
# Matches the first entry of LTX-2's STAGE_2_DISTILLED_SIGMA_VALUES in spirit
# (start the refine denoise from a high-noise level) without hard-wiring the
# LTX-2 distilled grid onto Wan's flow-match schedule.
DEFAULT_REFINE_SIGMA = 0.909375
@dataclass(frozen=True)
class RefinePlan:
"""Resolved stage-1 / stage-2 geometry for a two-pass refine run."""
target_height: int
target_width: int
stage1_height: int
stage1_width: int
spatial_scale: int
vae_spatial_compression: int
vae_temporal_compression: int
num_frames: int
@property
def stage1_latent_height(self) -> int:
"""Return the stage-1 latent height after VAE spatial compression."""
return self.stage1_height // self.vae_spatial_compression
@property
def stage1_latent_width(self) -> int:
"""Return the stage-one latent width after VAE spatial compression."""
return self.stage1_width // self.vae_spatial_compression
@property
def stage2_latent_height(self) -> int:
"""Calculate the target-resolution latent height.
Returns:
int: The target height divided by the VAE spatial compression factor.
"""
return self.target_height // self.vae_spatial_compression
@property
def stage2_latent_width(self) -> int:
"""Return the target image width in latent-space units."""
return self.target_width // self.vae_spatial_compression
@property
def latent_frames(self) -> int:
"""Calculate the number of latent frames after VAE temporal compression.
Returns:
int: The compressed latent frame count.
"""
return (self.num_frames - 1) // self.vae_temporal_compression + 1
def plan_refine_resolutions(
*,
height: int,
width: int,
num_frames: int,
spatial_scale: int = 2,
vae_spatial_compression: int = 8,
vae_temporal_compression: int = 4,
patch_size: tuple[int, int, int] = (1, 2, 2),
enabled: bool = True,
mode_label: str = "Refine",
) -> RefinePlan:
"""
Validate the requested dimensions and create the stage-1 and target-resolution refinement plan.
Parameters:
height (int): Target image height in pixels.
width (int): Target image width in pixels.
num_frames (int): Number of frames in the input sequence.
spatial_scale (int): Factor used to reduce spatial dimensions for stage 1.
vae_spatial_compression (int): Spatial compression factor of the VAE.
vae_temporal_compression (int): Temporal compression factor of the VAE.
patch_size (tuple[int, int, int]): Temporal and spatial patch dimensions used to validate latent-grid alignment.
enabled (bool): Whether to use two-pass refinement.
mode_label (str): Name of the calling mode, used to prefix validation
errors so ``--fast-spatial`` failures do not read as refine failures.
Returns:
RefinePlan: The validated stage-1 and target-resolution plan.
"""
if height <= 0 or width <= 0:
raise ValueError(f"height/width must be positive, got {height}x{width}")
if spatial_scale < 1:
raise ValueError(f"spatial_scale must be >= 1, got {spatial_scale}")
if num_frames <= 0:
raise ValueError(f"num_frames must be positive, got {num_frames}")
if vae_spatial_compression < 1 or vae_temporal_compression < 1:
raise ValueError("VAE compression factors must be positive")
if height % vae_spatial_compression != 0 or width % vae_spatial_compression != 0:
raise ValueError(f"height/width must be divisible by vae_spatial_compression={vae_spatial_compression} "
f"(got {height}x{width}).")
if (num_frames - 1) % vae_temporal_compression != 0:
raise ValueError(f"num_frames must be 1 modulo vae_temporal_compression={vae_temporal_compression} "
f"(got {num_frames}).")
if not enabled or spatial_scale == 1:
plan = RefinePlan(
target_height=height,
target_width=width,
stage1_height=height,
stage1_width=width,
spatial_scale=1,
vae_spatial_compression=vae_spatial_compression,
vae_temporal_compression=vae_temporal_compression,
num_frames=num_frames,
)
_validate_plan(plan, patch_size=patch_size, mode_label=mode_label)
return plan
if height % spatial_scale != 0 or width % spatial_scale != 0:
raise ValueError(f"{mode_label} requires height/width divisible by spatial_scale={spatial_scale} "
f"(got {height}x{width}).")
stage1_height = height // spatial_scale
stage1_width = width // spatial_scale
# Stage-1 must land on a VAE-aligned grid so the first denoise produces
# valid latents; the LTX-2 init stage enforces the same constraint.
if (stage1_height % vae_spatial_compression != 0 or stage1_width % vae_spatial_compression != 0):
raise ValueError(f"{mode_label} requires height/width divisible by "
f"{spatial_scale * vae_spatial_compression} "
f"(got {height}x{width}, vae_spatial={vae_spatial_compression}).")
plan = RefinePlan(
target_height=height,
target_width=width,
stage1_height=stage1_height,
stage1_width=stage1_width,
spatial_scale=spatial_scale,
vae_spatial_compression=vae_spatial_compression,
vae_temporal_compression=vae_temporal_compression,
num_frames=num_frames,
)
_validate_plan(plan, patch_size=patch_size, mode_label=mode_label)
logger.info(
"[MLX refine] enabled: stage1=%dx%d stage2=%dx%d scale=%dx",
stage1_width,
stage1_height,
width,
height,
spatial_scale,
)
return plan
def _validate_plan(plan: RefinePlan, *, patch_size: tuple[int, int, int], mode_label: str = "Refine") -> None:
"""
Validate that both refinement stages have latent dimensions aligned to the patch grid.
Parameters:
patch_size (tuple[int, int, int]): Temporal, height, and width patch dimensions.
Raises:
ValueError: If a stage's spatial latent dimensions or the temporal latent
dimension is not divisible by the corresponding patch dimension.
"""
pt, ph, pw = patch_size
for label, lh, lw in (
("stage1", plan.stage1_latent_height, plan.stage1_latent_width),
("stage2", plan.stage2_latent_height, plan.stage2_latent_width),
):
if lh % ph != 0 or lw % pw != 0:
raise ValueError(f"{mode_label} {label} latent grid {lh}x{lw} is not divisible by "
f"patch spatial size {ph}x{pw}.")
if plan.latent_frames % pt != 0:
raise ValueError(f"{mode_label} latent_frames={plan.latent_frames} is not divisible by "
f"patch temporal size {pt}.")
def upsample_latents_spatial(
latents: Any,
*,
scale: int = 2,
mode: str = "bilinear",
) -> Any:
"""
Upsample the spatial dimensions of 5-D latent arrays while preserving the batch, channel, and temporal dimensions.
Parameters:
latents (Any): Latents with shape ``(B, C, T, H, W)``.
scale (int): Integer factor for enlarging the spatial dimensions.
mode (str): Interpolation mode, either ``"nearest"`` or ``"bilinear"``.
Returns:
Any: Latents with shape ``(B, C, T, H * scale, W * scale)``.
"""
if scale < 1:
raise ValueError(f"scale must be >= 1, got {scale}")
if scale == 1:
return latents
# Accept both mx.array and np.ndarray so unit tests can run without MLX.
is_mlx = hasattr(latents, "dtype") and type(latents).__module__.startswith("mlx")
if is_mlx:
return _upsample_latents_mlx(latents, scale=scale, mode=mode)
return _upsample_latents_numpy(np.asarray(latents), scale=scale, mode=mode)
def _upsample_latents_numpy(
latents: np.ndarray,
*,
scale: int,
mode: str,
) -> np.ndarray:
"""Upsample 5-D latent arrays spatially using nearest-neighbor or bilinear interpolation.
Parameters:
latents (np.ndarray): Latents with shape ``(B, C, T, H, W)``.
scale (int): Spatial upsampling factor.
mode (str): Interpolation mode, either ``"nearest"`` or ``"bilinear"``.
Returns:
np.ndarray: Spatially upsampled latents with preserved batch, channel, and temporal dimensions.
Raises:
ValueError: If the latents are not five-dimensional or the interpolation mode is unsupported.
"""
if latents.ndim != 5:
raise ValueError(f"Expected 5-D latents (B,C,T,H,W), got shape {latents.shape}")
b, c, t, h, w = latents.shape
if mode == "nearest":
# (B,C,T,H,1,W,1) -> broadcast to (B,C,T,H,scale,W,scale) -> merge.
out = np.repeat(np.repeat(latents, scale, axis=3), scale, axis=4)
return out
if mode != "bilinear":
raise ValueError(f"Unsupported upsample mode: {mode}")
# Bilinear over the spatial plane. Flatten (B,C,T) into a batch of 2-D
# maps so a single vectorized gather covers every frame/channel.
src = latents.reshape(b * c * t, h, w).astype(np.float32, copy=False)
out_h, out_w = h * scale, w * scale
# Map output pixel centers onto the input grid (align_corners=False).
ys = (np.arange(out_h, dtype=np.float32) + 0.5) * (h / out_h) - 0.5
xs = (np.arange(out_w, dtype=np.float32) + 0.5) * (w / out_w) - 0.5
ys = np.clip(ys, 0.0, h - 1.0)
xs = np.clip(xs, 0.0, w - 1.0)
y0 = np.floor(ys).astype(np.int64)
x0 = np.floor(xs).astype(np.int64)
y1 = np.minimum(y0 + 1, h - 1)
x1 = np.minimum(x0 + 1, w - 1)
wy = (ys - y0.astype(np.float32))[:, None]
wx = (xs - x0.astype(np.float32))[None, :]
# Gather the four corners: shape (N, out_h, out_w).
Ia = src[:, y0[:, None], x0[None, :]]
Ib = src[:, y0[:, None], x1[None, :]]
Ic = src[:, y1[:, None], x0[None, :]]
Id = src[:, y1[:, None], x1[None, :]]
wa = (1.0 - wy) * (1.0 - wx)
wb = (1.0 - wy) * wx
wc = wy * (1.0 - wx)
wd = wy * wx
out = wa * Ia + wb * Ib + wc * Ic + wd * Id
return out.reshape(b, c, t, out_h, out_w).astype(latents.dtype, copy=False)
def _upsample_latents_mlx(latents: mx.array, *, scale: int, mode: str) -> mx.array:
"""
Upsample MLX latent tensors along their spatial dimensions.
Parameters:
latents (mx.array): A latent tensor with shape `(B, C, T, H, W)`.
scale (int): The integer spatial upsampling factor.
mode (str): The interpolation mode, such as `"nearest"` or `"bilinear"`.
Returns:
mx.array: The spatially upsampled latent tensor with its original data type.
"""
import mlx.core as mx
# Route through NumPy for the interpolation math. Latent tensors at Mac
# resolutions are small (e.g. 1×16×21×30×52 ≈ 1 MB) so the host hop is
# cheaper than carrying a bespoke Metal bilinear kernel, and it keeps
# the CPU-only unit tests and the MLX path on one implementation.
np_latents = np.array(latents.astype(mx.float32))
up = _upsample_latents_numpy(np_latents, scale=scale, mode=mode)
return mx.array(up).astype(latents.dtype)
def prepare_refine_latents(
clean_latents: Any,
*,
scale: int = 2,
sigma: float = DEFAULT_REFINE_SIGMA,
noise: Any | None = None,
add_noise_flag: bool = True,
upsample_mode: str = "bilinear",
seed: int | None = None,
) -> Any:
"""
Upsample clean latents spatially and optionally mix them with Gaussian noise.
Parameters:
clean_latents: The stage-1 latent tensor.
sigma: Noise mixing factor between 0 and 1.
noise: Optional noise tensor to mix with the upsampled latents.
add_noise_flag: Whether to apply noise mixing.
upsample_mode: Spatial interpolation mode.
seed: Optional seed for generated noise.
Returns:
The upsampled latents, optionally mixed with noise.
Raises:
ValueError: If sigma is outside the range from 0 to 1.
"""
if sigma < 0.0 or sigma > 1.0:
raise ValueError(f"sigma must be in [0, 1], got {sigma}")
upsampled = upsample_latents_spatial(clean_latents, scale=scale, mode=upsample_mode)
if not add_noise_flag or sigma == 0.0:
return upsampled
is_mlx = hasattr(upsampled, "dtype") and type(upsampled).__module__.startswith("mlx")
if noise is None:
noise = _draw_noise_like(upsampled, seed=seed, is_mlx=is_mlx)
return add_noise(upsampled, noise, float(sigma))
def refine_sigma_from_schedule(
schedule: MLXDMDSchedule,
timesteps: Sequence[float | int],
) -> float:
"""Derive the refinement noise level from the first refinement timestep.
Parameters:
schedule (MLXDMDSchedule): Schedule used to map timesteps to noise levels.
timesteps (Sequence[float | int]): Refinement timesteps, whose first value determines the sigma.
Returns:
float: Sigma corresponding to the first refinement timestep.
Raises:
ValueError: If `timesteps` is empty.
"""
if not timesteps:
raise ValueError("timesteps must be non-empty to derive a refine sigma")
return float(schedule.sigma_for(float(timesteps[0])))
def default_refine_timesteps(
schedule: MLXDMDSchedule,
timesteps: Sequence[float | int],
) -> list[float]:
"""Derive stage-2 timesteps from the stage-1 DMD grid.
The stage-2 pass must start *below* full noise, otherwise the hand-off
``(1 - sigma) * upsampled + sigma * noise`` weights stage 1 at zero and
the refine pass silently becomes a plain full-resolution generation at
twice the cost. FastWan's stage-1 grid opens at ``t=1000`` (``sigma``
exactly 1.0), so reusing it verbatim — which is what happens when
``--refine-dmd-denoising-steps`` is left unset — discards stage 1.
Dropping the leading full-noise entries keeps the pass on timesteps the
distilled student was actually trained on (no off-grid ``t`` the DiT has
never seen) while letting the stage-1 structure through.
Parameters:
schedule (MLXDMDSchedule): Schedule used to map timesteps to noise levels.
timesteps (Sequence[float | int]): The stage-1 DMD timestep grid.
Returns:
list[float]: The stage-1 grid with leading full-noise timesteps removed.
Raises:
ValueError: If every timestep in the grid is at full noise, leaving no
usable refine step.
"""
steps = [float(step) for step in timesteps]
first = 0
while first < len(steps) and schedule.sigma_for(steps[first]) >= 1.0:
first += 1
if first == len(steps):
raise ValueError(f"No usable refine timesteps in {steps}: every entry is at sigma >= 1 "
"(full noise), which would discard the stage-1 result. Pass "
"explicit stage-2 timesteps below the full-noise step.")
return steps[first:]
def run_dmd_loop(
*,
dit: Any,
latents: Any,
encoder_hidden_states: Any,
freqs_cis: tuple[Any, Any],
timesteps: Sequence[float | int],
schedule: MLXDMDSchedule,
mx_dtype: Any,
seed: int | None = None,
step_callback: Callable[[int, int], None] | None = None,
label: str = "denoise",
) -> Any:
"""
Denoise latents over the supplied timesteps using the DMD schedule.
Parameters:
timesteps (Sequence[float | int]): Denoising timesteps in execution order.
seed (int | None): Seed for reproducible intermediate noise generation.
step_callback (Callable[[int, int], None] | None): Callback receiving the
completed step number and total step count.
label (str): Label used for progress output when no callback is provided.
Returns:
Any: The denoised latents.
"""
import mlx.core as mx
renoise_rng = np.random.default_rng(seed) if seed is not None else None
latents_out = latents
n_steps = len(timesteps)
for step_index, timestep in enumerate(timesteps):
noise_input = latents_out
ts_val = float(timestep)
timestep_mx = mx.array([ts_val]).astype(mx.float32)
noise_pred = dit(
latents_out.astype(mx_dtype),
encoder_hidden_states,
timestep_mx,
freqs_cis,
)
noise_input_f32 = noise_input.astype(mx.float32)
pred_noise_f32 = noise_pred.astype(mx.float32)
if step_index < n_steps - 1:
next_ts: float | None = float(timesteps[step_index + 1])
if renoise_rng is not None:
renoise = mx.array(renoise_rng.standard_normal(tuple(noise_input_f32.shape)).astype(np.float32))
else:
renoise = mx.random.normal(noise_input_f32.shape).astype(mx.float32)
else:
next_ts, renoise = None, None
latents_out = dmd_step(
latents=noise_input_f32,
noise_input_latent=noise_input_f32,
pred_noise=pred_noise_f32,
schedule=schedule,
timestep=ts_val,
next_timestep=next_ts,
noise=renoise,
).astype(mx_dtype)
mx.eval(latents_out)
if step_callback is not None:
step_callback(step_index + 1, n_steps)
else:
print(f"{label} step {step_index + 1}/{n_steps} complete")
return latents_out
@dataclass(frozen=True)
class TwoPassResult:
"""Outputs of :func:`run_two_pass_dmd`."""
latents: Any
stage1_latents: Any
plan: RefinePlan
refine_sigma: float
def run_two_pass_dmd(
*,
dit: Any,
encoder_hidden_states: Any,
noise_latents_stage1: Any,
freqs_cis_stage1: tuple[Any, Any],
freqs_cis_stage2: tuple[Any, Any] | None,
plan: RefinePlan,
schedule: MLXDMDSchedule,
timesteps: Sequence[float | int],
refine_timesteps: Sequence[float | int] | None = None,
mx_dtype: Any,
seed: int = 0,
add_noise_flag: bool = True,
upsample_mode: str = "bilinear",
refine_sigma: float | None = None,
step_callback: Callable[[str, int, int], None] | None = None,
) -> TwoPassResult:
"""
Run base denoising and, when enabled, spatial refinement denoising.
Parameters:
dit: DiT callable used for both denoising passes.
encoder_hidden_states: Prompt embeddings shared across both passes.
noise_latents_stage1: Initial stage-1 noise latents.
freqs_cis_stage1: RoPE tables for the stage-1 resolution.
freqs_cis_stage2: RoPE tables for the stage-2 resolution, required when refinement is enabled.
plan: Refinement geometry and configuration.
schedule: Flow-matching schedule used by both passes.
timesteps: Stage-1 denoising timesteps.
refine_timesteps: Stage-2 denoising timesteps. Uses `timesteps` when omitted.
mx_dtype: MLX dtype used for DiT inputs and outputs.
seed: Base seed for reproducible noise generation.
add_noise_flag: Whether to add noise to the upsampled stage-1 latents.
upsample_mode: Spatial upsampling mode, either `"bilinear"` or `"nearest"`.
refine_sigma: Stage-2 starting noise level. Derived from the first refinement timestep when omitted.
step_callback: Optional callback receiving the phase name, step index, and total step count.
Returns:
TwoPassResult containing the final latents, stage-1 latents, refinement plan, and applied refinement sigma.
Raises:
ValueError: If refinement is enabled without stage-2 RoPE tables, without refinement timesteps, or if upsampled latents do not match the planned stage-2 dimensions.
"""
stage1_cb = None
stage2_cb = None
if step_callback is not None:
stage1_cb = lambda i, n: step_callback("stage1", i, n) # noqa: E731
stage2_cb = lambda i, n: step_callback("stage2", i, n) # noqa: E731
stage1_latents = run_dmd_loop(
dit=dit,
latents=noise_latents_stage1,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis_stage1,
timesteps=timesteps,
schedule=schedule,
mx_dtype=mx_dtype,
seed=seed,
step_callback=stage1_cb,
label="stage1 denoise",
)
if plan.spatial_scale == 1:
return TwoPassResult(
latents=stage1_latents,
stage1_latents=stage1_latents,
plan=plan,
refine_sigma=0.0,
)
if freqs_cis_stage2 is None:
raise ValueError("freqs_cis_stage2 is required when refine spatial_scale > 1")
if refine_timesteps is not None:
stage2_timesteps = [float(step) for step in refine_timesteps]
if not stage2_timesteps:
raise ValueError("refine_timesteps must be non-empty when refine is enabled")
else:
# Not `list(timesteps)`: the stage-1 grid opens at full noise, which
# would weight the stage-1 result at zero. See default_refine_timesteps.
stage2_timesteps = default_refine_timesteps(schedule, timesteps)
grid_sigma = refine_sigma_from_schedule(schedule, stage2_timesteps)
sigma = float(refine_sigma) if refine_sigma is not None else grid_sigma
if refine_sigma is not None and abs(sigma - grid_sigma) > 1e-6:
# The loop tells the DiT `stage2_timesteps[0]`, which implies grid_sigma.
# Overriding the hand-off noise level breaks that correspondence, so the
# model is denoising from a level it was not told about. Useful for
# exploring schedules that bottom out too high, but say so out loud.
logger.warning(
"[MLX refine] refine_sigma=%.4f overrides the schedule's %.4f for timestep %g; "
"the DiT is told a timestep that no longer matches the noise it receives.",
sigma,
grid_sigma,
stage2_timesteps[0],
)
# A hand-off at sigma >= 1 is `0 * upsampled + 1 * noise`: stage 1 is
# thrown away and refine degrades to a plain full-res run at 2x the cost.
# Fail loudly rather than silently burning the first pass.
if add_noise_flag and sigma >= 1.0:
raise ValueError(f"Refine hand-off sigma={sigma:.4f} (from stage-2 timestep "
f"{stage2_timesteps[0]:g}) discards the stage-1 result entirely: "
"the upsampled latents are weighted (1 - sigma) = 0. Start the "
"stage-2 grid below the full-noise timestep, or pass "
"add_noise_flag=False to hand off the clean upsample.")
stage2_input = prepare_refine_latents(
stage1_latents,
scale=plan.spatial_scale,
sigma=sigma,
add_noise_flag=add_noise_flag,
upsample_mode=upsample_mode,
seed=seed + 1,
)
# Shape guard: upsampled latents must match the stage-2 RoPE grid.
expected_h = plan.stage2_latent_height
expected_w = plan.stage2_latent_width
got_h, got_w = int(stage2_input.shape[-2]), int(stage2_input.shape[-1])
if got_h != expected_h or got_w != expected_w:
raise ValueError(f"Refine upsample produced {got_h}x{got_w} latents, expected "
f"{expected_h}x{expected_w} for target "
f"{plan.target_height}x{plan.target_width}.")
logger.info(
"[MLX refine] stage2 start: latent=%dx%d sigma=%.4f steps=%d",
expected_w,
expected_h,
sigma,
len(stage2_timesteps),
)
stage2_latents = run_dmd_loop(
dit=dit,
latents=stage2_input,
encoder_hidden_states=encoder_hidden_states,
freqs_cis=freqs_cis_stage2,
timesteps=stage2_timesteps,
schedule=schedule,
mx_dtype=mx_dtype,
seed=seed + 2,
step_callback=stage2_cb,
label="stage2 refine",
)
return TwoPassResult(
latents=stage2_latents,
stage1_latents=stage1_latents,
plan=plan,
refine_sigma=sigma,
)
__all__ = [
"DEFAULT_REFINE_SIGMA",
"RefinePlan",
"TwoPassResult",
"default_refine_timesteps",
"plan_refine_resolutions",
"prepare_refine_latents",
"refine_sigma_from_schedule",
"run_dmd_loop",
"run_two_pass_dmd",
"upsample_latents_spatial",
]
def _draw_noise_like(like: Any, *, seed: int | None, is_mlx: bool) -> Any:
"""Generate Gaussian noise with the shape and array type of the input."""
shape = tuple(int(s) for s in like.shape)
if seed is not None:
rng = np.random.default_rng(seed)
noise_np = rng.standard_normal(shape).astype(np.float32)
if is_mlx:
import mlx.core as mx
return mx.array(noise_np).astype(mx.float32)
return noise_np.astype(np.asarray(like).dtype, copy=False)
if is_mlx:
import mlx.core as mx
return mx.random.normal(shape).astype(mx.float32)
return np.random.standard_normal(shape).astype(np.asarray(like).dtype, copy=False)
-137
View File
@@ -1,137 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Small MLX RIFE wrapper for frame interpolation experiments.
The backend is the Apple-Silicon-native ``rife-mlx`` package, using the
``mlx-community/RIFE-4.25`` weights. Frames are HWC RGB ``uint8`` arrays.
"""
from __future__ import annotations
from collections.abc import Iterable
from functools import lru_cache
import numpy as np
from huggingface_hub.utils import LocalEntryNotFoundError
class RIFEBackendError(RuntimeError):
"""Raised when the MLX RIFE backend cannot be loaded or run."""
class RIFEWeightsUnavailableError(RIFEBackendError):
"""Raised when uncached RIFE weights cannot be downloaded."""
def aligned_keyframe_count(target_frames: int, factor: int, temporal_compression: int = 4) -> int:
"""Return the smallest VAE-aligned keyframe count that RIFE can expand to the target."""
if target_frames < 1:
raise ValueError(f"target_frames must be >= 1, got {target_frames}")
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
if temporal_compression < 1:
raise ValueError(f"temporal_compression must be >= 1, got {temporal_compression}")
required_intervals = (target_frames - 1 + factor - 1) // factor
aligned_intervals = ((required_intervals + temporal_compression - 1) // temporal_compression * temporal_compression)
return aligned_intervals + 1
def _require_hwc_rgb(frame: np.ndarray, index: int) -> np.ndarray:
array = np.asarray(frame)
if array.ndim != 3 or array.shape[2] != 3:
raise ValueError(f"frame {index} must have shape HxWx3, got {array.shape}")
if array.dtype != np.uint8:
array = np.clip(array, 0, 255).astype(np.uint8)
return np.ascontiguousarray(array)
@lru_cache(maxsize=2)
def load_model(version: str = "4.25", weights_dir: str | None = None):
"""Load the MLX-native RIFE model.
``weights_dir`` is passed through to ``build_model`` in the vendored ``rife_mlx``.
When it is ``None``, the package downloads/uses the Hugging Face
``mlx-community/RIFE-4.25`` snapshot.
"""
try:
from fastvideo.third_party.rife_mlx.utils.weights import build_model
except ImportError:
# Fall back to a separately installed upstream package, for anyone who
# already has one in the environment.
try:
from rife_mlx.utils.weights import build_model
except ImportError as exc:
raise RIFEBackendError("MLX RIFE backend is unavailable. It ships vendored under "
"fastvideo/third_party/rife_mlx, so this usually means MLX "
"itself is missing: install with `uv pip install -e '.[mlx]'`.") from exc
try:
return build_model(version, weights_dir=weights_dir)
except LocalEntryNotFoundError as exc:
raise RIFEWeightsUnavailableError(f"MLX RIFE {version} weights are unavailable: {exc}") from exc
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
raise RIFEBackendError(f"Failed to load MLX RIFE {version}: {exc}") from exc
def interpolate_pair(
frame_a: np.ndarray,
frame_b: np.ndarray,
timestep: float = 0.5,
*,
model=None,
scale: float = 1.0,
) -> np.ndarray:
"""Interpolate one RGB frame between two input RGB frames."""
if not 0.0 < timestep < 1.0:
raise ValueError(f"timestep must be inside (0, 1), got {timestep}")
img0 = _require_hwc_rgb(frame_a, 0)
img1 = _require_hwc_rgb(frame_b, 1)
if img0.shape != img1.shape:
raise ValueError(f"frame shapes must match, got {img0.shape} and {img1.shape}")
if model is None:
model = load_model()
try:
try:
from fastvideo.third_party.rife_mlx.pipeline_mlx import interpolate_pair as _interpolate_pair
except ImportError:
from rife_mlx.pipeline_mlx import interpolate_pair as _interpolate_pair
return _interpolate_pair(model, img0, img1, timestep=timestep, scale=scale)
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
raise RIFEBackendError(f"MLX RIFE interpolation failed at timestep={timestep}: {exc}") from exc
def interpolate(
frames: list[np.ndarray] | Iterable[np.ndarray],
factor: int = 2,
*,
model=None,
scale: float = 1.0,
) -> list[np.ndarray]:
"""Return an Nx interpolated frame list.
For ``len(frames)=41`` and ``factor=2``, the output length is 81:
``(41 - 1) * 2 + 1``. Original keyframes are preserved in order and RIFE
fills ``factor - 1`` intermediate timesteps between each adjacent pair.
"""
frame_list = [_require_hwc_rgb(frame, idx) for idx, frame in enumerate(frames)]
if factor < 1:
raise ValueError(f"factor must be >= 1, got {factor}")
if len(frame_list) < 2 or factor == 1:
return [frame.copy() for frame in frame_list]
first_shape = frame_list[0].shape
for idx, frame in enumerate(frame_list[1:], start=1):
if frame.shape != first_shape:
raise ValueError(f"all frames must have the same shape; frame 0={first_shape}, frame {idx}={frame.shape}")
if model is None:
model = load_model()
out: list[np.ndarray] = []
for left, right in zip(frame_list[:-1], frame_list[1:], strict=True):
out.append(left)
for step in range(1, factor):
out.append(interpolate_pair(left, right, step / factor, model=model, scale=scale))
out.append(frame_list[-1])
return out
-139
View File
@@ -1,139 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""On-device (MLX) DMD sampling for the FastWan runtime.
The hybrid proof-of-concept ran the FastWan DiT in MLX but bounced every
denoising step back through torch/NumPy to run the DMD scheduler math
(``MLX -> np.array -> torch (CPU) -> np.array -> MLX``). That host round-trip
forces a full device sync per step and defeats MLX's lazy graph execution.
This module mirrors the exact DMD arithmetic from
``fastvideo/models/utils.py::pred_noise_to_pred_video`` and
``FlowMatchEulerDiscreteScheduler.add_noise`` while keeping every large tensor
on the MLX device. The schedule lookup (``argmin`` over the ~1000-entry
training schedule) is done once on the host in NumPy: it is tiny, it is the
same value torch would compute, and it sidesteps the reduction-index quirk that
affects ``argmin`` on the Metal/MPS backends (see the CPU fallbacks in
``fastvideo/models/utils.py`` and ``scheduling_flow_match_euler_discrete.py``).
Because the DMD loop applies a single scalar timestep per step, ``sigma`` is a
scalar and the update is a plain elementwise affine combination — no
permute/flatten reshaping is required.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
import numpy as np
if TYPE_CHECKING: # pragma: no cover - typing only
import mlx.core as mx
@dataclass(frozen=True)
class MLXDMDSchedule:
"""Host-side copy of a flow-match scheduler's ``(sigmas, timesteps)``.
Holds the full training schedule so a DMD timestep (e.g. one of
``1000, 757, 522``) can be mapped to its flow-match ``sigma`` with the same
nearest-timestep lookup the torch path uses.
"""
sigmas: np.ndarray
timesteps: np.ndarray
@classmethod
def from_torch_scheduler(cls, scheduler: Any) -> MLXDMDSchedule:
"""Snapshot ``scheduler.sigmas`` / ``scheduler.timesteps`` to NumPy.
Matches ``pred_noise_to_pred_video`` / ``add_noise``, which index the
scheduler's *full* training schedule (not the per-inference subset).
"""
sigmas = scheduler.sigmas.detach().to("cpu").double().numpy()
timesteps = scheduler.timesteps.detach().to("cpu").double().numpy()
return cls(sigmas=np.asarray(sigmas), timesteps=np.asarray(timesteps))
def sigma_for(self, timestep: float) -> float:
"""
Find the sigma associated with the scheduled timestep nearest to the given timestep.
Parameters:
timestep (float): Timestep for which to find the nearest scheduled sigma.
Returns:
float: Sigma associated with the nearest scheduled timestep.
"""
idx = int(np.argmin(np.abs(self.timesteps - float(timestep))))
return float(self.sigmas[idx])
def pred_noise_to_pred_video(
pred_noise: mx.array,
noise_input_latent: mx.array,
sigma: float,
) -> mx.array:
"""
Compute the clean latent prediction from a flow-matching noise prediction.
Parameters:
pred_noise (mx.array): Predicted noise.
noise_input_latent (mx.array): Noised latent input.
sigma (float): Noise level used for the prediction.
Returns:
mx.array: Predicted clean latent.
"""
return noise_input_latent - sigma * pred_noise
def add_noise(
clean_latent: mx.array,
noise: mx.array,
sigma: float,
) -> mx.array:
"""Flow-match forward noising, mirroring the scheduler's ``add_noise``.
``sample = (1 - sigma) * clean_latent + sigma * noise``.
"""
return (1.0 - sigma) * clean_latent + sigma * noise
def dmd_step(
*,
latents: mx.array,
noise_input_latent: mx.array,
pred_noise: mx.array,
schedule: MLXDMDSchedule,
timestep: float,
next_timestep: float | None,
noise: mx.array | None = None,
) -> mx.array:
"""
Compute one DMD sampling update, optionally re-noising the clean latent prediction.
Args:
latents: Retained for call-site compatibility and not used in the update.
noise_input_latent: Noisy latent used to compute the clean prediction.
pred_noise: Predicted noise or velocity.
schedule: Flow-matching schedule used to map timesteps to sigmas.
timestep: Current sampling timestep.
next_timestep: Timestep for the next update, or `None` for the final step.
noise: Fresh noise used for re-noising intermediate steps.
Returns:
The re-noised latent for the next step or the clean latent prediction on
the final step.
Raises:
ValueError: If `next_timestep` is provided without `noise`.
"""
del latents # symmetry with the torch loop; not needed for the math.
sigma = schedule.sigma_for(timestep)
pred_video = pred_noise_to_pred_video(pred_noise, noise_input_latent, sigma)
if next_timestep is None:
return pred_video
if noise is None:
raise ValueError("dmd_step requires `noise` when `next_timestep` is set (re-noise step).")
sigma_next = schedule.sigma_for(next_timestep)
return add_noise(pred_video, noise, sigma_next)
-191
View File
@@ -1,191 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Optional TAEHV decode helpers for Apple Silicon FastWan experiments.
The TAEHV module itself is vendored at ``fastvideo/third_party/taehv`` (MIT,
madebyollin/taehv), so no source code is downloaded or executed at runtime.
Only the ``taew2_1.pth`` checkpoint is fetched on demand, and its sha256 is
verified before use.
"""
from __future__ import annotations
import hashlib
import importlib.util
import urllib.request
from pathlib import Path
import numpy as np
TAEW2_1_CHECKPOINT_URL = "https://raw.githubusercontent.com/madebyollin/taehv/main/taew2_1.pth"
# sha256 of the upstream taew2_1.pth this module was validated against
# (fetched 2026-07-02). If upstream publishes a new checkpoint, revalidate the
# decode path and update this pin.
TAEW2_1_CHECKPOINT_SHA256 = "d26151e76cdc2c9424bef988de874b33d9a53f30ef3060cd556c429c469c797e"
# Wan2.2 5B (z_dim=48) — see madebyollin/taehv taew2_2.pth; prefer
# ``fastvideo.mlx_runtime.wan_vae.ensure_taehv_checkpoint(z_dim=48)`` for new code.
TAEW2_2_CHECKPOINT_URL = ("https://raw.githubusercontent.com/madebyollin/taehv/"
"563f40bdc820ed86bcad72ea515ee48f06bd22ec/taew2_2.pth")
def _default_cache_dir() -> Path:
"""Return the default directory used to cache TAEHV checkpoints.
Returns:
Path: The TAEHV checkpoint cache directory under the user's home directory.
"""
return Path.home() / ".cache" / "fastvideo" / "taehv"
def _sha256(path: Path) -> str:
"""
Compute the SHA-256 digest of a file.
Parameters:
path (Path): The file whose contents are hashed.
Returns:
str: The file's SHA-256 digest in hexadecimal form.
"""
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def _verify_checkpoint(path: Path) -> None:
"""Verify that a TAEW2.1 checkpoint matches the expected SHA-256 digest.
Parameters:
path (Path): Path to the checkpoint file.
Raises:
RuntimeError: If the checkpoint digest does not match the expected value.
"""
actual = _sha256(path)
if actual != TAEW2_1_CHECKPOINT_SHA256:
raise RuntimeError(f"TAEHV checkpoint at {path} failed sha256 verification "
f"(expected {TAEW2_1_CHECKPOINT_SHA256}, got {actual}). "
"Delete the file to re-download it, or pass --taehv-checkpoint-path "
"pointing at a checkpoint you trust.")
def ensure_taew2_1_checkpoint(checkpoint_path: Path | None = None) -> Path:
"""
Ensure the TAEW2.1 checkpoint is available locally.
A caller-provided path is treated as trusted and is only checked for existence.
The module-managed cached checkpoint is verified against the pinned SHA-256 digest
after downloading or before reuse.
Parameters:
checkpoint_path (Path | None): Optional path to a caller-provided checkpoint.
Returns:
Path: The available checkpoint path.
Raises:
FileNotFoundError: If a caller-provided checkpoint does not exist.
RuntimeError: If a module-managed checkpoint fails verification.
"""
if checkpoint_path is not None:
if not checkpoint_path.exists():
raise FileNotFoundError(f"TAEHV checkpoint not found: {checkpoint_path}")
return checkpoint_path
checkpoint_path = _default_cache_dir() / "taew2_1.pth"
if not checkpoint_path.exists():
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
print(f"Downloading {TAEW2_1_CHECKPOINT_URL} -> {checkpoint_path}")
import socket
import tempfile
# Download to a temporary file, verify, then atomically rename.
with tempfile.NamedTemporaryFile(
mode="wb",
dir=checkpoint_path.parent,
prefix=".tmp_taew2_1_",
suffix=".pth",
delete=False,
) as tmp_file:
tmp_path = Path(tmp_file.name)
try:
old_timeout = socket.getdefaulttimeout()
socket.setdefaulttimeout(300)
try:
urllib.request.urlretrieve(
TAEW2_1_CHECKPOINT_URL,
tmp_path, # noqa: S310 - pinned public artifact, hash-verified below.
)
finally:
socket.setdefaulttimeout(old_timeout)
_verify_checkpoint(tmp_path)
tmp_path.replace(checkpoint_path)
except Exception:
tmp_path.unlink(missing_ok=True)
raise
else:
_verify_checkpoint(checkpoint_path)
return checkpoint_path
def _load_taehv_class(source_path: Path | None):
"""Load the TAEHV class from the vendored implementation or a local source override.
Parameters:
source_path (Path | None): Path to a local Python file defining `TAEHV`; `None` selects the vendored implementation.
Returns:
The loaded `TAEHV` class.
Raises:
RuntimeError: If the specified source cannot be loaded.
"""
if source_path is None:
from fastvideo.third_party.taehv import TAEHV
return TAEHV
# Explicit local override for experimenting with a modified TAEHV; this is
# a user-supplied file on disk, never something this module downloads.
spec = importlib.util.spec_from_file_location("fastvideo_external_taehv", source_path)
if spec is None or spec.loader is None:
raise RuntimeError(f"Could not load TAEHV source from {source_path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.TAEHV
def decode_latents_to_video_taehv(
*,
latents_np: np.ndarray,
output_path: Path,
fps: int,
device,
dtype,
parallel: bool,
source_path: Path | None = None,
checkpoint_path: Path | None = None,
) -> None:
"""Decode Wan/FastWan diffusion latents with TAEW2.1 and export MP4.
TAEHV's Wan wrapper expects the diffusion latents directly, without applying
the standard Wan VAE's `latents_mean` / `latents_std` shift.
"""
import torch
from diffusers.utils import export_to_video
checkpoint_path = ensure_taew2_1_checkpoint(checkpoint_path)
TAEHV = _load_taehv_class(source_path)
taehv = TAEHV(str(checkpoint_path)).to(device=device, dtype=dtype)
taehv.eval()
latents = torch.from_numpy(latents_np).to(device=device, dtype=dtype)
with torch.no_grad():
video_ntchw = taehv.decode_video(
latents.transpose(1, 2),
parallel=parallel,
show_progress_bar=False,
)
video = video_ntchw.transpose(1, 2)
video_np = video[0].permute(1, 2, 3, 0).float().cpu().numpy()
output_path.parent.mkdir(parents=True, exist_ok=True)
export_to_video(video_np, str(output_path), fps=fps)
-286
View File
@@ -1,286 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Wan2.2-TI2V-5B dense MLX runtime — Track D.
The Wan2.2 TI2V-5B (FullAttn) differs from the ported Wan2.1-T2V only in:
- **Scale** (24 heads x 128, hidden 3072, ffn 14336) — pure config, block math
identical, so the dense loader ``mlx_dit_from_diffusers_safetensors`` loads the
weights unchanged and we re-wrap the blocks here.
- **Per-token timestep conditioning** (``expand_timesteps=True``): the timestep is
``[batch, seq_len]`` (a level per patch token — how TI2V keeps the conditioning
image frame at t=0 while the video frames are noised). ``timestep_proj`` becomes
``[batch, seq_len, 6, dim]`` and the block/output modulation is per-token
(``[B, L, dim]``), a direct broadcast — this module implements exactly that.
I2V rides on the same forward: encode the image, replace the first latent frame,
and set that frame's timestep to 0 (handled by the caller / sampler). See
``docs/design/ti2v_5b_port_guide.md``.
"""
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING, Any
from collections.abc import Callable
from fastvideo.logger import init_logger
from fastvideo.mlx_runtime.fastwan import (
MLXWanT2VCrossAttention,
gelu_tanh,
layer_norm,
linear,
mlx_dit_from_diffusers_safetensors,
rms_norm,
silu,
timestep_embedding,
weight_dtype,
)
if TYPE_CHECKING:
import mlx.core as mx
logger = init_logger(__name__)
class MLXWan22TransformerBlock:
"""Dense Wan block with per-token (``[B, L, dim]``) timestep modulation."""
def __init__(
self,
weights: dict[str, mx.array],
*,
dim: int,
ffn_dim: int,
num_heads: int,
eps: float = 1e-6,
):
self.weights = weights
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.eps = eps
self.attn2 = MLXWanT2VCrossAttention(weights, dim=dim, num_heads=num_heads, eps=eps)
def __call__(self, hidden_states, encoder_hidden_states, timestep_proj, cos, sin) -> mx.array:
import mlx.core as mx
orig_dtype = hidden_states.dtype
batch = hidden_states.shape[0]
# timestep_proj: [B, L, 6, dim] -> six per-token [B, L, dim] modulations.
e = self.weights["scale_shift_table"][None] + timestep_proj.astype(mx.float32)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = [
part.squeeze(2) for part in mx.split(e, 6, axis=2)
]
# 1. Self-attention (dense, bidirectional) with per-token modulation.
norm_hidden = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm_hidden = (norm_hidden * (1.0 + scale_msa) + shift_msa).astype(orig_dtype)
query = linear(norm_hidden, self.weights["to_q.weight"], self.weights.get("to_q.bias"))
key = linear(norm_hidden, self.weights["to_k.weight"], self.weights.get("to_k.bias"))
value = linear(norm_hidden, self.weights["to_v.weight"], self.weights.get("to_v.bias"))
query = rms_norm(query, self.weights["norm_q.weight"],
eps=self.eps).reshape(batch, -1, self.num_heads, self.head_dim)
key = rms_norm(key, self.weights["norm_k.weight"], eps=self.eps).reshape(batch, -1, self.num_heads,
self.head_dim)
value = value.reshape(batch, -1, self.num_heads, self.head_dim)
from fastvideo.mlx_runtime.fastwan import apply_rotary_emb
query = apply_rotary_emb(query, cos, sin, is_neox_style=False)
key = apply_rotary_emb(key, cos, sin, is_neox_style=False)
attn = mx.fast.scaled_dot_product_attention(
query.transpose(0, 2, 1, 3),
key.transpose(0, 2, 1, 3),
value.transpose(0, 2, 1, 3),
scale=self.head_dim**-0.5,
).transpose(0, 2, 1, 3)
attn = attn.reshape(batch, -1, self.dim)
attn = linear(attn, self.weights["to_out.weight"], self.weights.get("to_out.bias"))
hidden_states = hidden_states + (attn * gate_msa).astype(orig_dtype)
norm_hidden = layer_norm(hidden_states.astype(mx.float32),
weight=self.weights["self_attn_residual_norm.norm.weight"],
bias=self.weights["self_attn_residual_norm.norm.bias"],
eps=self.eps).astype(orig_dtype)
# 2. Cross-attention, then per-token shift/scale modulation.
cross = self.attn2(norm_hidden, encoder_hidden_states)
hidden_states = hidden_states + cross
norm_hidden = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm_hidden = (norm_hidden * (1.0 + c_scale_msa) + c_shift_msa).astype(orig_dtype)
# 3. Feed-forward with per-token gate.
ff = linear(norm_hidden, self.weights["ffn.fc_in.weight"], self.weights.get("ffn.fc_in.bias"))
ff = gelu_tanh(ff)
ff = linear(ff, self.weights["ffn.fc_out.weight"], self.weights.get("ffn.fc_out.bias"))
hidden_states = hidden_states + (ff * c_gate_msa).astype(orig_dtype)
return hidden_states.astype(orig_dtype)
class MLXWan22DiT:
"""Wan2.2-TI2V-5B dense DiT with per-token timestep conditioning."""
def __init__(
self,
weights: dict[str, mx.array],
blocks: list[MLXWan22TransformerBlock],
config: dict,
*,
compile: bool = False,
) -> None:
import os
self.weights = weights
self.blocks = blocks
self.config = config
self.num_heads = int(config["num_attention_heads"])
self.head_dim = int(config["attention_head_dim"])
self.hidden_size = self.num_heads * self.head_dim
self.freq_dim = int(config["freq_dim"])
self.patch_size = tuple(config["patch_size"])
self.out_channels = int(config["out_channels"])
self.eps = float(config.get("eps", 1e-6))
self._enable_compile = compile or os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1"
self._compiled_forward: Callable[..., Any] | None = None
self._compiled_signature: tuple | None = None
def _patch_embed(self, hidden_states) -> mx.array:
batch, channels, frames, height, width = hidden_states.shape
pt, ph, pw = self.patch_size
patch_dim = channels * pt * ph * pw
x = hidden_states.reshape(batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw)
x = x.transpose(0, 2, 4, 6, 1, 3, 5, 7).reshape(batch, -1, patch_dim)
return linear(x, self.weights["patch_embedding.weight"], self.weights.get("patch_embedding.bias"))
def _condition(self, timestep, encoder_hidden_states) -> tuple:
"""Per-token conditioning. ``timestep`` is ``[B, L]`` (one level per token)."""
batch, seq = timestep.shape
t_freq = timestep_embedding(timestep.reshape(-1), self.freq_dim).astype(
weight_dtype(self.weights["condition_embedder.time_embedder.linear_1.weight"]))
temb = linear(t_freq, self.weights["condition_embedder.time_embedder.linear_1.weight"],
self.weights["condition_embedder.time_embedder.linear_1.bias"])
temb = silu(temb)
temb = linear(temb, self.weights["condition_embedder.time_embedder.linear_2.weight"],
self.weights["condition_embedder.time_embedder.linear_2.bias"])
timestep_proj = linear(silu(temb), self.weights["condition_embedder.time_proj.weight"],
self.weights["condition_embedder.time_proj.bias"])
timestep_proj = timestep_proj.reshape(batch, seq, 6, self.hidden_size)
ehs = linear(encoder_hidden_states, self.weights["condition_embedder.text_embedder.linear_1.weight"],
self.weights["condition_embedder.text_embedder.linear_1.bias"])
ehs = gelu_tanh(ehs)
ehs = linear(ehs, self.weights["condition_embedder.text_embedder.linear_2.weight"],
self.weights["condition_embedder.text_embedder.linear_2.bias"])
temb_out = temb.reshape(batch, seq, self.hidden_size)
return temb_out, timestep_proj, ehs
def _output(self, hidden_states, temb_out, *, batch, frames, height, width) -> mx.array:
import mlx.core as mx
pt, ph, pw = self.patch_size
post_pt, post_ph, post_pw = frames // pt, height // ph, width // pw
# Per-token output modulation: scale_shift_table[1,2,dim] + temb[B,L,1,dim].
e = self.weights["scale_shift_table"][None] + temb_out[:, :, None, :].astype(mx.float32)
shift, scale = [part.squeeze(2) for part in mx.split(e, 2, axis=2)]
norm = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
norm = (norm * (1.0 + scale) + shift).astype(weight_dtype(self.weights["proj_out.weight"]))
out = linear(norm, self.weights["proj_out.weight"], self.weights["proj_out.bias"])
out = out.reshape(batch, post_pt, post_ph, post_pw, pt, ph, pw, self.out_channels)
out = out.transpose(0, 7, 1, 4, 2, 5, 3, 6)
return out.reshape(batch, self.out_channels, frames, height, width)
def _forward(self, hidden_states, encoder_hidden_states, timestep, cos, sin) -> mx.array:
batch, _, frames, height, width = hidden_states.shape
hidden = self._patch_embed(hidden_states)
temb_out, timestep_proj, ehs = self._condition(timestep, encoder_hidden_states)
for block in self.blocks:
hidden = block(hidden, ehs, timestep_proj, cos, sin)
return self._output(hidden, temb_out, batch=batch, frames=frames, height=height, width=width)
def __call__(self, hidden_states, encoder_hidden_states, timestep, freqs_cis) -> mx.array:
cos, sin = freqs_cis
if self._enable_compile and cos is not None:
import mlx.core as mx
# One traced graph per input signature, and each pins its own copy of
# the quantized weights. --refine denoises at two resolutions, so
# keeping both alive doubles resident DiT memory. Retire the previous
# graph when the signature changes.
signature = (hidden_states.shape, encoder_hidden_states.shape, timestep.shape)
if self._compiled_forward is not None and signature != self._compiled_signature:
self._compiled_forward = None
self._compiled_signature = None
mx.clear_cache()
if self._compiled_forward is None:
self._compiled_forward = mx.compile(self._forward)
self._compiled_signature = signature
compiled_forward = self._compiled_forward
try:
return compiled_forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
except Exception as exc: # noqa: BLE001 - some graphs may not trace; fall back to eager.
logger.warning(
"Wan2.2 mx.compile forward failed (%s); falling back to eager execution.",
exc,
)
self._enable_compile = False
self._compiled_forward = None
self._compiled_signature = None
return self._forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
def mlx_wan22_dit_from_diffusers_safetensors(
checkpoint_path: str | Path,
config_path: str | Path,
*,
dtype: str = "fp16",
num_blocks: int | None = None,
quantization=None,
compile: bool = False,
) -> MLXWan22DiT:
"""Load Wan2.2-TI2V-5B (FullAttn) into ``MLXWan22DiT`` via the dense loader."""
dense = mlx_dit_from_diffusers_safetensors(checkpoint_path,
config_path,
dtype=dtype,
num_blocks=num_blocks,
quantization=quantization)
inner_dim = int(dense.config["num_attention_heads"]) * int(dense.config["attention_head_dim"])
blocks = [
MLXWan22TransformerBlock(block.weights,
dim=inner_dim,
ffn_dim=int(dense.config["ffn_dim"]),
num_heads=int(dense.config["num_attention_heads"]),
eps=float(dense.config.get("eps", 1e-6))) for block in dense.blocks
]
return MLXWan22DiT(dense.weights, blocks, dense.config, compile=compile)
def mlx_wan22_dit_from_mlx_checkpoint(
checkpoint_dir: str | Path,
*,
compile: bool = False,
) -> MLXWan22DiT:
"""Rewrap a persisted MLX DiT checkpoint with Wan2.2 conditioning.
The generic checkpoint loader intentionally rebuilds ``MLXWanDiT`` because
it is also used by the Wan2.1 runtime. Wan2.2 TI2V has the same weight
layout but needs per-token timestep modulation, so callers must rewrap the
loaded weights and blocks as :class:`MLXWan22DiT` before sampling.
"""
from fastvideo.mlx_runtime.checkpoint import load_mlx_dit_checkpoint
dense = load_mlx_dit_checkpoint(checkpoint_dir)
inner_dim = int(dense.config["num_attention_heads"]) * int(dense.config["attention_head_dim"])
blocks = [
MLXWan22TransformerBlock(
block.weights,
dim=inner_dim,
ffn_dim=int(dense.config["ffn_dim"]),
num_heads=int(dense.config["num_attention_heads"]),
eps=float(dense.config.get("eps", 1e-6)),
) for block in dense.blocks
]
return MLXWan22DiT(dense.weights, blocks, dense.config, compile=compile)
-113
View File
@@ -1,113 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Dense DMD sampling for MLXWan22DiT (Wan2.2 per-token timestep).
Matches the FastVideo pipeline's warped DMD schedule (``warp_denoising_step=True``,
``dmd_denoising_steps=[1000,757,522]``, ``flow_shift=5.0`` for TI2V-5B) rather
than treating raw step indices as continuous timesteps (a bug in early demos).
"""
from __future__ import annotations
from typing import TYPE_CHECKING
from collections.abc import Sequence
import numpy as np
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step, pred_noise_to_pred_video
if TYPE_CHECKING:
import mlx.core as mx
from fastvideo.mlx_runtime.wan22 import MLXWan22DiT
def build_wan22_dmd_schedule(
dmd_denoising_steps: Sequence[int] | None = None,
*,
flow_shift: float = 5.0,
warp_denoising_step: bool = True,
) -> tuple[MLXDMDSchedule, list[float]]:
"""
Build the flow-matching schedule and continuous timesteps used for Wan2.2 DMD sampling.
Parameters:
dmd_denoising_steps (Sequence[int] | None): Denoising step values to use; defaults to 1000, 757, and 522.
flow_shift (float): Flow-matching shift applied when constructing the schedule.
warp_denoising_step (bool): Whether to convert denoising steps to scheduler-warped continuous timesteps.
Returns:
tuple[MLXDMDSchedule, list[float]]: The DMD schedule and corresponding continuous timesteps.
"""
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
steps = list(dmd_denoising_steps or [1000, 757, 522])
scheduler = FlowMatchEulerDiscreteScheduler(shift=flow_shift)
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
step_idx = torch.tensor(steps, dtype=torch.long)
if warp_denoising_step:
warped = torch.cat((scheduler.timesteps.cpu(), torch.tensor([0.0], dtype=torch.float32)))
timesteps = [float(t) for t in warped[1000 - step_idx]]
else:
timesteps = [float(s) for s in steps]
return schedule, timesteps
def sample_wan22_dmd(
model: MLXWan22DiT,
encoder_hidden_states: mx.array,
noise_latents: mx.array,
freqs_cis: tuple,
*,
dmd_denoising_steps: Sequence[int] | None = None,
flow_shift: float = 5.0,
warp_denoising_step: bool = True,
seed: int = 0,
) -> mx.array:
"""
Generate clean video latents from noisy latents using iterative DMD denoising.
Parameters:
noise_latents (mx.array): Initial noisy video latents.
freqs_cis (tuple): Rotary positional frequency tensors used by the model.
dmd_denoising_steps (Sequence[int] | None): DMD denoising steps, or the default schedule when omitted.
flow_shift (float): Flow-matching schedule shift.
warp_denoising_step (bool): Whether to warp the denoising timesteps.
seed (int): Seed for reproducible intermediate re-noising.
Returns:
mx.array: Denoised video latents.
"""
import mlx.core as mx
schedule, timesteps = build_wan22_dmd_schedule(dmd_denoising_steps,
flow_shift=flow_shift,
warp_denoising_step=warp_denoising_step)
# NumPy RNG so re-noise is bit-reproducible across MLX / torch A/B dumps.
renoise_rng = np.random.default_rng(seed)
latents = noise_latents
batch, _c, frames, height, width = latents.shape
pt, ph, pw = model.patch_size
tokens = (frames // pt) * (height // ph) * (width // pw)
last = len(timesteps) - 1
for i, t in enumerate(timesteps):
ts = mx.full((batch, tokens), float(t), dtype=mx.float32)
pred = model(latents.astype(mx.float16), encoder_hidden_states, ts, freqs_cis)
ni = latents.astype(mx.float32)
pn = pred.astype(mx.float32)
if i < last:
renoise = mx.array(renoise_rng.standard_normal(tuple(latents.shape)).astype(np.float32))
latents = dmd_step(
latents=ni,
noise_input_latent=ni,
pred_noise=pn,
schedule=schedule,
timestep=float(t),
next_timestep=float(timesteps[i + 1]),
noise=renoise,
).astype(latents.dtype)
else:
latents = pred_noise_to_pred_video(pn, ni, schedule.sigma_for(float(t))).astype(latents.dtype)
mx.eval(latents)
return latents
-557
View File
@@ -1,557 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Wan VAE decode helpers for Apple Silicon MLX inference.
Two decode backends:
1. **TAEHV (primary / fast)** — Tiny AutoEncoder (madebyollin/taehv). Fully
MLX-native Conv2d path. ``taew2_1.pth`` for Wan2.1 (z_dim=16),
``taew2_2.pth`` for Wan2.2 5B (z_dim=48, patch_size=2). Expected decode
wall-clock ~seconds vs ~minutes for the full 3D VAE on MPS.
2. **Full AutoencoderKLWan (reference / quality)** — denormalize with
``latents_mean`` / ``latents_std`` then torch decode (MPS preferred). Used
for parity gates and when TAEHV is unavailable. A pure-MLX 3D-conv port of
the residual Wan2.2 decoder is left as follow-up (causal feat-cache +
residual up blocks are large); TAEHV covers the product latency path.
Diffusion latents from the DiT are **not** mean/std-normalized for TAEHV
(matching ``taehv_decode.py``); full VAE decode **does** denormalize first
(matching ``mlx_wan_prompt_to_video.decode_latents_to_video``).
"""
from __future__ import annotations
import hashlib
import json
import urllib.request
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Literal
import numpy as np
from fastvideo.mlx_runtime.memory import cleanup_mlx, cleanup_torch_mps
GIB = 1024**3
TAEW2_1_URL = "https://raw.githubusercontent.com/madebyollin/taehv/main/taew2_1.pth"
TAEW2_2_URL = ("https://raw.githubusercontent.com/madebyollin/taehv/"
"563f40bdc820ed86bcad72ea515ee48f06bd22ec/taew2_2.pth")
# Validated 2026-07-02 / 2026-07-09 against upstream madebyollin/taehv.
TAEW2_1_SHA256 = "d26151e76cdc2c9424bef988de874b33d9a53f30ef3060cd556c429c469c797e"
TAEW2_2_SHA256 = "d053e216ca50e2bb837bbcd79b85f0366bea00e5938025572382a773b74c559a"
DecodeBackend = Literal["taehv", "taehv-torch", "wan-vae"]
def _cache_dir() -> Path:
"""Return the local directory used to cache TAEHV files."""
return Path.home() / ".cache" / "fastvideo" / "taehv"
def _sha256(path: Path) -> str:
"""Compute the SHA-256 digest of a file.
Parameters:
path (Path): Path to the file to hash.
Returns:
str: Lowercase hexadecimal SHA-256 digest.
"""
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def _verify_checkpoint(path: Path, expected_digest: str) -> None:
"""
Verify a checkpoint's SHA-256 digest against a required expected digest.
Parameters:
path (Path): Path to the checkpoint file.
expected_digest (str): Expected lowercase SHA-256 digest.
Raises:
RuntimeError: If verification is enabled and the checkpoint digest does not match.
"""
if len(expected_digest) != 64 or any(char not in "0123456789abcdef" for char in expected_digest):
raise ValueError("A valid lowercase SHA-256 digest is required for bundled TAEHV checkpoints")
actual = _sha256(path)
if actual != expected_digest:
raise RuntimeError(f"TAEHV checkpoint at {path} failed sha256 verification "
f"(expected {expected_digest}, got {actual}). "
"Delete the file to re-download it.")
def ensure_taehv_checkpoint(*, z_dim: int, checkpoint_path: Path | None = None) -> Path:
"""
Return a validated TAEHV checkpoint for the specified latent channel count.
Parameters:
z_dim (int): Number of latent channels, supported values are 16 and 48.
checkpoint_path (Path | None): Optional existing checkpoint path to validate and use.
Returns:
Path: Path to the validated TAEHV checkpoint.
Raises:
FileNotFoundError: If the supplied checkpoint path does not exist.
ValueError: If no checkpoint is mapped to the specified latent channel count.
"""
if checkpoint_path is not None:
if not checkpoint_path.exists():
raise FileNotFoundError(f"TAEHV checkpoint not found: {checkpoint_path}")
return checkpoint_path
if z_dim == 16:
name, url, expect = "taew2_1.pth", TAEW2_1_URL, TAEW2_1_SHA256
elif z_dim == 48:
name, url, expect = "taew2_2.pth", TAEW2_2_URL, TAEW2_2_SHA256
else:
raise ValueError(f"No TAEHV checkpoint mapped for z_dim={z_dim} (supported: 16, 48)")
path = _cache_dir() / name
if not path.exists():
path.parent.mkdir(parents=True, exist_ok=True)
print(f"Downloading {url} -> {path}")
import socket
import tempfile
# Download to a temporary file, verify, then atomically rename.
with tempfile.NamedTemporaryFile(
mode="wb",
dir=path.parent,
prefix=f".tmp_{name}_",
suffix=".pth",
delete=False,
) as tmp_file:
tmp_path = Path(tmp_file.name)
try:
old_timeout = socket.getdefaulttimeout()
socket.setdefaulttimeout(300)
try:
urllib.request.urlretrieve(url, tmp_path) # noqa: S310 - public pinned artifact.
finally:
socket.setdefaulttimeout(old_timeout)
_verify_checkpoint(tmp_path, expect)
tmp_path.replace(path)
except Exception:
tmp_path.unlink(missing_ok=True)
raise
else:
_verify_checkpoint(path, expect)
return path
@dataclass(frozen=True)
class WanVAEConfigView:
"""Minimal config fields needed for denormalize + spatial scale."""
z_dim: int
latents_mean: tuple[float, ...]
latents_std: tuple[float, ...]
scale_factor_spatial: int = 8
scale_factor_temporal: int = 4
patch_size: int | None = None
vae_dir: Path | None = None
@classmethod
def from_vae_dir(cls, vae_dir: Path) -> WanVAEConfigView:
"""
Load Wan VAE configuration values from a directory.
Parameters:
vae_dir (Path): Directory containing the VAE ``config.json`` file.
Returns:
WanVAEConfigView: Configuration loaded from the VAE directory.
"""
cfg = json.loads((vae_dir / "config.json").read_text())
return cls(
z_dim=int(cfg["z_dim"]),
latents_mean=tuple(float(x) for x in cfg["latents_mean"]),
latents_std=tuple(float(x) for x in cfg["latents_std"]),
scale_factor_spatial=int(cfg.get("scale_factor_spatial", 8)),
scale_factor_temporal=int(cfg.get("scale_factor_temporal", 4)),
patch_size=cfg.get("patch_size"),
vae_dir=vae_dir,
)
def denormalize_latents_np(latents: np.ndarray, config: WanVAEConfigView) -> np.ndarray:
"""
Denormalize Wan VAE latent values using the configured means and standard deviations.
Parameters:
latents (np.ndarray): Latent values in normalized form.
config (WanVAEConfigView): Wan VAE latent statistics.
Returns:
np.ndarray: Denormalized latent values as float32.
"""
mean = np.asarray(config.latents_mean, dtype=np.float32).reshape(1, -1, 1, 1, 1)
std = np.asarray(config.latents_std, dtype=np.float32).reshape(1, -1, 1, 1, 1)
return latents.astype(np.float32) * std + mean
# ---------------------------------------------------------------------------
# MLX TAEHV decoder (Conv2d stack — primary fully-MLX product path)
# ---------------------------------------------------------------------------
def _mlx_conv2d(x: Any, weight: Any, bias: Any, *, stride: int = 1) -> Any:
import mlx.core as mx
# x: NCHW, weight: OIHW
y = mx.conv2d(x.transpose(0, 2, 3, 1), weight.transpose(0, 2, 3, 1), stride=stride, padding=1)
y = y.transpose(0, 3, 1, 2)
if bias is not None:
y = y + bias.reshape(1, -1, 1, 1)
return y
def _mlx_conv2d_1x1(x: Any, weight: Any, bias: Any = None) -> Any:
"""
Applies a 1×1 convolution to an MLX tensor in channel-first layout.
Parameters:
x (Any): Input tensor with shape [batch, channels, height, width].
weight (Any): Convolution weights.
bias (Any, optional): Optional output-channel bias.
Returns:
Any: The convolved tensor with shape [batch, output_channels, height, width].
"""
import mlx.core as mx
y = mx.conv2d(x.transpose(0, 2, 3, 1), weight.transpose(0, 2, 3, 1), stride=1, padding=0)
y = y.transpose(0, 3, 1, 2)
if bias is not None:
y = y + bias.reshape(1, -1, 1, 1)
return y
def _load_torch_state(path: Path) -> dict[str, np.ndarray]:
"""Load a PyTorch state dictionary as NumPy arrays.
Parameters:
path (Path): Path to the PyTorch checkpoint.
Returns:
dict[str, np.ndarray]: State dictionary with tensors converted to NumPy arrays.
"""
import torch
sd = torch.load(path, map_location="cpu", weights_only=True)
return {k: v.detach().float().cpu().numpy() for k, v in sd.items()}
class MLXTAEHVDecoder:
"""Minimal MLX port of TAEHV ``decoder`` (parallel-over-time MemBlocks)."""
def __init__(self, checkpoint_path: Path, *, z_dim: int) -> None:
"""Initialize the TAEHV decoder from a checkpoint for the specified latent dimensionality.
Parameters:
checkpoint_path (Path): Path to the TAEHV checkpoint.
z_dim (int): Number of latent channels, determining the decoder patch size.
"""
import mlx.core as mx
self.checkpoint_path = Path(checkpoint_path)
self.latent_channels = z_dim
# Derive patch_size from z_dim: 48 channels → patch_size=2, 16 → patch_size=1
self.patch_size = 2 if z_dim == 48 else 1
self.image_channels = 3
self.frames_to_trim = 3 # TGrow strides (1,2,2) → 2**2 - 1 for w2.1/w2.2 defaults
sd = _load_torch_state(self.checkpoint_path)
# Patch TGrow kernels like upstream TAEHV.patch_tgrow_layers.
self.weights = {k: mx.array(v) for k, v in sd.items()}
self._n_f = [256, 128, 64, 64]
def decode_ntchw(self, latents_ntchw: Any) -> Any:
"""
Decode latent video batches into clipped RGB frames.
Parameters:
latents_ntchw (Any): Latents with shape ``[N, T, C, H, W]`` and the
decoder's configured latent channel count.
Returns:
Any: Decoded frames with shape ``[N, T_out, 3, H_out, W_out]`` and values
clipped to the range ``[0, 1]``.
"""
import mlx.core as mx
x = latents_ntchw
n, t, c, h, w = x.shape
if c != self.latent_channels:
raise ValueError(f"expected C={self.latent_channels}, got {c}")
x = x.reshape(n * t, c, h, w)
x = self._run_decoder_parallel(x, n=n)
# Pixel-shuffle if patch_size > 1: (NT, 3*p*p, H, W) -> (NT, 3, H*p, W*p)
if self.patch_size > 1:
p = self.patch_size
nt, c_out, hh, ww = x.shape
x = x.reshape(nt, self.image_channels, p, p, hh, ww)
x = x.transpose(0, 1, 4, 2, 5, 3).reshape(nt, self.image_channels, hh * p, ww * p)
_, c_out, hh, ww = x.shape
t_out = x.shape[0] // n
x = x.reshape(n, t_out, c_out, hh, ww)
if self.frames_to_trim > 0 and t_out > self.frames_to_trim:
x = x[:, self.frames_to_trim:]
return mx.clip(x, 0.0, 1.0)
def _run_decoder_parallel(self, x: Any, *, n: int) -> Any:
"""
Apply the TAEHV decoder stack to flattened batch and temporal frames while preserving temporal memory.
"""
import mlx.core as mx
w = self.weights
def memblock(base: int, xx: Any, past: Any) -> Any:
"""
Apply a temporal memory block to the current and past feature tensors.
Parameters:
base (int): Decoder block index used to select the block weights.
xx (Any): Current feature tensor.
past (Any): Past feature tensor concatenated with the current features.
Returns:
Any: Activated feature tensor produced by the memory block.
"""
cat = mx.concatenate([xx, past], axis=1)
h = _mlx_conv2d(cat, w[f"decoder.{base}.conv.0.weight"], w.get(f"decoder.{base}.conv.0.bias"))
h = mx.maximum(h, 0.0)
h = _mlx_conv2d(h, w[f"decoder.{base}.conv.2.weight"], w.get(f"decoder.{base}.conv.2.bias"))
h = mx.maximum(h, 0.0)
h = _mlx_conv2d(h, w[f"decoder.{base}.conv.4.weight"], w.get(f"decoder.{base}.conv.4.bias"))
skip_key = f"decoder.{base}.skip.weight"
skip = _mlx_conv2d_1x1(xx, w[skip_key], None) if skip_key in w else xx
return mx.maximum(h + skip, 0.0)
def upsample2(xx: Any) -> Any:
"""
Upsample a four-dimensional tensor by a factor of two along its spatial dimensions.
Parameters:
xx (Any): Tensor with shape `(N, C, H, W)`.
Returns:
Any: Tensor with shape `(N, C, 2H, 2W)` containing replicated spatial values.
"""
nt, c, h, ww = xx.shape
xx = xx.reshape(nt, c, h, 1, ww, 1)
xx = mx.broadcast_to(xx, (nt, c, h, 2, ww, 2))
return xx.reshape(nt, c, h * 2, ww * 2)
def tgrow(base: int, xx: Any, stride: int) -> Any:
wt = w[f"decoder.{base}.conv.weight"]
out_ch = int(xx.shape[1]) * stride
if int(wt.shape[0]) > out_ch:
wt = wt[-out_ch:]
y = _mlx_conv2d_1x1(xx, wt, None)
if stride == 1:
return y
# TGrow.forward: (NT, C*stride, H, W) -> (NT*stride, C, H, W)
nt, c, h, ww = y.shape
c_in = c // stride
y = y.reshape(nt, stride, c_in, h, ww).transpose(0, 1, 2, 3, 4)
return y.reshape(nt * stride, c_in, h, ww)
def mem_past(xx: Any) -> Any:
"""
Build a temporal memory tensor containing a zero frame followed by the preceding frame at each time step.
Parameters:
xx (Any): Flattened batch and temporal tensor with shape ``(batch * time, channels, height, width)``.
Returns:
Any: Tensor with the same shape as ``xx`` containing the preceding frame for each temporal position.
"""
nt, c, h, ww = xx.shape
t_cur = nt // n
x_ = xx.reshape(n, t_cur, c, h, ww)
# pad one zero frame at t=0, align past[t] = x[t-1]
past = mx.concatenate([mx.zeros_like(x_[:, :1]), x_[:, :-1]], axis=1)
return past.reshape(nt, c, h, ww)
# 0 Clamp, 1 conv, 2 ReLU
x = mx.tanh(x / 3.0) * 3.0
x = _mlx_conv2d(x, w["decoder.1.weight"], w.get("decoder.1.bias"))
x = mx.maximum(x, 0.0)
for mem_idx in (3, 4, 5):
x = memblock(mem_idx, x, mem_past(x))
x = upsample2(x)
x = tgrow(7, x, 1)
x = _mlx_conv2d(x, w["decoder.8.weight"], w.get("decoder.8.bias"))
for mem_idx in (9, 10, 11):
x = memblock(mem_idx, x, mem_past(x))
x = upsample2(x)
x = tgrow(13, x, 2)
x = _mlx_conv2d(x, w["decoder.14.weight"], w.get("decoder.14.bias"))
for mem_idx in (15, 16, 17):
x = memblock(mem_idx, x, mem_past(x))
x = upsample2(x)
x = tgrow(19, x, 2)
x = _mlx_conv2d(x, w["decoder.20.weight"], w.get("decoder.20.bias"))
x = mx.maximum(x, 0.0)
x = _mlx_conv2d(x, w["decoder.22.weight"], w.get("decoder.22.bias"))
return x
def decode_latents_taehv_mlx(
latents_np: np.ndarray,
*,
z_dim: int | None = None,
checkpoint_path: Path | None = None,
) -> np.ndarray:
"""
Decode latent representations with the MLX TAEHV decoder.
Parameters:
latents_np (np.ndarray): Latents arranged as [B, C, T, H, W].
z_dim (int | None): Latent channel dimension used to select the decoder checkpoint.
checkpoint_path (Path | None): Optional path to a TAEHV checkpoint.
Returns:
np.ndarray: Decoded pixels arranged as [B, T, H, W, 3] with values in [0, 1].
Raises:
ValueError: If `latents_np` does not have five dimensions.
"""
import mlx.core as mx
if latents_np.ndim != 5:
raise ValueError(f"expected [B,C,T,H,W], got {latents_np.shape}")
c = latents_np.shape[1]
z = z_dim if z_dim is not None else c
ckpt = ensure_taehv_checkpoint(z_dim=z, checkpoint_path=checkpoint_path)
dec = MLXTAEHVDecoder(ckpt, z_dim=z)
# NTCHW
x = mx.array(latents_np.transpose(0, 2, 1, 3, 4).astype(np.float32))
out = dec.decode_ntchw(x) # N T C H W
mx.eval(out)
arr = np.array(out)
# B T H W C
return arr.transpose(0, 1, 3, 4, 2)
def decode_latents_wan_vae_torch(
latents_np: np.ndarray,
*,
vae_dir: Path,
device: str = "auto",
dtype_name: str = "fp16",
) -> np.ndarray:
"""Full AutoencoderKLWan decode on torch (MPS/CPU) with mean/std denormalize.
Returns pixels ``[B, T, H, W, 3]`` float in [0, 1].
"""
import torch
from diffusers import AutoencoderKLWan
from diffusers.video_processor import VideoProcessor
if device == "auto":
device = "mps" if torch.backends.mps.is_available() else "cpu"
dtype = torch.float16 if dtype_name == "fp16" and device == "mps" else torch.float32
config = WanVAEConfigView.from_vae_dir(vae_dir)
vae = AutoencoderKLWan.from_pretrained(vae_dir, torch_dtype=dtype, local_files_only=True).to(device)
vae.eval()
latents = torch.from_numpy(latents_np.astype(np.float32)).to(device=device, dtype=dtype)
mean = torch.tensor(config.latents_mean, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
inv_std = (1.0 / torch.tensor(config.latents_std, device=device, dtype=dtype)).view(1, -1, 1, 1, 1)
latents = latents / inv_std + mean # matches prompt_to_video path
with torch.no_grad():
video = vae.decode(latents, return_dict=False)[0]
video = VideoProcessor(vae_scale_factor=config.scale_factor_spatial).postprocess_video(video, output_type="np")
return video # [B, T, H, W, 3]
def decode_latents_to_video(
latents_np: np.ndarray,
output_path: Path,
*,
fps: int = 16,
backend: DecodeBackend = "taehv",
vae_dir: Path | None = None,
z_dim: int | None = None,
taehv_checkpoint: Path | None = None,
torch_device: str = "auto",
) -> dict[str, Any]:
"""Decode latent video frames and export them as an MP4 file.
Parameters:
latents_np (np.ndarray): Latent video representation to decode.
output_path (Path): Destination path for the MP4 file.
fps (int): Output video frame rate.
backend (DecodeBackend): Decoder backend to use.
vae_dir (Path | None): Directory containing the full Wan VAE when using
the ``wan-vae`` backend.
z_dim (int | None): Latent channel count for TAEHV decoding.
taehv_checkpoint (Path | None): Optional TAEHV checkpoint path.
torch_device (str): PyTorch device selection for PyTorch-based decoding.
Returns:
dict[str, Any]: Decode time in seconds, backend name, output path, frame
count, and video resolution.
Raises:
ValueError: If the full VAE backend lacks ``vae_dir`` or the backend is
unknown.
"""
import time
from diffusers.utils import export_to_video
t0 = time.perf_counter()
if backend in ("taehv", "taehv-torch"):
c = latents_np.shape[1] if z_dim is None else z_dim
if backend == "taehv":
video = decode_latents_taehv_mlx(latents_np, z_dim=c, checkpoint_path=taehv_checkpoint)
else:
# torch TAEHV (regression / parity reference)
import torch
from fastvideo.third_party.taehv import TAEHV
ckpt = ensure_taehv_checkpoint(z_dim=c, checkpoint_path=taehv_checkpoint)
if torch_device == "auto":
torch_device = "mps" if torch.backends.mps.is_available() else "cpu"
dtype = torch.float16 if torch_device == "mps" else torch.float32
model = TAEHV(str(ckpt)).to(device=torch_device, dtype=dtype).eval()
lat = torch.from_numpy(latents_np).to(device=torch_device, dtype=dtype)
with torch.no_grad():
out = model.decode_video(lat.transpose(1, 2), parallel=True, show_progress_bar=False)
video = out[0].permute(0, 2, 3, 1).float().cpu().numpy()[None, ...]
# out is NTCHW -> need BTHWC; decode_video returns NTCHW for batch
if video.ndim == 4:
video = video[None]
elif backend == "wan-vae":
if vae_dir is None:
raise ValueError("vae_dir required for wan-vae backend")
video = decode_latents_wan_vae_torch(latents_np, vae_dir=vae_dir, device=torch_device)
else:
raise ValueError(f"unknown backend {backend}")
decode_s = time.perf_counter() - t0
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
# export_to_video expects list/array of frames HxWxC
frames = video[0]
frames = np.clip(frames, 0.0, 1.0)
export_to_video(frames, str(output_path), fps=fps)
if backend == "taehv":
cleanup_mlx()
else:
if backend == "taehv-torch":
del model, lat, out
cleanup_torch_mps()
return {
"decode_s": decode_s,
"backend": backend,
"output_path": str(output_path),
"num_frames": int(frames.shape[0]),
"resolution": f"{frames.shape[2]}x{frames.shape[1]}" if frames.ndim == 4 else None,
}
-223
View File
@@ -1,223 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Chunked non-causal sliding-window self-attention for MLX scaling studies.
This module is intentionally standalone (``mlx.core`` + stdlib only) so it can
be micro-benchmarked without pulling in the DiT / FastVideo stack.
Window policy
-------------
**Symmetric** sliding window (non-causal). For query index ``i`` the allowed
key indices are:
sinks: ``j in [0, sink)`` (always visible to every query, if ``sink > 0``)
local: ``j in [max(0, i - half), min(S, i + half + 1))``
where ``half = window // 2``
so each query sees roughly ``window + 1`` local keys (plus any sinks outside
that range). This is appropriate for a dense, bidirectional DiT denoise pass.
Implementation note (FLOPs)
---------------------------
A full-size additive attention mask still materialises an ``O(S^2)`` score
matrix inside SDPA and does **not** reduce work. Instead we tile the sequence
into query blocks and run ``mx.fast.scaled_dot_product_attention`` only against
the union of keys that block needs (local slice ± sinks). That makes
per-block work ``O(chunk * (window + sink) * D)`` and total work
``O(S * (window + sink) * D)``.
"""
from __future__ import annotations
import mlx.core as mx
def _default_scale(head_dim: int, scale: float | None) -> float:
"""
Determine the attention scaling factor from an explicit value or head dimension.
Parameters:
head_dim (int): The attention head dimension used to derive the default scale.
scale (Optional[float]): An explicit scaling factor.
Returns:
float: The explicit scale converted to a float, or the reciprocal square root of `head_dim`.
Raises:
ValueError: If `scale` is not provided and `head_dim` is not positive.
"""
if scale is not None:
return float(scale)
if head_dim <= 0:
raise ValueError(f"head_dim must be positive, got {head_dim}")
return head_dim**-0.5
def _validate_qkv(q: mx.array, k: mx.array, v: mx.array) -> tuple[int, int, int, int]:
"""
Validate compatible rank-4 query, key, and value tensors.
Parameters:
q (mx.array): Query tensor shaped `(B, H, S, D)`.
k (mx.array): Key tensor with the same shape as `q`.
v (mx.array): Value tensor with the same shape as `q`.
Returns:
tuple[int, int, int, int]: Batch size, head count, sequence length, and head dimension.
Raises:
ValueError: If the tensors are not rank 4, do not have identical shapes, or have an empty sequence or head dimension.
"""
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
raise ValueError(f"q/k/v must be rank-4 (B, H, S, D); got shapes "
f"{q.shape}, {k.shape}, {v.shape}")
b, h, s, d = q.shape
if k.shape != (b, h, s, d) or v.shape != (b, h, s, d):
raise ValueError(f"q/k/v shapes must match exactly; got q={q.shape}, k={k.shape}, v={v.shape}")
if s == 0:
raise ValueError("sequence length S must be > 0")
if d == 0:
raise ValueError("head dim D must be > 0")
return b, h, s, d
def full_attention(
q: mx.array,
k: mx.array,
v: mx.array,
scale: float | None = None,
) -> mx.array:
"""
Compute dense scaled dot-product attention over the full sequence.
Parameters:
scale (float, optional): Attention scaling factor. If omitted, uses the
inverse square root of the head dimension.
Returns:
mx.array: Attention output with shape ``(B, H, S, D)``.
"""
_, _, _, d = _validate_qkv(q, k, v)
sc = _default_scale(d, scale)
return mx.fast.scaled_dot_product_attention(q, k, v, scale=sc)
def _concat_kv_slices(
k: mx.array,
v: mx.array,
ranges: list[tuple[int, int]],
) -> tuple[mx.array, mx.array]:
"""Concatenate non-overlapping ``[start, end)`` key/value slices along seq."""
if not ranges:
raise ValueError("ranges must be non-empty")
if len(ranges) == 1:
s0, e0 = ranges[0]
return k[:, :, s0:e0, :], v[:, :, s0:e0, :]
k_parts = [k[:, :, s:e, :] for s, e in ranges]
v_parts = [v[:, :, s:e, :] for s, e in ranges]
return mx.concatenate(k_parts, axis=2), mx.concatenate(v_parts, axis=2)
def _key_ranges_for_block(
qs: int,
qe: int,
seq_len: int,
half: int,
sink: int,
) -> list[tuple[int, int]]:
"""Return ordered, non-overlapping key ranges for a query block ``[qs, qe)``.
Symmetric local window over every query in the block, plus global sinks
``[0, sink)``. Overlap is merged into a single contiguous range when
possible so we avoid double-counting sink tokens.
"""
local_start = max(0, qs - half)
# Last query index is ``qe - 1``; its right edge is ``qe - 1 + half + 1 = qe + half``.
local_end = min(seq_len, qe + half)
if local_start >= local_end:
# Degenerate (should not happen for valid qs < qe); fall back to sinks only.
if sink > 0:
return [(0, min(sink, seq_len))]
raise ValueError(f"empty local key range for query block [{qs}, {qe})")
if sink <= 0:
return [(local_start, local_end)]
sink_end = min(sink, seq_len)
if local_start <= sink_end:
# Sinks abut or overlap the local window — one contiguous slice from 0.
return [(0, max(local_end, sink_end))]
# Gap between sinks and local window: two slices, concat at SDPA time.
return [(0, sink_end), (local_start, local_end)]
def windowed_attention(
q: mx.array,
k: mx.array,
v: mx.array,
window: int,
sink: int = 0,
scale: float | None = None,
*,
chunk_size: int | None = None,
) -> mx.array:
"""
Apply symmetric sliding-window self-attention with optional global sink positions.
Parameters:
q (mx.array): Query tensor shaped `(B, H, S, D)`.
k (mx.array): Key tensor shaped `(B, H, S, D)`.
v (mx.array): Value tensor shaped `(B, H, S, D)`.
window (int): Symmetric attention window width in tokens; must be at least 1.
sink (int): Number of leading key positions available to every query; must
be between 0 and the sequence length.
scale (Optional[float]): Softmax scale. Defaults to `1 / sqrt(D)`.
chunk_size (Optional[int]): Query block length used for chunked processing.
Defaults to the smaller of `window` and 512.
Returns:
mx.array: Attention output with the same shape as `q`.
Raises:
ValueError: If the inputs or attention parameters are invalid.
RuntimeError: If a query block has no available keys.
"""
_, _, seq_len, d = _validate_qkv(q, k, v)
if window < 1:
raise ValueError(f"window must be >= 1, got {window}")
if sink < 0:
raise ValueError(f"sink must be >= 0, got {sink}")
if sink > seq_len:
raise ValueError(f"sink ({sink}) cannot exceed sequence length ({seq_len})")
sc = _default_scale(d, scale)
half = window // 2
# When the requested window is at least the sequence length, every query can
# see every key under a symmetric policy — fall back to one dense SDPA.
# (Sinks are redundant once the full key set is used.)
if window >= seq_len:
return mx.fast.scaled_dot_product_attention(q, k, v, scale=sc)
chunk = min(window, 512) if chunk_size is None else int(chunk_size)
if chunk < 1:
raise ValueError(f"chunk_size must be >= 1, got {chunk}")
chunk = min(chunk, seq_len)
outputs: list[mx.array] = []
for qs in range(0, seq_len, chunk):
qe = min(seq_len, qs + chunk)
q_block = q[:, :, qs:qe, :]
ranges = _key_ranges_for_block(qs, qe, seq_len, half, sink)
k_block, v_block = _concat_kv_slices(k, v, ranges)
if k_block.shape[2] == 0:
raise RuntimeError(f"empty key set for query block [{qs}, {qe}) with window={window}, sink={sink}")
query_positions = mx.arange(qs, qe)[:, None]
key_positions = mx.array([position for start, end in ranges for position in range(start, end)])[None, :]
mask = mx.abs(query_positions - key_positions) <= half
if sink > 0:
mask = mask | (key_positions < sink)
out_block = mx.fast.scaled_dot_product_attention(q_block, k_block, v_block, scale=sc, mask=mask)
outputs.append(out_block)
return mx.concatenate(outputs, axis=2)
+19
View File
@@ -333,6 +333,7 @@ class GELUApprox(nn.Module):
self,
in_features: int,
out_features: int,
bias: bool = True,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
@@ -340,6 +341,7 @@ class GELUApprox(nn.Module):
self.proj = ReplicatedLinear(
in_features,
out_features,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.fc_in",
)
@@ -357,6 +359,7 @@ class FeedForward(nn.Module):
dim: int,
dim_out: int,
mult: int = 4,
bias: bool = True,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
@@ -365,12 +368,14 @@ class FeedForward(nn.Module):
project_in = GELUApprox(
dim,
inner_dim,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.ffn",
)
project_out = ReplicatedLinear(
inner_dim,
dim_out,
bias=bias,
quant_config=quant_config,
prefix=f"{prefix}.ffn.fc_out",
)
@@ -1206,6 +1211,9 @@ class TransformerConfig:
# LTX-2.3 gated extensions (default OFF == LTX-2.0 behavior).
apply_gated_attention: bool = False
cross_attention_adaln: bool = False
# FFN bias is per stream: a checkpoint may drop it on video while keeping
# it on audio. Default True preserves existing behavior.
ff_bias: bool = True
class LTXDistributedAttention(DistributedAttention):
@@ -1846,6 +1854,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.ff = FeedForward(
video.dim,
dim_out=video.dim,
bias=video.ff_bias,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{idx}",
)
@@ -1884,6 +1893,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
self.audio_ff = FeedForward(
audio.dim,
dim_out=audio.dim,
bias=audio.ff_bias,
quant_config=quant_config,
prefix=f"{prefix}.blocks.{idx}.audio",
)
@@ -2353,6 +2363,8 @@ class LTXModel(torch.nn.Module):
cross_attention_adaln: bool = False,
caption_proj_before_connector: bool = False,
apply_gated_attention: bool = False,
ff_bias: bool = True,
audio_ff_bias: bool = True,
stg_block_idx: int = 29,
use_distributed_attention: bool = False,
quant_config: QuantizationConfig | None = None,
@@ -2364,6 +2376,9 @@ class LTXModel(torch.nn.Module):
self.cross_attention_adaln = cross_attention_adaln
self.caption_proj_before_connector = caption_proj_before_connector
self.apply_gated_attention = apply_gated_attention
# Per-stream FFN bias; default True preserves existing behavior.
self.ff_bias = ff_bias
self.audio_ff_bias = audio_ff_bias
self.stg_block_idx = stg_block_idx
self.use_middle_indices_grid = use_middle_indices_grid
self.rope_type = rope_type
@@ -2584,6 +2599,7 @@ class LTXModel(torch.nn.Module):
context_dim=cross_attention_dim,
apply_gated_attention=self.apply_gated_attention,
cross_attention_adaln=self.cross_attention_adaln,
ff_bias=self.ff_bias,
) if self.model_type.is_video_enabled() else None)
audio_config = (TransformerConfig(
dim=self.audio_inner_dim,
@@ -2592,6 +2608,7 @@ class LTXModel(torch.nn.Module):
context_dim=audio_cross_attention_dim,
apply_gated_attention=self.apply_gated_attention,
cross_attention_adaln=self.cross_attention_adaln,
ff_bias=self.audio_ff_bias,
) if self.model_type.is_audio_enabled() else None)
self.use_distributed_attention = use_distributed_attention
self.transformer_blocks = torch.nn.ModuleList([
@@ -2778,6 +2795,8 @@ class LTX2Transformer3DModel(BaseDiT):
cross_attention_adaln=arch.cross_attention_adaln,
caption_proj_before_connector=arch.caption_proj_before_connector,
apply_gated_attention=arch.apply_gated_attention,
ff_bias=arch.ff_bias,
audio_ff_bias=arch.audio_ff_bias,
stg_block_idx=arch.stg_block_idx,
use_distributed_attention=use_distributed_attention,
quant_config=config.quant_config,
+27 -137
View File
@@ -10,7 +10,6 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo import envs
from fastvideo.attention import DistributedAttention
from fastvideo.attention.layer import DistributedAttention_VSA
from fastvideo.attention.selector import get_attn_backend
@@ -24,50 +23,12 @@ from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.quantization import QuantizationConfig
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_fusions import (
HAVE_TRITON,
fused_qknorm_rope,
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
minimax_h3_swiglu,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
from fastvideo.utils import get_compute_dtype
logger = init_logger(__name__)
MINIMAX_H3_MODALITY_NUM = 3
_CFG = MiniMaxH3Config()
_MINIMAX_H3_FUSION_NAMES = frozenset({"modulate", "qknorm_rope", "swiglu"})
def _enabled_minimax_h3_fusions(value: str | None = None) -> frozenset[str]:
"""Parse the independently switchable inference fusion set."""
raw = envs.FASTVIDEO_MINIMAX_H3_FUSIONS if value is None else value
normalized = raw.strip().lower()
if normalized in {"", "0", "none"}:
return frozenset()
if normalized in {"1", "all"}:
return _MINIMAX_H3_FUSION_NAMES
enabled = frozenset(item.strip() for item in normalized.split(",") if item.strip())
unknown = enabled - _MINIMAX_H3_FUSION_NAMES
if unknown:
supported = ",".join(sorted(_MINIMAX_H3_FUSION_NAMES))
raise ValueError(f"Unknown MiniMax H3 fusion(s) {sorted(unknown)}; expected a subset of {supported}.")
return enabled
def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool:
"""Triton kernels are inference-only and stay outside Dynamo capture.
The ``HAVE_TRITON`` check makes the eager fallback exact: on a CUDA build
whose Triton failed to import, an enabled fusion falls back instead of
hitting the strict wrappers' hard RuntimeError mid-forward.
"""
return (HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() and not torch.compiler.is_compiling())
class MiniMaxH3RotaryPosEmbed(nn.Module):
@@ -101,7 +62,6 @@ class MiniMaxH3FeedForward(nn.Module):
ffn_dim: int,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.fc_in = ReplicatedLinear(
@@ -118,15 +78,11 @@ class MiniMaxH3FeedForward(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.fc_out",
)
self.fuse_swiglu = fuse_swiglu
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc_in(hidden_states)
if self.fuse_swiglu and _can_run_minimax_h3_fusion(hidden_states):
hidden_states = minimax_h3_swiglu(hidden_states)
else:
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
hidden_states, _ = self.fc_out(hidden_states)
return hidden_states
@@ -143,7 +99,6 @@ class MiniMaxH3Attention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...],
quant_config: QuantizationConfig | None,
prefix: str,
fuse_qknorm_rope: bool = False,
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -179,7 +134,6 @@ class MiniMaxH3Attention(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.to_out",
)
self.fuse_qknorm_rope = fuse_qknorm_rope
# VSA carries a learned gate on its pooled-compression branch. The H3
# checkpoint has no such weight, so the loader zero-initializes it
# (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
@@ -257,18 +211,11 @@ class MiniMaxH3Attention(nn.Module):
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))
if (self.fuse_qknorm_rope and rotary_emb is not None and _can_run_minimax_h3_fusion(query)):
cos, sin = rotary_emb
cos = cos.to(query.dtype)
sin = sin.to(query.dtype)
query = fused_qknorm_rope(query, self.norm_q.weight, cos, sin, self.norm_q.eps)
key = fused_qknorm_rope(key, self.norm_k.weight, cos, sin, self.norm_k.eps)
else:
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
# branch cannot express. Apply it above, then pass no RoPE here.
@@ -450,9 +397,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
quant_config: QuantizationConfig | None,
prefix: str,
adaln_apply_silu: bool = True,
fuse_modulate: bool = False,
fuse_qknorm_rope: bool = False,
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
@@ -464,7 +408,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
supported_attention_backends,
quant_config,
prefix=f"{prefix}.attn",
fuse_qknorm_rope=fuse_qknorm_rope,
)
self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
self.ff = MiniMaxH3FeedForward(
@@ -472,7 +415,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
ffn_dim,
quant_config=quant_config,
prefix=f"{prefix}.ff",
fuse_swiglu=fuse_swiglu,
)
self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
time_embed_dim,
@@ -481,7 +423,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
prefix=f"{prefix}.adaln_proj",
apply_silu=adaln_apply_silu,
)
self.fuse_modulate = fuse_modulate
def forward(
self,
@@ -494,39 +435,19 @@ class MiniMaxH3TransformerBlock(nn.Module):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
t.to(hidden_states.dtype) for t in self.adaln_proj(temb))
use_modulate_fusion = self.fuse_modulate and _can_run_minimax_h3_fusion(hidden_states)
if use_modulate_fusion:
norm_hidden_states = fused_rmsnorm_modulate(
hidden_states,
self.norm1.weight,
scale_msa,
shift_msa,
adaln_indices,
self.norm1.eps,
)
else:
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
residual = hidden_states
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
if use_modulate_fusion:
hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate(
hidden_states,
attention_output,
gate_msa,
self.norm2.weight,
scale_mlp,
shift_mlp,
adaln_indices,
self.norm2.eps,
)
else:
hidden_states = hidden_states + gate_msa.index_select(0, adaln_indices) * attention_output
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)
hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attention_output
residual = hidden_states
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)
feed_forward_output = self.ff(norm_hidden_states)
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
class MiniMaxH3Transformer3DModel(BaseDiT):
@@ -572,17 +493,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
super().__init__(config, hf_config)
arch = config.arch_config
self.enabled_fusions = _enabled_minimax_h3_fusions()
if self.enabled_fusions:
if HAVE_TRITON:
logger.info(
"MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled and "
"torch.compile-captured forwards fall back to eager).",
",".join(sorted(self.enabled_fusions)))
else:
logger.warning(
"FASTVIDEO_MINIMAX_H3_FUSIONS requested %s but Triton is unavailable; "
"every forward stays on the eager path.", ",".join(sorted(self.enabled_fusions)))
sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1
if arch.num_attention_heads % sp_world_size:
raise ValueError(f"MiniMax H3 attention heads ({arch.num_attention_heads}) must be divisible by "
@@ -636,7 +546,7 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
"parameter, but factorized AdaLN weights are pinned to FP16 "
"(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank "
"checkpoint instead, then re-fit the basis with "
"scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
"tools/minimax_h3/fit_adaln_basis.py.")
adaln_dim = self.adaln_rank or arch.time_embed_dim
self.adaln_basis = ReplicatedLinear(
arch.time_embed_dim,
@@ -680,9 +590,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
config.quant_config,
prefix=f"{config.prefix}.transformer_blocks.{index}",
adaln_apply_silu=self.adaln_rank is None,
fuse_modulate="modulate" in self.enabled_fusions,
fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions,
fuse_swiglu="swiglu" in self.enabled_fusions,
) for index in range(arch.num_layers)
])
self.norm_out = MiniMaxH3AdaLayerNormOut(
@@ -709,20 +616,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
)
self.__post_init__()
def prepare_for_compile(self) -> None:
"""Pipeline hook, called once right before torch.compile wraps the blocks.
Dynamo capture traces the eager branch of every fusion guard, so an
enabled ``FASTVIDEO_MINIMAX_H3_FUSIONS`` set is silently inert inside
compiled block forwards (H3 compiles per-block by default). Say so
once instead of leaving the flag looking active.
"""
if self.enabled_fusions:
logger.warning(
"torch.compile is enabled for MiniMax H3, so the requested inference fusions (%s) are "
"inert inside compiled block forwards; the compiled eager path runs instead.",
",".join(sorted(self.enabled_fusions)))
def materialize_non_persistent_buffers(
self,
device: torch.device,
@@ -841,17 +734,14 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
rotary_emb = (rotary_cos, rotary_sin)
# 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}"):
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
for block in self.transformer_blocks:
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
packed_hidden_states = self.norm_out(
packed_hidden_states,
@@ -1,20 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Inference-only MiniMax H3 fusions adapted from NVlabs/Sana Sol-Engine.
Source: https://github.com/NVlabs/Sana/tree/sol-engine/models/minimax_h3/GB200
"""
from .modulation import (
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
)
from .qknorm_rope import HAVE_TRITON, fused_qknorm_rope
from .swiglu import minimax_h3_swiglu
__all__ = [
"HAVE_TRITON",
"fused_qknorm_rope",
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
"minimax_h3_swiglu",
]
@@ -1,302 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax H3 RMSNorm and row-indexed modulation fusions."""
from __future__ import annotations
import math
import torch
try:
import triton
import triton.language as tl
except ImportError as exc: # pragma: no cover - depends on the runtime image
triton = None
tl = None
_TRITON_IMPORT_ERROR: ImportError | None = exc
else:
_TRITON_IMPORT_ERROR = None
__all__ = [
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
]
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
_rmsnorm_modulate_kernel = None
_residual_gate_rmsnorm_modulate_kernel = None
if triton is not None:
@triton.jit
def _rmsnorm_modulate_kernel(
out_ptr,
x_ptr,
weight_ptr,
scale_ptr,
shift_ptr,
index_ptr,
n_cols,
n_index,
eps,
stride_x_row,
stride_scale_row,
stride_shift_row,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK)
mask = cols < n_cols
x_offsets = row * stride_x_row + cols
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
x = tl.load(x_ptr + x_offsets, mask=mask, other=0.0).to(tl.float32)
variance = tl.sum(x * x, axis=0) / n_cols
normed = x * tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
scale = tl.load(
scale_ptr + table_row * stride_scale_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
shift = tl.load(
shift_ptr + table_row * stride_shift_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
output = normed * weight * (1.0 + scale) + shift
tl.store(out_ptr + x_offsets, output.to(out_ptr.dtype.element_ty), mask=mask)
@triton.jit
def _residual_gate_rmsnorm_modulate_kernel(
hidden_out_ptr,
normed_out_ptr,
residual_ptr,
branch_ptr,
gate_ptr,
weight_ptr,
scale_ptr,
shift_ptr,
index_ptr,
n_cols,
n_index,
eps,
stride_input_row,
stride_gate_row,
stride_scale_row,
stride_shift_row,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK)
mask = cols < n_cols
input_offsets = row * stride_input_row + cols
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
residual = tl.load(residual_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
branch = tl.load(branch_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(
gate_ptr + table_row * stride_gate_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
hidden = residual + gate * branch
tl.store(hidden_out_ptr + input_offsets, hidden.to(hidden_out_ptr.dtype.element_ty), mask=mask)
variance = tl.sum(hidden * hidden, axis=0) / n_cols
normed = hidden * tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
scale = tl.load(
scale_ptr + table_row * stride_scale_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
shift = tl.load(
shift_ptr + table_row * stride_shift_row + cols,
mask=mask,
other=0.0,
).to(tl.float32)
output = normed * weight * (1.0 + scale) + shift
tl.store(normed_out_ptr + input_offsets, output.to(normed_out_ptr.dtype.element_ty), mask=mask)
def _validate_contract(
x: torch.Tensor,
weight: torch.Tensor,
tables: tuple[torch.Tensor, ...],
index: torch.Tensor,
eps: float,
) -> None:
if x.ndim < 2:
raise ValueError(f"x must have shape (..., sequence_length, hidden_size), got {tuple(x.shape)}.")
if x.numel() == 0 or x.shape[-1] == 0:
raise ValueError("x must not be empty.")
if x.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"x must use float16, bfloat16, or float32, got {x.dtype}.")
hidden_size = x.shape[-1]
sequence_length = x.shape[-2]
if weight.shape != (hidden_size, ):
raise ValueError(f"weight must have shape ({hidden_size},), got {tuple(weight.shape)}.")
if weight.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"weight must use float16, bfloat16, or float32, got {weight.dtype}.")
if index.ndim != 1 or index.numel() != sequence_length:
raise ValueError(
f"index must have shape ({sequence_length},) so it can wrap over batch rows, got {tuple(index.shape)}."
)
if index.dtype not in (torch.int32, torch.int64):
raise TypeError(f"index must use int32 or int64, got {index.dtype}.")
if not isinstance(eps, (float, int)) or isinstance(eps, bool) or not math.isfinite(eps) or eps <= 0:
raise ValueError(f"eps must be a positive finite number, got {eps!r}.")
table_rows = tables[0].shape[0] if tables and tables[0].ndim == 2 else None
for name, table in zip(("gate", "scale", "shift")[-len(tables):], tables, strict=True):
if table.ndim != 2 or table.shape[1] != hidden_size:
raise ValueError(f"{name} must have shape (table_rows, {hidden_size}), got {tuple(table.shape)}.")
if table.shape[0] == 0 or table.shape[0] != table_rows:
raise ValueError("all modulation tables must have the same non-zero row count.")
if table.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"{name} must use float16, bfloat16, or float32, got {table.dtype}.")
tensors = (x, weight, *tables, index)
if any(tensor.device != x.device for tensor in tensors[1:]):
raise ValueError("x, weight, modulation tables, and index must be on the same device.")
def _validate_residual_branch(residual: torch.Tensor, branch: torch.Tensor) -> None:
if branch.shape != residual.shape:
raise ValueError(f"branch must match residual shape {tuple(residual.shape)}, got {tuple(branch.shape)}.")
if branch.dtype != residual.dtype:
raise TypeError(f"branch dtype must match residual dtype {residual.dtype}, got {branch.dtype}.")
if branch.device != residual.device:
raise ValueError("branch and residual must be on the same device.")
def _require_triton_cuda(x: torch.Tensor) -> None:
if triton is None:
detail = f": {_TRITON_IMPORT_ERROR}" if _TRITON_IMPORT_ERROR is not None else ""
raise RuntimeError(f"MiniMax H3 modulation fusion requires Triton{detail}.")
if x.device.type != "cuda":
raise RuntimeError(f"MiniMax H3 modulation fusion requires CUDA tensors, got device {x.device}.")
def _require_forward_only(*tensors: torch.Tensor) -> None:
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in tensors):
raise RuntimeError("MiniMax H3 modulation fusion is forward-only and does not support autograd.")
def _next_power_of_two(value: int) -> int:
return 1 << (value - 1).bit_length()
def _num_warps(block_size: int) -> int:
if block_size >= 8192:
return 16
if block_size >= 2048:
return 8
return 4
def _row_addressable(table: torch.Tensor) -> torch.Tensor:
return table if table.stride(-1) == 1 else table.contiguous()
def fused_rmsnorm_modulate(
x: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> torch.Tensor:
"""Run RMSNorm and row-indexed modulation in one strict Triton kernel.
``index`` values must lie in ``[0, table_rows)``. Unlike eager
``index_select``, the kernel does not raise on out-of-range values (a
device-side bounds check would synchronize); callers are safe by
construction (``timestep_indices * 3 + token_tags``, SP pads with 0).
"""
_validate_contract(x, weight, (scale, shift), index, eps)
_require_forward_only(x, weight, scale, shift)
_require_triton_cuda(x)
hidden_size = x.shape[-1]
flat_x = x.reshape(-1, hidden_size).contiguous()
weight = weight.contiguous()
scale = _row_addressable(scale)
shift = _row_addressable(shift)
index = index.contiguous()
output = torch.empty_like(flat_x)
block_size = _next_power_of_two(hidden_size)
_rmsnorm_modulate_kernel[(flat_x.shape[0], )](
output,
flat_x,
weight,
scale,
shift,
index,
hidden_size,
index.numel(),
eps,
flat_x.stride(0),
scale.stride(0),
shift.stride(0),
BLOCK=block_size,
num_warps=_num_warps(block_size),
)
return output.view_as(x)
def fused_residual_gate_rmsnorm_modulate(
residual: torch.Tensor,
branch: torch.Tensor,
gate: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Fuse residual update, row-indexed gate, RMSNorm, and modulation.
``index`` values must lie in ``[0, table_rows)``; see
:func:`fused_rmsnorm_modulate` for why the wrapper does not check them.
"""
_validate_residual_branch(residual, branch)
_validate_contract(residual, weight, (gate, scale, shift), index, eps)
_require_forward_only(residual, branch, gate, weight, scale, shift)
_require_triton_cuda(residual)
hidden_size = residual.shape[-1]
flat_residual = residual.reshape(-1, hidden_size).contiguous()
flat_branch = branch.reshape(-1, hidden_size).contiguous()
weight = weight.contiguous()
gate = _row_addressable(gate)
scale = _row_addressable(scale)
shift = _row_addressable(shift)
index = index.contiguous()
hidden = torch.empty_like(flat_residual)
modulated = torch.empty_like(flat_residual)
block_size = _next_power_of_two(hidden_size)
_residual_gate_rmsnorm_modulate_kernel[(flat_residual.shape[0], )](
hidden,
modulated,
flat_residual,
flat_branch,
gate,
weight,
scale,
shift,
index,
hidden_size,
index.numel(),
eps,
flat_residual.stride(0),
gate.stride(0),
scale.stride(0),
shift.stride(0),
BLOCK=block_size,
num_warps=_num_warps(block_size),
)
return hidden.view_as(residual), modulated.view_as(residual)
@@ -1,174 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Fused per-head RMSNorm and partial rotary embedding for MiniMax H3."""
from __future__ import annotations
import math
import torch
try:
import triton
import triton.language as tl
HAVE_TRITON = True
except ImportError: # pragma: no cover - exercised only in environments without Triton
triton = None
tl = None
HAVE_TRITON = False
if HAVE_TRITON:
@triton.jit
def _qknorm_partial_rope_kernel(
out_ptr,
x_ptr,
weight_ptr,
cos_ptr,
sin_ptr,
head_dim,
rotary_dim,
half_rotary_dim,
num_heads,
seq_len,
eps,
BLOCK_SIZE: tl.constexpr,
):
# int64, like the sibling kernels: with int32 program ids,
# ``row * head_dim`` wraps once the flattened input reaches 2**31
# elements (H3's 56 heads x 128 head_dim crosses that at
# batch*seq >= 299_593 tokens per rank) and the loads/stores below
# become out-of-bounds. ``seq_index`` inherits int64 from ``row``.
row = tl.program_id(0).to(tl.int64)
seq_index = (row // num_heads) % seq_len
cols = tl.arange(0, BLOCK_SIZE)
head_mask = cols < head_dim
row_offset = row * head_dim
x = tl.load(x_ptr + row_offset + cols, mask=head_mask, other=0.0).to(tl.float32)
variance = tl.sum(x * x, axis=0) / head_dim
inv_rms = tl.math.rsqrt(variance + eps)
weight = tl.load(weight_ptr + cols, mask=head_mask, other=0.0).to(tl.float32)
normalized = x * inv_rms * weight
rotary_mask = cols < rotary_dim
first_half = cols < half_rotary_dim
partner_col = tl.where(first_half, cols + half_rotary_dim, cols - half_rotary_dim)
partner_x = tl.load(
x_ptr + row_offset + partner_col,
mask=rotary_mask,
other=0.0,
).to(tl.float32)
partner_weight = tl.load(weight_ptr + partner_col, mask=rotary_mask, other=0.0).to(tl.float32)
partner_normalized = partner_x * inv_rms * partner_weight
rotated = tl.where(first_half, -partner_normalized, partner_normalized)
table_offset = seq_index * rotary_dim + cols
cos = tl.load(cos_ptr + table_offset, mask=rotary_mask, other=1.0).to(tl.float32)
sin = tl.load(sin_ptr + table_offset, mask=rotary_mask, other=0.0).to(tl.float32)
rotary_output = normalized * cos + rotated * sin
output = tl.where(rotary_mask, rotary_output, normalized)
tl.store(out_ptr + row_offset + cols, output.to(out_ptr.dtype.element_ty), mask=head_mask)
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
def _validate_inputs(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> tuple[int, int, int, int, int]:
for name, tensor in (("x", x), ("weight", weight), ("cos", cos), ("sin", sin)):
if not isinstance(tensor, torch.Tensor):
raise TypeError(f"{name} must be a torch.Tensor, got {type(tensor).__name__}")
if x.ndim != 4:
raise ValueError(f"x must have shape (batch, seq, heads, head_dim), got {tuple(x.shape)}")
batch, seq_len, num_heads, head_dim = x.shape
if min(batch, seq_len, num_heads, head_dim) <= 0:
raise ValueError(f"x dimensions must all be positive, got {tuple(x.shape)}")
if weight.shape != (head_dim, ):
raise ValueError(f"weight must have shape ({head_dim},), got {tuple(weight.shape)}")
if cos.ndim != 2:
raise ValueError(f"cos must have shape (seq, rotary_dim), got {tuple(cos.shape)}")
if sin.shape != cos.shape:
raise ValueError(f"sin must match cos shape {tuple(cos.shape)}, got {tuple(sin.shape)}")
if cos.shape[0] != seq_len:
raise ValueError(f"cos/sin sequence length must be {seq_len}, got {cos.shape[0]}")
rotary_dim = cos.shape[1]
if rotary_dim <= 0:
raise ValueError(f"rotary_dim must be positive, got {rotary_dim}")
if rotary_dim > head_dim:
raise ValueError(f"rotary_dim must not exceed head_dim, got rotary_dim={rotary_dim}, head_dim={head_dim}")
if rotary_dim % 2:
raise ValueError(f"rotary_dim must be even, got {rotary_dim}")
if x.dtype not in _SUPPORTED_DTYPES:
raise TypeError(f"x dtype must be float16, bfloat16, or float32, got {x.dtype}")
for name, tensor in (("weight", weight), ("cos", cos), ("sin", sin)):
if tensor.dtype != x.dtype:
raise TypeError(f"{name} dtype must match x dtype {x.dtype}, got {tensor.dtype}")
if tensor.device != x.device:
raise ValueError(f"{name} device must match x device {x.device}, got {tensor.device}")
if not isinstance(eps, (float, int)) or not math.isfinite(float(eps)) or eps <= 0:
raise ValueError(f"eps must be a positive finite number, got {eps!r}")
return batch, seq_len, num_heads, head_dim, rotary_dim
def fused_qknorm_rope(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
"""Run per-head RMSNorm and partial RoPE in one Sol-Engine-style kernel.
RMSNorm reduction and RoPE arithmetic stay in FP32 registers until the
final store. Triton's reduction order and the absence of eager's BF16
intermediate materializations can produce small, expected rounding drift.
Row offsets are computed in int64, so inputs beyond 2**31 total elements
(about 300k tokens per rank at H3's 56 heads x 128 head_dim) address
correctly.
"""
batch, seq_len, num_heads, head_dim, rotary_dim = _validate_inputs(x, weight, cos, sin, eps)
if not weight.is_contiguous():
raise ValueError("weight must be contiguous")
if not cos.is_contiguous() or not sin.is_contiguous():
raise ValueError("cos and sin must be contiguous (seq, rotary_dim) tables")
if not x.is_cuda:
raise RuntimeError("fused_qknorm_rope requires CUDA tensors")
if not HAVE_TRITON:
raise RuntimeError("fused_qknorm_rope requires Triton")
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in (x, weight, cos, sin)):
raise RuntimeError("fused_qknorm_rope is inference-only and does not implement autograd")
flat_x = x.reshape(-1, head_dim).contiguous()
flat_out = torch.empty_like(flat_x)
block_size = 1 << (head_dim - 1).bit_length()
_qknorm_partial_rope_kernel[(flat_x.shape[0], )](
flat_out,
flat_x,
weight,
cos,
sin,
head_dim,
rotary_dim,
rotary_dim // 2,
num_heads,
seq_len,
eps,
BLOCK_SIZE=block_size,
num_warps=4,
)
return flat_out.view(batch, seq_len, num_heads, head_dim)
__all__ = ["HAVE_TRITON", "fused_qknorm_rope"]
@@ -1,104 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""MiniMax H3's value-first packed SwiGLU fusion."""
from __future__ import annotations
import torch
try:
import triton
import triton.language as tl
HAVE_TRITON = True
except ImportError: # pragma: no cover - exercised only in environments without Triton
triton = None
tl = None
HAVE_TRITON = False
def _validate_input(x: torch.Tensor) -> int:
if x.ndim == 0:
raise ValueError("MiniMax H3 SwiGLU expects at least one dimension")
packed_width = x.shape[-1]
if packed_width == 0 or packed_width % 2 != 0:
raise ValueError(
"MiniMax H3 SwiGLU expects a positive even last dimension containing packed (value, gate) halves, "
f"got {packed_width}"
)
if not x.is_floating_point():
raise TypeError(f"MiniMax H3 SwiGLU expects a floating-point tensor, got {x.dtype}")
return packed_width // 2
if HAVE_TRITON:
@triton.jit
def _minimax_h3_swiglu_kernel(
out_ptr,
x_ptr,
ffn_dim,
stride_in_row,
stride_out_row,
BLOCK_SIZE: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK_SIZE)
mask = cols < ffn_dim
value = tl.load(x_ptr + row * stride_in_row + cols, mask=mask, other=0.0).to(tl.float32)
gate = tl.load(x_ptr + row * stride_in_row + ffn_dim + cols, mask=mask, other=0.0).to(tl.float32)
# Match Sol-Engine: keep the complete SwiGLU expression in FP32 and
# convert only the final output store.
out = value * (gate * tl.sigmoid(gate))
tl.store(out_ptr + row * stride_out_row + cols, out.to(out_ptr.dtype.element_ty), mask=mask)
else:
_minimax_h3_swiglu_kernel = None
def _num_warps(block_size: int) -> int:
if block_size >= 8192:
return 16
if block_size >= 2048:
return 8
return 4
def minimax_h3_swiglu(x: torch.Tensor) -> torch.Tensor:
"""Run the forward-only Triton fusion over an H3 ``(..., 2 * ffn_dim)`` input.
This is intentionally a strict kernel wrapper: callers own fallback policy and
must only invoke it for a supported CUDA inference path.
"""
ffn_dim = _validate_input(x)
if not x.is_cuda:
raise ValueError("MiniMax H3 fused SwiGLU requires a CUDA tensor")
if x.dtype not in (torch.float16, torch.bfloat16, torch.float32):
raise TypeError(f"MiniMax H3 fused SwiGLU supports float16, bfloat16, and float32, got {x.dtype}")
if torch.is_grad_enabled() and x.requires_grad:
raise RuntimeError("MiniMax H3 fused SwiGLU is forward-only and does not implement autograd")
if _minimax_h3_swiglu_kernel is None:
raise RuntimeError("MiniMax H3 fused SwiGLU requires Triton")
packed_width = x.shape[-1]
flat = x.reshape(-1, packed_width).contiguous()
output_shape = (*x.shape[:-1], ffn_dim)
if flat.shape[0] == 0:
return torch.empty(output_shape, dtype=x.dtype, device=x.device)
out = torch.empty((flat.shape[0], ffn_dim), dtype=x.dtype, device=x.device)
block_size = triton.next_power_of_2(ffn_dim)
_minimax_h3_swiglu_kernel[(flat.shape[0],)](
out,
flat,
ffn_dim,
flat.stride(0),
out.stride(0),
BLOCK_SIZE=block_size,
num_warps=_num_warps(block_size),
)
return out.view(output_shape)
__all__ = ["HAVE_TRITON", "minimax_h3_swiglu"]
+8 -8
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from dataclasses import field
from typing import Any, Generic, TypeVar
import torch
from torch import nn
@@ -9,16 +8,11 @@ from torch import nn
from fastvideo.configs.models.encoders import (BaseEncoderOutput, ImageEncoderConfig, TextEncoderConfig)
from fastvideo.platforms import AttentionBackendEnum
TextEncoderOutputT = TypeVar("TextEncoderOutputT")
class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
"""Base for native encoders with a model-specific forward output contract."""
class TextEncoder(nn.Module, ABC):
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = TextEncoderConfig()._supported_attention_backends
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -29,7 +23,13 @@ class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
@abstractmethod
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
def forward(self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput:
pass
@property
+40 -4
View File
@@ -7,7 +7,7 @@ from typing import Any, Iterable
import torch
from torch import nn
from transformers import AutoTokenizer, Gemma3ForConditionalGeneration
from transformers import AutoModelForImageTextToText, AutoTokenizer
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
from fastvideo.models.encoders.base import TextEncoder
@@ -398,7 +398,7 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
self.gemma_model_path = arch.gemma_model_path
self.gemma_dtype = arch.gemma_dtype
self.padding_side = arch.padding_side
self._gemma_model: Gemma3ForConditionalGeneration | None = None
self._gemma_model: nn.Module | None = None
def named_parameters(self, prefix: str = "", recurse: bool = True):
for name, param in super().named_parameters(prefix=prefix, recurse=recurse):
@@ -436,17 +436,22 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
)
@property
def gemma_model(self) -> Gemma3ForConditionalGeneration:
def gemma_model(self) -> nn.Module:
if self._gemma_model is None:
gemma_path = self.gemma_model_path
if not gemma_path:
raise ValueError("gemma_model_path must be set (expected text_encoder/gemma).")
dtype = getattr(torch, self.gemma_dtype, torch.bfloat16)
self._gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
# Resolve the class from the root config's own ``model_type``
# instead of hardcoding one: a checkpoint may declare its own
# encoder pairing. The auto class still resolves to the previously
# hardcoded class for the roots shipped so far.
self._gemma_model = AutoModelForImageTextToText.from_pretrained(
gemma_path,
local_files_only=True,
torch_dtype=dtype,
)
self._sync_arch_from_gemma_config(self._gemma_model.config)
# Configure model-level attention implementation when using TORCH_SDPA.
# Note: torch.backends.cuda.enable_*_sdp() settings should be configured
# at application/pipeline initialization level, not here, to avoid
@@ -464,6 +469,33 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
self._gemma_model.eval()
return self._gemma_model
def _sync_arch_from_gemma_config(self, gemma_config: Any) -> None:
"""Take language-model geometry from the loaded Gemma config.
A multimodal root nests it under ``text_config``; a text-only config
carries it at the root. The dataclass defaults hold for every encoder
root shipped so far, so this is not a behavior change -- it just stops
the defaults from being authoritative.
The feature-extractor linears are sized in ``__init__`` from
``feature_extractor_in_features``, before the Gemma weights load, so a
geometry mismatch is raised here rather than surfacing later as an
opaque matmul shape error.
"""
text_config = getattr(gemma_config, "text_config", gemma_config)
arch = self.config.arch_config
arch.hidden_size = text_config.hidden_size
arch.num_hidden_layers = text_config.num_hidden_layers
# +1: the stacked hidden states include the embedding output.
expected_in_features = arch.hidden_size * (arch.num_hidden_layers + 1)
if expected_in_features != arch.feature_extractor_in_features:
raise ValueError(
"Gemma geometry does not match the configured feature "
f"extractor: loaded config gives hidden_size={arch.hidden_size} "
f"x (num_hidden_layers={arch.num_hidden_layers} + 1) = "
f"{expected_in_features}, but feature_extractor_in_features is "
f"{arch.feature_extractor_in_features}."
)
def _run_feature_extractor(
self,
hidden_states: tuple[torch.Tensor, ...],
@@ -670,6 +702,10 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
name = "feature_extractor_linear.aggregate_embed.weight"
elif name.startswith("video_connector."):
name = name.replace("video_connector.", "embeddings_connector.", 1)
elif name.startswith("video_embeddings_connector."):
# Some checkpoints name the video connector sub-tree after its
# modality; the module is just ``embeddings_connector``.
name = name.replace("video_embeddings_connector.", "embeddings_connector.", 1)
elif name.startswith("audio_connector."):
name = name.replace("audio_connector.", "audio_embeddings_connector.", 1)
if name not in params_dict:
@@ -1,453 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Serialized block-FP8 execution for the MiniMax-H3 Qwen3-VL encoder."""
from typing import Any
import torch
from torch import nn
from torch.nn.parameter import Parameter
try:
import triton
import triton.language as tl
except ImportError:
triton = None
tl = None
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.linear import LinearBase, LinearMethodBase
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.layers.quantization.fp8_config import FP8_DTYPE
from fastvideo.models.utils import set_weight_attrs
class MiniMaxH3SerializedFP8Config(QuantizationConfig):
"""Serialized 128x128 block-FP8 contract for the H3 text encoder."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
if weight_block_size != (128, 128):
raise ValueError("MiniMax-H3 serialized FP8 requires weight_block_size=[128, 128], "
f"got {list(weight_block_size)}")
self.weight_block_size = weight_block_size
self.is_checkpoint_fp8_serialized = True
self.activation_scheme = "dynamic"
@classmethod
def get_name(cls) -> str:
return "fp8"
@classmethod
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
return [torch.bfloat16]
@classmethod
def get_min_capability(cls) -> int:
return 100
@staticmethod
def get_config_filenames() -> list[str]:
return []
@classmethod
def from_config(cls, config: dict[str, Any]) -> "MiniMaxH3SerializedFP8Config":
quant_method = str(config.get("quant_method", "")).lower()
if quant_method != "fp8":
raise ValueError(f"MiniMax-H3 only supports serialized FP8 text-encoder checkpoints, got {quant_method!r}")
if str(config.get("activation_scheme", "")).lower() != "dynamic":
raise ValueError("MiniMax-H3 serialized FP8 requires dynamic activation quantization")
if str(config.get("fmt", "e4m3")).lower() not in ("e4m3", "float8_e4m3fn"):
raise ValueError(f"MiniMax-H3 serialized FP8 requires E4M3 weights, got {config.get('fmt')!r}")
block_size = config.get("weight_block_size")
if not isinstance(block_size, list | tuple) or len(block_size) != 2:
raise ValueError("MiniMax-H3 serialized FP8 requires a two-dimensional weight_block_size")
ignored_layers = config.get("modules_to_not_convert", config.get("ignored_layers", []))
if not isinstance(ignored_layers, list | tuple):
raise ValueError("MiniMax-H3 serialized FP8 modules_to_not_convert must be a sequence")
language_exclusions = [
name for name in ignored_layers
if isinstance(name, str) and (name.startswith("language_model.") or ".language_model." in name)
]
if language_exclusions:
raise ValueError("MiniMax-H3 does not support partially quantized language stacks; "
f"ignored language layers: {language_exclusions[:3]}")
if not any(isinstance(name, str) and "visual" in name for name in ignored_layers):
raise ValueError("MiniMax-H3 serialized FP8 requires the vision stack to be listed in "
"modules_to_not_convert")
return cls((int(block_size[0]), int(block_size[1])))
def validate_runtime(self, device: torch.device) -> None:
if device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires a CUDA device; "
f"got {device.type!r}")
capability = torch.cuda.get_device_capability(device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < self.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{self.get_min_capability()} or newer, got sm{capability_number}")
if capability[0] not in (10, 12):
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 currently adapts SGLang's Blackwell "
f"FlashInfer path; got unsupported sm{capability_number}")
_require_sglang_per_token_group_fp8_quantization()
_get_flashinfer_groupwise_fp8_gemm()
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
if isinstance(layer, LinearBase) and ".language_model.layers." in prefix:
return MiniMaxH3SerializedFP8LinearMethod(self.weight_block_size)
return None
# Copyright 2024 SGLang Team
# Licensed under the Apache License, Version 2.0.
# Adapted from SGLang's per-token-group quantization kernels and Blackwell
# FlashInfer dispatch at commit f99c62063c7dcfcd06784b885dc08cb52cf23865:
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/kernels/ops/quantization/fp8_kernel.py
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/srt/layers/quantization/fp8_utils.py
if triton is not None:
@triton.jit
def _h3_per_token_group_quant_fp8_row_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
scale_ptr += group_id
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
@triton.jit
def _h3_per_token_group_quant_fp8_column_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
input_columns,
scale_column_stride,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
groups_per_row = input_columns // group_size
scale_column = group_id % groups_per_row
scale_row = group_id // groups_per_row
scale_ptr += scale_column * scale_column_stride + scale_row
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
else:
_h3_per_token_group_quant_fp8_row_major = None
_h3_per_token_group_quant_fp8_column_major = None
def _require_sglang_per_token_group_fp8_quantization() -> None:
if (triton is None or _h3_per_token_group_quant_fp8_row_major is None
or _h3_per_token_group_quant_fp8_column_major is None):
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires Triton for SGLang-compatible "
"per-token-group activation quantization")
def _sglang_per_token_group_quant_fp8(
input_tensor: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""SGLang-compatible dynamic FP8 quantization for contiguous 2-D activations."""
_require_sglang_per_token_group_fp8_quantization()
if input_tensor.ndim != 2:
raise ValueError(f"per-token-group FP8 quantization expects 2-D input, got {input_tensor.ndim}-D")
if not input_tensor.is_contiguous():
raise ValueError("per-token-group FP8 quantization requires contiguous input")
if input_tensor.shape[-1] % group_size:
raise ValueError(f"activation width {input_tensor.shape[-1]} is not divisible by group_size={group_size}")
quantized = torch.empty_like(input_tensor, dtype=FP8_DTYPE)
rows, columns = input_tensor.shape
groups_per_row = columns // group_size
if column_major_scales:
scales = torch.empty(
(groups_per_row, rows),
device=input_tensor.device,
dtype=torch.float32,
).permute(1, 0)
else:
scales = torch.empty(
(rows, groups_per_row),
device=input_tensor.device,
dtype=torch.float32,
)
if rows:
num_groups = input_tensor.numel() // group_size
block = triton.next_power_of_2(group_size)
num_warps = min(max(block // 256, 1), 8)
if column_major_scales:
_h3_per_token_group_quant_fp8_column_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
columns,
scales.stride(1),
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
else:
_h3_per_token_group_quant_fp8_row_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
return quantized, scales
def _get_flashinfer_groupwise_fp8_gemm():
try:
from flashinfer.gemm import gemm_fp8_nt_groupwise
except (AttributeError, ImportError) as error:
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires "
"flashinfer.gemm.gemm_fp8_nt_groupwise (validated with flashinfer-python==0.6.8). "
"FastVideo will not re-quantize this checkpoint to tensorwise FP8.") from error
return gemm_fp8_nt_groupwise
def _get_flashinfer_groupwise_backend(device: torch.device) -> str:
capability = torch.cuda.get_device_capability(device)
if capability[0] >= 12:
return "cutlass"
if capability[0] == 10:
return "trtllm"
capability_number = capability[0] * 10 + capability[1]
raise RuntimeError(f"FlashInfer groupwise FP8 requires a Blackwell GPU, got sm{capability_number}")
def _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor: torch.Tensor,
weight: torch.Tensor,
block_size: tuple[int, int],
weight_scale: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
input_2d = input_tensor.view(-1, input_tensor.shape[-1])
output_shape = [*input_tensor.shape[:-1], weight.shape[0]]
backend = _get_flashinfer_groupwise_backend(input_tensor.device)
if input_2d.dtype != torch.bfloat16:
raise RuntimeError("MiniMax-H3 FlashInfer groupwise FP8 requires BF16 activations; "
f"got {input_2d.dtype}. The SGLang FP16 Triton GEMM fallback is not enabled for H3.")
if backend == "trtllm" and input_2d.shape[1] < 256:
raise RuntimeError("MiniMax-H3 FlashInfer TRTLLM groupwise FP8 requires K >= 256; "
f"got K={input_2d.shape[1]}. The SGLang Triton GEMM fallback is not enabled for H3.")
gemm_fp8_nt_groupwise = _get_flashinfer_groupwise_fp8_gemm()
block_n, block_k = block_size
q_input, x_scale = _sglang_per_token_group_quant_fp8(
input_2d,
block_k,
column_major_scales=(backend == "trtllm"),
)
if backend == "cutlass":
m, k = input_2d.shape
n = weight.shape[0]
expected_x_scale_shape = (k // block_k, m)
expected_weight_scale_shape = (k // block_k, n // block_n)
if x_scale.shape == (m, k // block_k):
x_scale = x_scale.transpose(-1, -2).contiguous()
if weight_scale.shape == (n // block_n, k // block_k):
weight_scale = weight_scale.transpose(-1, -2).contiguous()
if x_scale.shape != expected_x_scale_shape or weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scale layout mismatch: "
f"x_scale={tuple(x_scale.shape)}, weight_scale={tuple(weight_scale.shape)}, "
f"expected={expected_x_scale_shape}/{expected_weight_scale_shape}")
if x_scale.dtype != torch.float32 or weight_scale.dtype != torch.float32:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scales must be float32")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale.contiguous(),
weight_scale.contiguous(),
out_dtype=input_2d.dtype,
backend="cutlass",
scale_major_mode="MN",
)
else:
expected_x_scale_shape = (input_2d.shape[0], input_2d.shape[1] // block_k)
expected_weight_scale_shape = (weight.shape[0] // block_n, weight.shape[1] // block_k)
if x_scale.shape != expected_x_scale_shape or x_scale.stride(0) != 1:
raise RuntimeError("FlashInfer TRTLLM block-FP8 activation scale layout mismatch: "
f"shape={tuple(x_scale.shape)}, stride={x_scale.stride()}, "
f"expected column-major {expected_x_scale_shape}")
if weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer TRTLLM block-FP8 weight scale layout mismatch: "
f"shape={tuple(weight_scale.shape)}, expected={expected_weight_scale_shape}")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale,
weight_scale,
out_dtype=input_2d.dtype,
backend="trtllm",
)
if bias is not None:
output += bias
return output.to(dtype=input_2d.dtype).view(*output_shape)
class MiniMaxH3SerializedFP8LinearMethod(LinearMethodBase):
"""Execute serialized 128x128 block-FP8 weights without re-quantizing them."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
self.weight_block_size = weight_block_size
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: list[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
output_size_per_partition = sum(output_partition_sizes)
block_n, block_k = self.weight_block_size
tp_size = get_tp_world_size()
if tp_size > 1 and input_size // input_size_per_partition == tp_size:
if input_size_per_partition % block_k:
raise ValueError(f"Weight input_size_per_partition={input_size_per_partition} is not divisible "
f"by block_k={block_k}")
if tp_size > 1 and output_size // output_size_per_partition == tp_size:
for output_partition_size in output_partition_sizes:
if output_partition_size % block_n:
raise ValueError(f"Weight output_partition_size={output_partition_size} is not divisible "
f"by block_n={block_n}")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
layer.orig_dtype = params_dtype
weight_loader = extra_weight_attrs.get("weight_loader")
weight = Parameter(
torch.empty(output_size_per_partition, input_size_per_partition, dtype=FP8_DTYPE),
requires_grad=False,
)
set_weight_attrs(weight, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
layer.register_parameter("weight", weight)
scale = Parameter(
torch.empty((output_size_per_partition + block_n - 1) // block_n,
(input_size_per_partition + block_k - 1) // block_k,
dtype=torch.float32),
requires_grad=False,
)
set_weight_attrs(scale, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
scale.data.fill_(torch.finfo(torch.float32).min)
layer.register_parameter("weight_scale_inv", scale)
layer.register_parameter("input_scale", None)
def process_weights_after_loading(self, layer: nn.Module) -> None:
weight = getattr(layer, "weight", None)
block_scales = getattr(layer, "weight_scale_inv", None)
if weight is None or block_scales is None:
raise ValueError("Serialized MiniMax-H3 FP8 linear is missing weight or weight_scale_inv")
if weight.dtype != FP8_DTYPE:
raise ValueError(f"Serialized MiniMax-H3 FP8 weight must be {FP8_DTYPE}, got {weight.dtype}")
if block_scales.dtype != torch.float32:
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must be float32, "
f"got {block_scales.dtype}")
block_n, block_k = self.weight_block_size
output_size, input_size = weight.shape
if output_size % block_n or input_size % block_k:
raise ValueError("Serialized MiniMax-H3 FP8 weight dimensions must be divisible by the 128x128 block size; "
f"got {tuple(weight.shape)}")
expected_scale_shape = (output_size // block_n, input_size // block_k)
if tuple(block_scales.shape) != expected_scale_shape:
raise ValueError("Serialized MiniMax-H3 FP8 scale shape mismatch: "
f"expected {expected_scale_shape}, got {tuple(block_scales.shape)}")
if not bool(torch.isfinite(block_scales).all()) or bool((block_scales <= 0).any()):
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must contain finite positive values")
layer.weight.data = weight.data
layer.weight_scale_inv.data = block_scales.data
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
if x.device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 execution requires CUDA")
capability = torch.cuda.get_device_capability(x.device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < MiniMaxH3SerializedFP8Config.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{MiniMaxH3SerializedFP8Config.get_min_capability()} or newer, "
f"got sm{capability_number}")
if not x.is_contiguous():
x = x.contiguous()
return _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
x,
layer.weight,
self.weight_block_size,
layer.weight_scale_inv,
bias,
)
__all__ = [
"MiniMaxH3SerializedFP8Config",
"MiniMaxH3SerializedFP8LinearMethod",
]
@@ -8,13 +8,13 @@ import torch
import torch.nn.functional as F
from torch import nn
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import MiniMaxH3SerializedFP8Config
from fastvideo.models.loader.weight_utils import default_weight_loader
@@ -227,15 +227,10 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
)
override = config.num_hidden_layers_override
self.num_layers = (config.num_hidden_layers
if override is None else min(config.num_hidden_layers, override))
self.output_hidden_state_index = config.output_hidden_state_index
self.layers = nn.ModuleList(
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
for index in range(self.num_layers))
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
if self.num_layers == config.num_hidden_layers else None)
for index in range(config.num_hidden_layers))
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
def forward(
@@ -243,14 +238,18 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
inputs_embeds: torch.Tensor,
position_ids: torch.Tensor,
attention_mask: torch.Tensor | None,
output_hidden_states: bool,
visual_pos_masks: torch.Tensor | None,
deepstack_visual_embeds: list[torch.Tensor] | None,
) -> torch.Tensor:
) -> BaseEncoderOutput:
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
attention_mask = None
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
hidden_states = inputs_embeds
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
for layer_index, layer in enumerate(self.layers):
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
hidden_states = layer(hidden_states, position_embeddings, attention_mask)
if deepstack_visual_embeds is not None and layer_index < len(deepstack_visual_embeds):
if visual_pos_masks is None:
@@ -259,9 +258,10 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
visual = deepstack_visual_embeds[layer_index].to(hidden_states.device, hidden_states.dtype)
updated = hidden_states[mask].clone() + visual
hidden_states[mask] = updated
if layer_index + 1 == self.output_hidden_state_index:
return hidden_states
raise RuntimeError(f"MiniMax-H3 text stack did not reach hidden_states[{self.output_hidden_state_index}]")
hidden_states = self.norm(hidden_states)
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
@@ -499,18 +499,10 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
return self.merger(hidden_states), deepstack_features
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
supports_hf_from_pretrained = False
supported_checkpoint_quantization_methods = frozenset({"fp8"})
@classmethod
def checkpoint_quantization_config_from_metadata(
cls,
metadata: dict[str, Any],
) -> MiniMaxH3SerializedFP8Config:
return MiniMaxH3SerializedFP8Config.from_config(metadata)
def __init__(self, config: MiniMaxH3Qwen3VLConfig) -> None:
super().__init__(config)
@@ -526,10 +518,6 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
def num_hidden_layers(self) -> int:
return self.config.num_hidden_layers
@property
def num_built_hidden_layers(self) -> int:
return self.language_model.num_layers
def _get_rope_index(
self,
input_ids: torch.Tensor,
@@ -622,39 +610,35 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
f"tokens={int(mask.sum())}, features={features.shape[0]}")
return mask
# no_grad, NOT inference_mode: with text_encoder_cpu_offload=True (the
# FastVideoArgs default) the loader FSDP2-shards this conditioner, and
# FSDP2's wait_for_unshard reads tensor._version via
# _unsafe_preserve_version_counter - inference tensors do not track
# version counters, so inference_mode crashes the first encode. no_grad
# frees the same activation memory and keeps prompt_embeds ordinary
# tensors (safe for any future backward through the conditioning).
@torch.no_grad()
def encode_ids(
def forward(
self,
input_ids: torch.Tensor,
*,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
if input_ids.ndim != 1:
raise ValueError(f"MiniMax-H3 slim forward expects 1-D input_ids, got shape={tuple(input_ids.shape)}")
if (pixel_values is None) != (image_grid_thw is None):
raise ValueError("pixel_values and image_grid_thw must be provided together")
if (pixel_values_videos is None) != (video_grid_thw is None):
raise ValueError("pixel_values_videos and video_grid_thw must be provided together")
input_ids = input_ids.unsqueeze(0)
inputs_embeds = self.language_model.embed_tokens(input_ids)
mm_token_type_ids: torch.Tensor | None = None,
**kwargs: Any,
) -> BaseEncoderOutput:
del mm_token_type_ids, kwargs
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Exactly one of input_ids or inputs_embeds is required")
if inputs_embeds is None:
assert input_ids is not None
inputs_embeds = self.language_model.embed_tokens(input_ids)
if input_ids is None and (pixel_values is not None or pixel_values_videos is not None):
raise ValueError("Multimodal Qwen3-VL inputs require input_ids for placeholder matching")
image_mask = None
video_mask = None
image_deepstack = None
video_deepstack = None
if pixel_values is not None:
if image_grid_thw is None:
if input_ids is None or image_grid_thw is None:
raise ValueError("pixel_values require input_ids and image_grid_thw")
image_features, image_deepstack = self._visual_features(pixel_values, image_grid_thw)
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -662,7 +646,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"image")
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
if pixel_values_videos is not None:
if video_grid_thw is None:
if input_ids is None or video_grid_thw is None:
raise ValueError("pixel_values_videos require input_ids and video_grid_thw")
video_features, video_deepstack = self._visual_features(pixel_values_videos, video_grid_thw)
video_features = video_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -690,34 +674,25 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
visual_mask = video_mask
deepstack_features = video_deepstack
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
hidden_states = self.language_model(
if position_ids is None:
if input_ids is None:
sequence_length = inputs_embeds.shape[1]
position_ids = torch.arange(sequence_length,
device=inputs_embeds.device).view(1, 1,
-1).expand(3, inputs_embeds.shape[0], -1)
else:
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, attention_mask)
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
outputs = self.language_model(
inputs_embeds,
position_ids,
None,
attention_mask,
output_hidden_states,
visual_mask,
deepstack_features,
)
if hidden_states.ndim != 3 or hidden_states.shape[0] != 1:
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
return hidden_states[0]
def forward(
self,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
return self.encode_ids(
input_ids,
pixel_values=pixel_values,
image_grid_thw=image_grid_thw,
pixel_values_videos=pixel_values_videos,
video_grid_thw=video_grid_thw,
)
outputs.attention_mask = attention_mask
return outputs
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
parameters = dict(self.named_parameters())
@@ -727,8 +702,6 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
if source_name == "lm_head.weight":
continue
name = source_name[6:] if source_name.startswith("model.") else source_name
if self._is_omitted_checkpoint_key(name):
continue
if name not in parameters:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
parameter = parameters[name]
@@ -737,23 +710,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
loaded.add(name)
return loaded
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
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
return (index.isdigit() and language_model.num_layers <= int(index) < self.config.num_hidden_layers)
EntryClass = MiniMaxH3Qwen3VLConditioner
__all__ = [
"MiniMaxH3Qwen3VLConditioner",
"MiniMaxH3SerializedFP8Config",
]
__all__ = ["MiniMaxH3Qwen3VLConditioner"]
+169 -74
View File
@@ -30,12 +30,13 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.quantization import get_quantization_config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_resolve_text_encoder_checkpoint_path,
from fastvideo.models.loader.ltx_single_file import (
component_weights,
is_single_file_bundle,
read_ltx_metadata,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
@@ -142,6 +143,41 @@ class ComponentLoader(ABC):
return GenericComponentLoader(transformers_or_diffusers)
def _check_connector_widths(transformer_config: dict[str, Any]) -> None:
"""Fail loudly when a connector's declared factorization is inconsistent.
A connector's width is ``num_attention_heads * attention_head_dim``, and it
has to equal the cross-attention width of the stream it feeds. This is the
class of mistake that loads cleanly and computes wrong -- a factorization
that multiplies out to the wrong number produces correctly shaped tensors
everywhere except the one parameter it silently mis-sizes. Check it while
both numbers are still in hand rather than discovering it as a shape
assertion thousands of tensors later, or not at all.
Only checks what the checkpoint actually declares; a missing field means
the arch default applies and there is nothing to contradict.
"""
for heads_key, dim_key, width_key, label in (
("connector_num_attention_heads", "connector_attention_head_dim",
"cross_attention_dim", "video"),
("audio_connector_num_attention_heads",
"audio_connector_attention_head_dim", "audio_cross_attention_dim",
"audio"),
):
heads = transformer_config.get(heads_key)
head_dim = transformer_config.get(dim_key)
width = transformer_config.get(width_key)
if heads is None or head_dim is None or width is None:
continue
if heads * head_dim != width:
raise ValueError(
f"Checkpoint declares an inconsistent {label} connector shape: "
f"{heads_key}={heads} x {dim_key}={head_dim} = {heads * head_dim}, "
f"but {width_key}={width}. The connector feeds that stream, so "
"the two must agree; refusing to build a model whose weights "
"would load into mis-sized parameters.")
class TextEncoderLoader(ComponentLoader):
"""Loader for text encoders."""
@@ -255,6 +291,13 @@ class TextEncoderLoader(ComponentLoader):
# revision=fastvideo_args.revision,
# model_override_args=None,
# )
# For a bundle, `model_path` is the declared text-encoder root rather
# than a `text_encoder/` directory: the language model lives there, its
# projection and connectors live inside the bundle, and the fields the
# directory layout reads out of `<repo>/transformer/config.json` come
# from the bundle's metadata instead.
bundle_path = fastvideo_args.model_path
single_file = is_single_file_bundle(bundle_path)
model_config = get_diffusers_config(model=model_path)
model_config.pop("_name_or_path", None)
model_config.pop("transformers_version", None)
@@ -281,6 +324,15 @@ class TextEncoderLoader(ComponentLoader):
if gemma_path and not gemma_path_from_candidate:
if not os.path.isabs(gemma_path):
model_config["gemma_model_path"] = os.path.normpath(os.path.join(repo_root, gemma_path))
if single_file:
# The encoder root was declared, not discovered, so it wins over
# anything the probing above turned up.
model_config["gemma_model_path"] = model_path
# That root's config.json describes the language model it holds,
# which is not the class to build here: the root supplies the
# encoder's dimensions, the wrapper class supplies the connectors,
# and only the wrapper is buildable from this loader.
model_config["architectures"] = ["LTX2GemmaTextEncoderModel"]
transformer_config_path = os.path.join(repo_root, "transformer", "config.json")
if os.path.isfile(transformer_config_path):
try:
@@ -294,6 +346,40 @@ class TextEncoderLoader(ComponentLoader):
rope_type = transformer_config.get("rope_type")
if rope_type is not None:
model_config["connector_rope_type"] = rope_type
# Each per-modality projection feeds its connector, so its
# output width is that stream's cross-attention width, which
# only the transformer config declares. Absent from both
# configs -> the arch config default stands.
for src, dst in (
("cross_attention_dim",
"video_feature_extractor_out_features"),
("audio_cross_attention_dim",
"audio_feature_extractor_out_features"),
# The rest of the text stack's shape is declared under these
# exact names too. Leaving any of them to the arch default
# builds a DIFFERENT model than the checkpoint describes:
# the audio connector silently inherits the video
# connector's width, and the feature extractor is built in
# the wrong one of its two forms. Same names on both sides,
# so the mapping is identity.
*((name, name) for name in (
"connector_num_attention_heads",
"connector_attention_head_dim",
"connector_num_layers",
"connector_num_learnable_registers",
"connector_positional_embedding_theta",
"connector_positional_embedding_max_pos",
"connector_apply_gated_attention",
"audio_connector_num_attention_heads",
"audio_connector_attention_head_dim",
"audio_connector_num_layers",
# Selects which feature-extractor form is built.
"caption_proj_before_connector",
)),
):
if dst not in model_config and src in transformer_config:
model_config[dst] = transformer_config[src]
_check_connector_widths(transformer_config)
except json.JSONDecodeError:
pass
logger.info("HF Model config: %s", model_config)
@@ -336,6 +422,12 @@ class TextEncoderLoader(ComponentLoader):
fastvideo_args,
encoder_precision,
use_text_encoder_override=True,
# Everything this model owns is inside the bundle: the projection
# and both connectors. The language model at the encoder root is
# loaded separately and lazily by the model itself, and is filtered
# out of `named_parameters`, so it must not be routed through here.
weight_iterator=(component_weights(bundle_path, "text_encoder")
if single_file else None),
)
def load_model(
@@ -347,50 +439,27 @@ class TextEncoderLoader(ComponentLoader):
dtype: str = "fp16",
use_text_encoder_override: bool = False, # prevent subclasses from misusing
cpu_offload: bool | None = None,
weight_iterator: Iterable[tuple[str, torch.Tensor]] | None = None,
):
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
runtime_device = get_local_torch_device()
from fastvideo.platforms import current_platform
if cpu_offload:
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
# Set quantization config if specified
if (use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None):
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
checkpoint_path = _resolve_text_encoder_checkpoint_path(
model_path,
fastvideo_args,
use_text_encoder_override,
)
checkpoint_quant_config = _configure_text_encoder_quantization(
model_config,
model_cls,
checkpoint_path,
)
if checkpoint_quant_config is not None:
if fastvideo_args.override_text_encoder_quant is not None:
raise ValueError("Serialized checkpoint quantization is selected from checkpoint metadata; "
"override_text_encoder_quant is an online conversion option and must be unset")
requested_dtype = PRECISION_TO_TYPE[dtype]
if requested_dtype not in checkpoint_quant_config.get_supported_act_dtypes():
raise ValueError(f"Serialized {checkpoint_quant_config.get_name()} text encoder does not support "
f"activation dtype {requested_dtype}")
checkpoint_quant_config.validate_runtime(runtime_device)
logger.info(
"Selected serialized %s text-encoder checkpoint execution from %s",
checkpoint_quant_config.get_name(),
checkpoint_path,
)
elif use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
@@ -409,20 +478,17 @@ class TextEncoderLoader(ComponentLoader):
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):
override_weights = self._get_all_weights(
model,
checkpoint_path,
to_cpu=bool(cpu_offload),
)
else:
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
override_weights = safetensors_weights_iterator(
[checkpoint_path],
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
to_cpu=use_cpu_offload,
)
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
)) # type: ignore
elif weight_iterator is not None:
# Same contract as `maybe_load_fsdp_model`'s `weight_iterator`:
# already-routed `(name, cpu_tensor)` pairs, for checkpoints
# that are not a directory of per-component weight files.
self.counter_before_loading_weights = time.perf_counter()
loaded_weights: set[str] = model.load_weights(weight_iterator) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
@@ -437,10 +503,6 @@ class TextEncoderLoader(ComponentLoader):
self.counter_after_loading_weights - self.counter_before_loading_weights,
)
if checkpoint_quant_config is not None:
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
logger.info("Validated %d serialized blockwise FP8 text-encoder linears", processed_linears)
# Explicitly move model to target device after loading weights
model = model.to(target_device)
@@ -483,7 +545,7 @@ class TextEncoderLoader(ComponentLoader):
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
if weights_not_loaded and model_config.quant_config is None:
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
@@ -739,7 +801,17 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
if is_single_file_bundle(model_path):
# The bundle stores the VAE config flat; the converted directory
# layout nests the same fields under "vae" and hoists the class out
# (`convert_ltx2_weights.py::_wrap_component_config`), and the build
# path below keys off that nesting. Match the shape rather than
# forking the build path.
section = deepcopy(read_ltx_metadata(model_path).config["vae"])
declared_class = section.pop("_class_name", None)
config = {"_class_name": declared_class, "vae": section}
else:
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
config.pop("_name_or_path", None)
assert class_name is not None, (
@@ -858,15 +930,21 @@ class VAELoader(ComponentLoader):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(target_device)
# Find all safetensors files
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}")
# Common case: a single `.safetensors` checkpoint file.
# Some models may be sharded into multiple files; in that case we merge.
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
if is_single_file_bundle(model_path):
# `safetensors_load_file` takes a path, not an iterator, so the
# bundle's prefix-routed tensors are materialized here instead.
# The VAE is small enough for that (see `component_weights`).
loaded = dict(component_weights(model_path, "vae"))
else:
# Find all safetensors files
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}")
# Common case: a single `.safetensors` checkpoint file.
# Some models may be sharded into multiple files; in that case we merge.
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
# LTX-2 CausalVideoAutoencoder needs per_channel_statistics remapping
if class_name == "CausalVideoAutoencoder" and "vae" in config:
@@ -1015,7 +1093,16 @@ class TransformerLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the transformer based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
# LTX ships one bundle holding every component instead of a
# per-component directory: the config lives in the file's
# ``__metadata__`` and the transformer's tensors are selected by their
# top-level key prefix.
single_file = is_single_file_bundle(model_path)
if single_file:
config = deepcopy(
read_ltx_metadata(model_path).config["transformer"])
else:
config = get_diffusers_config(model=model_path)
hf_config = deepcopy(config)
cls_name = config.pop("_class_name")
config.pop("_name_or_path", None)
@@ -1049,9 +1136,12 @@ class TransformerLoader(ComponentLoader):
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
# Find all safetensors files
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}")
if single_file:
safetensors_list = [str(model_path)]
else:
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}")
# 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)
@@ -1098,12 +1188,7 @@ class TransformerLoader(ComponentLoader):
# so recording here makes the decision readable from the loaded
# transformer — and records the narrowed one for teacher/critic.
resolved = record_resolved_attention_backend(dit_config)
# Every worker records its resolved backend so distributed profile
# snapshots can prove that all ranks use the requested kernels.
logger.info("Worker %s transformer attention backend: %s",
os.environ.get("RANK", "0"),
resolved.name if resolved else "automatic selection",
local_main_process_only=False)
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
@@ -1111,6 +1196,11 @@ class TransformerLoader(ComponentLoader):
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
# Route the bundle's tensors to the transformer, prefix
# stripped. Skipped when custom init weights replaced the
# file list above.
weight_iterator=(component_weights(model_path, "transformer")
if single_file and not use_custom_weights else None),
device=get_local_torch_device(),
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
@@ -1155,7 +1245,12 @@ class SchedulerLoader(ComponentLoader):
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
"""Load the scheduler based on the model path, and inference args."""
config = get_diffusers_config(model=model_path)
# A bundle keeps the scheduler's config in its own metadata; there is
# no `scheduler/` directory holding a `scheduler_config.json`.
if is_single_file_bundle(model_path):
config = deepcopy(read_ltx_metadata(model_path).config["scheduler"])
else:
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, (
+9 -2
View File
@@ -7,7 +7,7 @@
from __future__ import annotations
import os
import contextlib
from collections.abc import Callable, Generator
from collections.abc import Callable, Generator, Iterable
from itertools import chain
from typing import Any
@@ -136,9 +136,15 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
weight_iterator: Iterable[tuple[str, torch.Tensor]] | None = None,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
``weight_iterator`` overrides reading ``weight_dir_list``, for checkpoints
whose keys need routing before they reach the model (see
``ltx_single_file.component_weights``). It must yield the same
``(name, cpu_tensor)`` pairs that ``safetensors_weights_iterator`` does.
"""
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
# manually casting the inputs to the model
@@ -201,7 +207,8 @@ def maybe_load_fsdp_model(
fsdp_shard_conditions=model._fsdp_shard_conditions,
pin_cpu_memory=pin_cpu_memory)
weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=True)
if weight_iterator is None:
weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=True)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
load_model_from_full_model_state_dict(
model,
+296
View File
@@ -0,0 +1,296 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-file LTX checkpoint support.
LTX ships one ``.safetensors`` bundle holding every component, so the two
assumptions the directory-based loaders make do not hold:
* there is no per-component ``config.json`` -- every component's config lives
in the file's ``__metadata__`` (safetensors stores an 8-byte little-endian
header length, then that many bytes of JSON whose ``__metadata__`` object is
a flat ``str -> str`` map; ``safe_open(...).metadata()`` returns it);
* there is no per-component directory -- components are told apart by the
top-level prefix on each tensor key.
This module reads that metadata and routes tensors by prefix. It does not
convert anything to a diffusers layout.
"""
from __future__ import annotations
import json
import os
from collections.abc import Generator
from dataclasses import dataclass
from typing import Any
import torch
from safetensors import safe_open
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Top-level tensor key prefix per component. Verified against the LTX bundle
# header; every tensor in the file starts with one of these except
# ``duration_head.``, which has no FastVideo component today.
COMPONENT_PREFIXES: dict[str, str] = {
"transformer": "model.diffusion_model.",
"vae": "vae.",
"audio_vae": "audio_vae.",
"vocoder": "vocoder.",
"text_encoder": "text_embedding_projection.",
}
# Both embeddings connectors are stored *under* the transformer prefix, but
# FastVideo builds and runs them in the text encoder
# (``LTX2GemmaTextEncoderModel``'s ``Embeddings1DConnector``), not in the DiT,
# so they are routed to the text encoder instead. Only the transformer prefix
# is stripped: the sub-tree name survives so the text encoder's own rename
# table can map it onto the module name, the same way the already-converted
# repo layout is remapped. Matches how ``convert_ltx2_weights.py`` splits them.
TEXT_STACK_SUBPREFIXES: tuple[str, ...] = (
"video_embeddings_connector.",
"audio_embeddings_connector.",
)
def is_single_file_bundle(path: str) -> bool:
"""True when a model path names a bundle rather than a component directory."""
return str(path).endswith(".safetensors")
@dataclass(frozen=True)
class LTXCheckpointMetadata:
"""Parsed ``__metadata__`` of a single-file LTX checkpoint.
``config`` holds the per-component config sections keyed by component
(``transformer``, ``vae``, ``scheduler``, ``audio_vae``, ``vocoder``).
``model_version`` and ``gemma_source_checkpoint`` are carried through so
callers can validate the text encoder against the checkpoint that trained
it; ``gemma_source_checkpoint`` is absent on some variants. ``variant`` is
the training variant the header declares, when it declares one (see
:func:`bundle_variant`).
"""
config: dict[str, Any]
model_version: str | None
gemma_source_checkpoint: dict[str, Any] | None
variant: str | None = None
def read_ltx_metadata(path: str) -> LTXCheckpointMetadata:
"""Read the ``__metadata__`` config out of a single-file LTX checkpoint.
Only the header is parsed; no tensor data is touched. The remaining
metadata entries are deliberately not returned or logged -- the bundle also
carries a full license text and an ``encrypted_wandb_properties`` blob.
"""
with safe_open(path, framework="pt") as f:
metadata = f.metadata() or {}
config = json.loads(metadata["config"])
transformer = config.get("transformer")
if transformer is not None:
# ``frequencies_precision`` is the checkpoint's name for what the arch
# config calls ``double_precision_rope``; without this the field would
# silently fall back to its dataclass default.
precision = transformer.get("frequencies_precision")
if precision is not None:
transformer["double_precision_rope"] = precision == "float64"
gemma_source = metadata.get("gemma_source_checkpoint")
return LTXCheckpointMetadata(
config=config,
model_version=metadata.get("model_version"),
gemma_source_checkpoint=(json.loads(gemma_source) if gemma_source is not None else None),
variant=metadata.get("variant"),
)
def bundle_variant(metadata: LTXCheckpointMetadata, path: str) -> str:
"""The bundle's training variant: ``"distilled"`` or ``"base"``.
Decides which sampling preset a bundle gets (a distilled model wants its
short no-CFG schedule; everything else wants the standard one), so it must
be answerable from the header alone, without loading weights. An explicit
``variant`` entry in the file metadata wins when the checkpoint declares
one.
ponytail: filename fallback -- the known bundles declare no variant marker
in their headers (a distilled file and its sft sibling differ only in
fields incidental to the variant), so the "distilled" token in the file
name is the only signal available today. Drop the fallback when
checkpoints start declaring ``variant``.
"""
declared = metadata.variant or os.path.basename(path)
return "distilled" if "distilled" in declared.lower() else "base"
def bundle_model_index(path: str) -> dict[str, Any]:
"""Build a ``model_index.json``-shaped dict out of a bundle's own metadata.
The pipeline loader is written against a diffusers repo layout, which
answers two questions: which components exist, and what class is each. A
bundle already answers both in its ``__metadata__``, so this only reshapes
the answer -- it mirrors the entries
``convert_ltx2_weights.py::_build_model_index`` writes for the converted
directory layout, including the library each is declared under.
A section that exists but declares no class is emitted as ``[None, None]``
rather than dropped. ``ComposedPipelineBase.load_modules`` already treats a
null library as "declared, but not something to build" and removes the
component from the required set; dropping the key instead would fail its
required-module check for a component the checkpoint does carry.
``text_encoder`` and ``tokenizer`` are always declared: they live outside
the bundle, but the pipeline needs both.
"""
model_index: dict[str, Any] = {
# ponytail: the pipeline class is pinned by the registry's bundle
# table (`registry._bundle_config_info`) or an explicit
# `override_pipeline_cls_name`, and `load_modules` pops both of these
# without reading them. Nothing here is entitled to name a pipeline,
# so these are placeholders -- they exist only because those pops have
# no default.
"_class_name": None,
"_diffusers_version": None,
}
for component, section in read_ltx_metadata(path).config.items():
cls_name = (section.get("_class_name") if isinstance(section, dict) else None)
model_index[component] = (["diffusers", cls_name] if cls_name else [None, None])
model_index["text_encoder"] = ["transformers", "LTX2GemmaTextEncoderModel"]
model_index["tokenizer"] = ["transformers", "AutoTokenizer"]
return model_index
def build_dit_config(metadata: LTXCheckpointMetadata) -> LTX2VideoConfig:
"""Build the LTX-2 DiT config from checkpoint metadata.
Reuses ``update_model_arch``, so the metadata keys that name an arch field
win and the rest of the section is ignored -- no hand-written constants.
"""
config = LTX2VideoConfig()
config.update_model_arch(metadata.config["transformer"])
return config
def component_weights(
path: str,
component: str,
) -> Generator[tuple[str, torch.Tensor], None, None]:
"""Yield ``(key_with_prefix_stripped, tensor)`` for one component.
The transformer prefix covers two owners -- see
``TEXT_STACK_SUBPREFIXES`` -- so keys under it are split before the
component filter applies.
``safe_open`` mmaps the file and ``get_tensor`` materializes one tensor at
a time, so iterating never holds more than a single tensor in RAM. Callers
that need a state dict for a small component can do
``dict(component_weights(path, "vocoder"))``; do not do that for the
transformer.
ponytail: every rank reads the file itself, unlike
``safetensors_weights_iterator``, which has local rank 0 read and broadcast.
Correct either way, but the ceiling is filesystem bandwidth -- N ranks read
N copies. Add the broadcast if loading a bundle over a shared filesystem
turns out to be the bottleneck.
"""
prefix = COMPONENT_PREFIXES[component]
transformer_prefix = COMPONENT_PREFIXES["transformer"]
with safe_open(path, framework="pt") as f:
for key in f.keys():
if key.startswith(transformer_prefix):
name = key[len(transformer_prefix):]
# str.startswith takes a tuple.
owner = ("text_encoder" if name.startswith(TEXT_STACK_SUBPREFIXES) else "transformer")
if owner != component:
continue
elif key.startswith(prefix):
name = key[len(prefix):]
else:
continue
yield name, f.get_tensor(key)
def resolve_text_encoder_root(
configured_path: str | None,
override_path: str | None = None,
metadata: LTXCheckpointMetadata | None = None,
) -> str:
"""Locate the text-encoder root that goes with a single-file bundle.
A bundle carries every component's *weights* but no pointer to the text
encoder, which lives outside it. The root therefore has to be declared,
and there are exactly two ways to declare it:
* ``override_path`` -- an explicit per-run argument, so a one-off run needs
no config edit;
* ``configured_path`` -- the pipeline config's encoder path, which is where
model composition already lives and which survives files being moved.
Deliberately absent: any search of the bundle's directory. Picking up an
encoder nobody asked for changes what gets loaded without being requested,
and silently picks the wrong one when two sit side by side. When neither
source is set this raises and names both, rather than guessing.
``metadata`` is used only to *validate* a declared root, never to find one:
a bundle may not declare an encoder pairing at all, so discovery cannot
depend on it.
"""
root = override_path or configured_path
if not root:
raise ValueError("A single-file checkpoint does not carry its text encoder, so the "
"encoder root must be declared. Set `gemma_model_path` in the "
"pipeline config, or pass the text-encoder path explicitly for "
"this run. It is not inferred from the checkpoint's directory: "
"loading whichever encoder happens to sit beside the file would "
"silently pick the wrong one when several are present.")
expected = (metadata.gemma_source_checkpoint or {}).get("gemma_version") if metadata is not None else None
if expected:
# ponytail: warn, don't raise -- the declared root is the user's
# explicit instruction and the pairing is advisory. Promote to a hard
# error if a mismatch ever turns out to produce silent garbage rather
# than an obvious shape failure.
actual = _read_encoder_version(root)
if actual is not None and actual != expected:
logger.warning(
"Checkpoint expects text-encoder version %r but the encoder at "
"%s reports %r; continuing with the declared root.", expected, root, actual)
return root
def _read_encoder_version(root: str) -> str | None:
"""The encoder's declared version, or None if it declares none."""
config_path = os.path.join(root, "config.json")
if not os.path.isfile(config_path):
return None
try:
with open(config_path, encoding="utf-8") as f:
return json.load(f).get("gemma_version")
except (OSError, json.JSONDecodeError):
return None
def model_index_and_component_path(model_path: str, module_type: str) -> tuple[dict[str, Any], str]:
"""``(model_index, component_path)`` for a directory repo OR a bundle.
The two differ in both halves: a repo answers "what components exist" from
``model_index.json`` and puts each in its own subdirectory, while a bundle
declares its components in its own metadata and holds them all in one file.
Callers that resolve those two things together should route through here so
the bundle case is handled once instead of at every site.
ponytail: a bundle's text encoder and tokenizer live OUTSIDE the file, so
this returns the bundle path for them too, which is wrong for those two
module types. Training does not load them (text embeddings are
preprocessed), so it does not arise. Give this the resolved encoder root the
day a caller needs them.
"""
if is_single_file_bundle(model_path):
return bundle_model_index(model_path), model_path
from fastvideo.utils import verify_model_config_and_directory
return verify_model_config_and_directory(model_path), os.path.join(model_path, module_type)
@@ -1,127 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Checkpoint-serialized quantization lifecycle for native text encoders."""
import json
import os
from itertools import chain
from typing import Any
import torch
import torch.nn as nn
from safetensors.torch import safe_open
from fastvideo.configs.models import EncoderConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.linear import LinearBase, UnquantizedLinearMethod
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.models.encoders.base import TextEncoder
def _resolve_text_encoder_checkpoint_path(
model_path: str,
fastvideo_args: FastVideoArgs,
use_text_encoder_override: bool,
) -> str:
override = fastvideo_args.override_text_encoder_safetensors if use_text_encoder_override else None
checkpoint_path = override or model_path
if not os.path.exists(checkpoint_path):
raise FileNotFoundError(f"Text-encoder checkpoint does not exist: {checkpoint_path}")
if not os.path.isdir(checkpoint_path) and not os.path.isfile(checkpoint_path):
raise ValueError(f"Text-encoder checkpoint must be a file or directory: {checkpoint_path}")
return checkpoint_path
def _read_text_encoder_checkpoint_quantization_config(checkpoint_path: str) -> dict[str, Any] | None:
checkpoint_dir = checkpoint_path if os.path.isdir(checkpoint_path) else os.path.dirname(checkpoint_path)
config_path = os.path.join(checkpoint_dir, "config.json")
if os.path.isfile(config_path):
try:
with open(config_path, encoding="utf-8") as config_file:
checkpoint_config = json.load(config_file)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid text-encoder checkpoint config: {config_path}") from error
quantization_config = checkpoint_config.get("quantization_config")
if quantization_config is not None:
if not isinstance(quantization_config, dict):
raise ValueError(f"quantization_config in {config_path} must be an object")
return quantization_config
if not os.path.isfile(checkpoint_path) or not checkpoint_path.endswith(".safetensors"):
return None
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint_file:
metadata = checkpoint_file.metadata() or {}
for key in ("quantization_config", "_quantization_metadata"):
serialized = metadata.get(key)
if serialized is None:
continue
try:
quantization_config = json.loads(serialized)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid {key} metadata in {checkpoint_path}") from error
if not isinstance(quantization_config, dict):
raise ValueError(f"{key} metadata in {checkpoint_path} must decode to an object")
return quantization_config
return None
def _configure_text_encoder_quantization(
model_config: EncoderConfig,
model_cls: type[nn.Module],
checkpoint_path: str,
) -> QuantizationConfig | None:
if not issubclass(model_cls, TextEncoder):
return None
checkpoint_quantization = _read_text_encoder_checkpoint_quantization_config(checkpoint_path)
if checkpoint_quantization is None:
return None
quant_method = str(checkpoint_quantization.get("quant_method", "")).lower()
if not quant_method:
raise ValueError(f"Quantized text-encoder checkpoint {checkpoint_path} does not declare quant_method")
supported_methods = getattr(model_cls, "supported_checkpoint_quantization_methods", frozenset())
if quant_method not in supported_methods:
supported = ", ".join(sorted(supported_methods)) or "none"
raise ValueError(f"Text encoder {model_cls.__name__} does not support serialized {quant_method!r} "
f"checkpoints (supported: {supported})")
factory = getattr(model_cls, "checkpoint_quantization_config_from_metadata", None)
if not callable(factory):
raise ValueError(f"Text encoder {model_cls.__name__} advertises serialized {quant_method!r} support "
"without a checkpoint quantization factory")
quant_config = factory(checkpoint_quantization)
model_config.quant_config = quant_config
return quant_config
def _module_tensor_device(module: nn.Module) -> torch.device | None:
devices = {
tensor.device
for tensor in chain(
module.parameters(recurse=False),
module.buffers(recurse=False),
)
}
if len(devices) > 1:
raise ValueError(f"Quantized text-encoder module {type(module).__name__} spans multiple devices: {devices}")
return next(iter(devices), None)
def _process_quantized_text_encoder_weights(model: nn.Module, process_device: torch.device) -> int:
"""Run quantized post-load hooks one linear at a time on ``process_device``."""
processed = 0
for module in model.modules():
if not isinstance(module, LinearBase) or isinstance(module.quant_method, UnquantizedLinearMethod):
continue
if module.quant_method is None:
continue
original_device = _module_tensor_device(module)
try:
module.to(process_device)
module.quant_method.process_weights_after_loading(module)
finally:
if original_device is not None:
module.to(original_device)
processed += 1
if processed == 0:
raise ValueError("Serialized quantized text-encoder checkpoint selected, but no quantized linear layers exist")
return processed
+2
View File
@@ -38,6 +38,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
# ``_class_name`` declared by single-file LTX checkpoint metadata.
"AVTransformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
"LingBotWorld2CausalFastTransformer3DModel": (
@@ -392,16 +392,9 @@ class MiniMaxH3AudioBigVGANDecoder(nn.Module):
return torch.clamp(hidden_states, min=-1.0, max=1.0)
def _is_minimax_h3_audio_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the audio decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3AudioBigVGANDecoder)
class MiniMaxH3AudioVAE(nn.Module):
"""DAC encoder plus BigVGAN decoder for mono 32 kHz waveforms."""
_compile_conditions = [_is_minimax_h3_audio_vae_decoder]
def __init__(self, config: MiniMaxH3AudioVAEConfig):
super().__init__()
self.config = config
@@ -1,378 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Sequence-parallel chunk scheduling for the MiniMax-H3 video VAE.
The H3 video VAE decodes a video as a series of temporal-chunk decoder
forwards whose outputs are joined by a short deterministic frame blend
(``AutoencoderKLMiniMaxH3._decode_chunks``), and encodes videos as fully
independent ``clip_length``-frame encoder forwards. Neither the chunk decode
nor the clip encode has any cross-chunk data dependency — only the *joining*
of decoded chunks (overlap blending, frame trimming) is sequential. This
module round-robins the chunk/clip forwards across the ranks of a
sequence-parallel group and replays the serial joining logic on the
assembling rank, reproducing the serial result bit for bit.
Bit-exactness contract:
- every rank holds an identical copy of the inputs (the H3 DiT all-gathers
its outputs, and reference pixels are prepared identically on all ranks);
- a chunk decoded on any rank is bitwise the tensor the serial loop would
produce (identical weights, inputs, and deterministic kernels on identical
GPUs), and NCCL transports it bitwise;
- every serialization point of the serial algorithm (overlap blending, frame
trimming, pixel denormalization, output-buffer copies, moment
concatenation and token-drop trimming) runs on the assembling rank in
serial order via the same VAE methods the serial path uses.
Collective safety: all group ranks must call these functions together with
identically shaped inputs. Work proceeds in rounds of one collective each;
ranks without a chunk in the final round contribute a placeholder tensor, so
participation is uniform by construction and no rank-dependent branch guards
a collective.
Caveat — compiled decoders (``enable_torch_compile_vae``): inductor autotunes
kernel configs per process at first call, so a compiled decoder is only
deterministic WITHIN a process, not across processes. Chunks decoded on other
ranks then differ from the serial rank's decode of the same chunk exactly as
two serial runs in different processes would (measured on GB200 at 124f:
max 63/255 on <0.5% of pixels, mean ~1e-2/255, first chunk bit-identical).
With the eager decoder — the pipeline default — parallel output is bitwise
equal to serial ``decode_to_pixels``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from fastvideo.models.vaes.minimax_h3_video import (
AutoencoderKLMiniMaxH3,
AutoencoderKLOutput,
DiagonalGaussianDistribution,
)
from fastvideo.profiler import nvtx_range
if TYPE_CHECKING:
from fastvideo.distributed.parallel_state import GroupCoordinator
# Collective used to move decoded chunk segments to the assembling rank.
# "gather" moves each segment once (destination-only); "all_gather" also
# leaves every rank with every segment. Both are exact; the default is the
# faster one measured on GB200 NVL72 (see the PR notes).
DECODE_GATHER_STRATEGIES = ("gather", "all_gather")
DEFAULT_DECODE_GATHER_STRATEGY = "gather"
def parallel_chunk_indices(num_chunks: int, world_size: int, rank_in_group: int) -> list[int]:
"""Round-robin chunk ownership: chunk ``i`` belongs to rank ``i % world_size``."""
if num_chunks < 0:
raise ValueError(f"num_chunks must be non-negative, got {num_chunks}.")
if world_size < 1:
raise ValueError(f"world_size must be positive, got {world_size}.")
if not 0 <= rank_in_group < world_size:
raise ValueError(f"rank_in_group {rank_in_group} out of range for world_size {world_size}.")
return list(range(rank_in_group, num_chunks, world_size))
def _num_rounds(num_chunks: int, world_size: int) -> int:
return -(-num_chunks // world_size)
def _decode_segment(vae: AutoencoderKLMiniMaxH3, z_padded: torch.Tensor, chunk_index: int) -> torch.Tensor:
"""Decode one temporal chunk's clip and keep the frames the join consumes.
The serial loop uses two spans of each decoded clip: the chunk body
``clip[:, :, frame_pre_padding:chunk_num_frames]`` and (when
``token_drop > 0``) the blend tail
``clip[:, :, chunk_num_frames + frame_pre_padding:]``. Everything from
``frame_pre_padding`` on covers both, so one contiguous slice per chunk
travels over the wire. ``.contiguous()`` also detaches the segment from
any decoder-owned storage (e.g. a compiled decoder's reuse pools) before
the next chunk decode can overwrite it.
"""
start = chunk_index * vae.tokens_chunk_size
with nvtx_range(f"minimax_h3.vae.parallel_chunk.{chunk_index}"):
clip = vae._decode_clip(z_padded[:, :, start:start + vae.tokens_chunk_size + vae.token_overlap])
return clip[:, :, vae.frame_pre_padding:].contiguous()
class _ChunkAssembler:
"""Replay the serial chunk-joining semantics of ``_decode_chunks`` +
``_decode_to_pixels`` on gathered chunk segments, in chunk order.
On CUDA the joining kernels and output copies run on a dedicated side
stream: they depend only on already-gathered segments, so running them
off the main stream keeps the assembling rank's next chunk decode (and
therefore every other rank's next collective) off the assembly's tail.
Stream placement cannot change values — the ops and their order are
identical — so bit-exactness with the serial path is unaffected.
"""
def __init__(self, vae: AutoencoderKLMiniMaxH3, output: torch.Tensor, output_num_frames: int,
non_blocking: bool, device: torch.device) -> None:
self._vae = vae
self._output = output
self._output_num_frames = output_num_frames
self._non_blocking = non_blocking
self._body_frames = vae.tokens_chunk_size * vae.temporal_compression_ratio - vae.frame_pre_padding
self._overlap: torch.Tensor | None = None
self._frame_start = 0
self._stream = torch.cuda.Stream(device) if device.type == "cuda" else None
def push(self, segment: torch.Tensor) -> None:
"""Consume the next chunk's segment (``clip[:, :, frame_pre_padding:]``)."""
if self._stream is None:
self._push(segment)
return
# The segment is produced on the current (collective) stream; hand it
# to the assembly stream and pin its storage until assembly reads it.
self._stream.wait_stream(torch.cuda.current_stream(segment.device))
segment.record_stream(self._stream)
with torch.cuda.stream(self._stream):
self._push(segment)
def _push(self, segment: torch.Tensor) -> None:
vae = self._vae
chunk = segment[:, :, :self._body_frames]
if self._overlap is not None:
chunk = vae._blend(self._overlap, chunk, vae.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], self._output_num_frames - self._frame_start)
chunk = chunk[:, :, :num_frames]
# The tail past the body (and its pre-padding gap) is the next
# chunk's blend overlap — the serial loop's ``next_overlap``.
self._overlap = segment[:, :, self._body_frames + vae.frame_pre_padding:] if vae.config.token_drop > 0 else None
if num_frames > 0:
self._emit(chunk)
def finalize(self) -> None:
"""Emit the final overlap tail exactly as the serial generator does."""
if self._overlap is not None and self._frame_start < self._output_num_frames:
tail = self._overlap[:, :, :self._output_num_frames - self._frame_start]
if self._stream is None:
self._emit(tail)
else:
with torch.cuda.stream(self._stream):
self._emit(tail)
if self._frame_start != self._output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {self._frame_start} frames into an output buffer expecting "
f"{self._output.shape[2]}.")
def synchronize(self) -> None:
"""Drain assembly kernels and output copies before the buffer is read."""
if self._stream is not None:
self._stream.synchronize()
def _emit(self, chunk: torch.Tensor) -> None:
pixels = self._vae.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._vae._copy_chunk_pixels(pixels, self._output, self._frame_start, self._non_blocking)
self._frame_start += pixels.shape[2]
def _broadcast_segment_meta(group: "GroupCoordinator",
segment: torch.Tensor | None) -> tuple[torch.dtype, tuple[int, ...]]:
"""Share the leader's real segment dtype/shape so placeholder tensors match.
The decoder's output dtype depends on the surrounding autocast context;
deriving it on the leader from an actually decoded segment (instead of
predicting it) keeps collective dtypes correct by construction.
"""
meta = (segment.dtype, tuple(segment.shape)) if segment is not None else None
meta = group.broadcast_object(meta, src=0)
if meta is None:
raise RuntimeError("MiniMax-H3 parallel VAE meta broadcast returned no leader metadata.")
return meta
def decode_to_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str = DEFAULT_DECODE_GATHER_STRATEGY,
) -> torch.Tensor | None:
"""Chunk-parallel ``decode_to_pixels`` across a sequence-parallel group.
All group ranks call this together with identical ``z``. Temporal chunks
are decoded round-robin across the group and their segments move to the
group's first rank, which assembles bitwise the serial
``decode_to_pixels`` result into ``output``. Only the first rank passes
``output`` (validated exactly like the serial API); other ranks pass
``None`` and receive ``None``.
"""
if strategy not in DECODE_GATHER_STRATEGIES:
raise ValueError(f"Unknown parallel-decode strategy {strategy!r}; expected one of {DECODE_GATHER_STRATEGIES}.")
is_leader = group.rank_in_group == 0
if is_leader:
if output is None:
raise ValueError("The first sequence-parallel rank must provide the CPU output buffer.")
expected_shape = vae.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
elif output is not None:
raise ValueError("Only the first sequence-parallel rank may provide an output buffer.")
if group.world_size == 1:
return vae.decode_to_pixels(z, output)
try:
if vae.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
slice_output = output[batch_index:batch_index + 1] if output is not None else None
_decode_single_parallel(vae, z_slice, slice_output, group, strategy)
else:
_decode_single_parallel(vae, z, output, group, strategy)
finally:
# Drain the leader's async chunk copies before the caller (or an
# exception handler) can read or release the pinned buffer.
if output is not None and vae._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def _decode_single_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str,
) -> None:
pad_tokens, num_chunks, output_num_frames = vae._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
world_size = group.world_size
rank = group.rank_in_group
# Every rank decodes its round-0 chunk BEFORE the metadata rendezvous so
# the first decodes run concurrently (a rank that waited on the broadcast
# first would idle a full chunk-decode behind the leader). The leader
# owns chunk 0 under round-robin assignment, so its segment supplies real
# dtype/shape for placeholder rounds instead of guessing autocast state.
first_segment = _decode_segment(vae, z, rank) if rank < num_chunks else None
segment_dtype, segment_shape = _broadcast_segment_meta(group, first_segment if rank == 0 else None)
assembler = None
if output is not None:
non_blocking = vae._streams_chunk_copies(z, output)
assembler = _ChunkAssembler(vae, output, output_num_frames, non_blocking, z.device)
try:
segment_frames = segment_shape[2]
for round_index in range(_num_rounds(num_chunks, world_size)):
chunk_index = round_index * world_size + rank
if chunk_index >= num_chunks:
segment = torch.zeros(segment_shape, dtype=segment_dtype, device=z.device)
elif round_index == 0 and first_segment is not None:
segment = first_segment
else:
segment = _decode_segment(vae, z, chunk_index)
with nvtx_range(f"minimax_h3.vae.parallel_{strategy}.{round_index}"):
if strategy == "gather":
gathered = group.gather(segment, dst=0, dim=2)
else:
gathered = group.all_gather(segment, dim=2)
if assembler is None or gathered is None:
continue
for slot in range(world_size):
if round_index * world_size + slot >= num_chunks:
break
assembler.push(gathered.narrow(2, slot * segment_frames, segment_frames))
if assembler is not None:
assembler.finalize()
finally:
# Drain assembly-stream copies into ``output`` even on the error path
# so an exception cannot leave an in-flight DMA into a buffer the
# caller may release.
if assembler is not None:
assembler.synchronize()
def _encode_clip_moments(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor, clip_index: int) -> torch.Tensor:
"""Encode one ``clip_length``-frame clip exactly as ``_encode_pixels`` does."""
clip_length = vae.config.clip_length
frame_start = clip_index * clip_length
with nvtx_range(f"minimax_h3.vae.parallel_encode_clip.{clip_index}"):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=vae.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = vae.normalize_pixels(clip)
return vae._encode_clip(clip).contiguous()
def encode_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
pixels: torch.Tensor,
group: "GroupCoordinator",
) -> AutoencoderKLOutput:
"""Clip-parallel ``encode_pixels`` across a sequence-parallel group.
Encoder clips have no cross-clip dependency (no overlap, no blending), so
ranks encode disjoint clips and all-gather the per-clip moment tensors.
Every rank returns the identical full posterior — preserving the serial
contract that all ranks hold the same encoded latents — bitwise equal to
``vae.encode_pixels(pixels)``. Moments are latent-sized (a few MB per
clip), so the all-gather is negligible next to the clip forwards.
"""
if pixels.ndim != 5 or pixels.shape[1] != vae.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {vae.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if group.world_size == 1:
return vae.encode_pixels(pixels)
if vae.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([_encode_single_parallel(vae, pixel_slice, group) for pixel_slice in pixels.split(1)])
else:
moments = _encode_single_parallel(vae, pixels, group)
return AutoencoderKLOutput(latent_dist=DiagonalGaussianDistribution(moments))
def _encode_single_parallel(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor,
group: "GroupCoordinator") -> torch.Tensor:
clip_length = vae.config.clip_length
num_clips = -(-pixels.shape[2] // clip_length)
world_size = group.world_size
rank = group.rank_in_group
# Same first-work-then-rendezvous ordering as the decode path: encode the
# round-0 clip before the metadata broadcast so first encodes overlap.
first_moments = _encode_clip_moments(vae, pixels, rank) if rank < num_clips else None
moment_dtype, moment_shape = _broadcast_segment_meta(group, first_moments if rank == 0 else None)
moment_tokens = moment_shape[2]
parts: list[torch.Tensor] = []
for round_index in range(_num_rounds(num_clips, world_size)):
clip_index = round_index * world_size + rank
if clip_index >= num_clips:
moments = torch.zeros(moment_shape, dtype=moment_dtype, device=vae.pixel_mean.device)
elif round_index == 0 and first_moments is not None:
moments = first_moments
else:
moments = _encode_clip_moments(vae, pixels, clip_index)
gathered = group.all_gather(moments, dim=2)
for slot in range(world_size):
if round_index * world_size + slot >= num_clips:
break
parts.append(gathered.narrow(2, slot * moment_tokens, moment_tokens))
encoded = torch.cat(parts, dim=2)
if vae.config.token_drop > 0:
encoded = encoded[:, :, :-vae.config.token_drop]
return encoded
__all__ = [
"DECODE_GATHER_STRATEGIES",
"DEFAULT_DECODE_GATHER_STRATEGY",
"decode_to_pixels_parallel",
"encode_pixels_parallel",
"parallel_chunk_indices",
]
+57 -298
View File
@@ -7,7 +7,6 @@ This module intentionally uses only PyTorch and FastVideo configuration types.
"""
import math
from collections.abc import Iterator
from dataclasses import dataclass
import torch
@@ -15,10 +14,7 @@ import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from fastvideo.attention import get_attn_backend
from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
class DiagonalGaussianDistribution:
@@ -295,7 +291,6 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
class MiniMaxH3VideoAttention(nn.Module):
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
"""Build projections and the selected dense FastVideo attention implementation."""
super().__init__()
self.heads = heads
self.dim_head = dim_head
@@ -307,34 +302,12 @@ class MiniMaxH3VideoAttention(nn.Module):
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
self.attn_impl = None
from fastvideo.platforms import current_platform
if current_platform.is_cuda_alike():
attention_backend = get_attn_backend(
dim_head,
# FlashAttention executes the FP32 VAE activations in BF16 and
# restores FP32 output, so resolve against the kernel dtype.
torch.bfloat16,
supported_attention_backends=(
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
),
)
self.attn_impl = attention_backend.get_impl_cls()(
num_heads=heads,
head_size=dim_head,
softmax_scale=dim_head**-0.5,
num_kv_heads=heads,
causal=False,
)
def forward(
self,
hidden_states: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""Apply dense self-attention to one spatial VAE token sequence."""
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
@@ -355,17 +328,9 @@ class MiniMaxH3VideoAttention(nn.Module):
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
if self.attn_impl is not None and query.device.type != "cpu":
# VAE decoding has no diffusion-step metadata, so call the selected
# backend implementation directly with dense BSHD tensors.
hidden_states = self.attn_impl.forward(query, key, value, None)
hidden_states = hidden_states.flatten(2, 3)
else:
# Keep CPU construction and execution available without requiring
# an accelerator attention backend.
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
return self.to_out[0](hidden_states)
@@ -468,7 +433,6 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Decode one latent spatial input through the H3 video transformer."""
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).reshape(
batch_size,
@@ -518,11 +482,6 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
)
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)
class AutoencoderKLMiniMaxH3(nn.Module):
"""MiniMax-H3 causal encoder and ViT decoder with exact release geometry."""
@@ -530,7 +489,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
_no_split_modules = ["MiniMaxH3VideoResnetBlock3d", "MiniMaxH3VideoTransformerBlock"]
_repeated_blocks = ["MiniMaxH3VideoTransformerBlock"]
_keep_in_fp32_modules = ["encoder", "decoder", "quant_conv", "post_quant_conv"]
_compile_conditions = [_is_minimax_h3_video_vae_decoder]
def __init__(self, config: MiniMaxH3VideoVAEConfig) -> None:
super().__init__()
@@ -696,15 +654,12 @@ class AutoencoderKLMiniMaxH3(nn.Module):
slice_rest[dim] = slice(blend_extent, None)
return torch.cat([blended, b[tuple(slice_rest)]], dim=dim)
# The fixed spatial tile grid reuses one compiled blend-and-concatenate graph.
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
def _stitch_tiles(
self,
tiles: list[list[torch.Tensor]],
height_overlaps: list[int],
width_overlaps: list[int],
) -> torch.Tensor:
"""Blend decoded tile overlaps and concatenate the spatial canvas."""
result_rows = []
for row_index, row in enumerate(tiles):
result_row = []
@@ -721,12 +676,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
result_rows.append(torch.cat(result_row, dim=-1))
return torch.cat(result_rows, dim=-2)
# Each fixed-shape latent tile reuses one compiled decoder-input projection.
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
def _project_decoder_tile(self, tile: torch.Tensor) -> torch.Tensor:
"""Project one spatial latent tile into the decoder input channels."""
return self.post_quant_conv(tile)
def _encode_clip(self, x: torch.Tensor) -> torch.Tensor:
if not self.use_tiling:
return self.quant_conv(self.encoder(x))
@@ -750,64 +699,36 @@ class AutoencoderKLMiniMaxH3(nn.Module):
rows.append(row)
latent_y_overlaps = [overlap // self.spatial_compression_ratio for overlap in y_overlaps]
latent_x_overlaps = [overlap // self.spatial_compression_ratio for overlap in x_overlaps]
# Under mode="reduce-overhead" the stitched canvas is a CUDA-graph
# static buffer that the next _stitch_tiles replay overwrites. Callers
# (_encode/_encode_pixels/encode_keyframe) collect per-clip results
# across replays before concatenating, so hand them a caller-owned
# tensor instead of cudagraph-pooled storage.
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps).clone()
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps)
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
"""Decode one temporal clip, with optional overlapping spatial tiles."""
with nvtx_range("minimax_h3.vae.decode_clip"):
if not self.use_tiling:
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"):
projected_clip = self.post_quant_conv(z)
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"):
return self.decoder(projected_clip)
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
with nvtx_range("minimax_h3.vae.decode_clip.split_tiles"):
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
# 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)):
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}"):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
projected_tile = self._project_decoder_tile(tile)
with nvtx_range("minimax_h3.vae.decode_clip.tile.decoder_forward"):
decoded_tile = self.decoder(projected_tile)
row.append(decoded_tile)
rows.append(row)
with nvtx_range("minimax_h3.vae.decode_clip.stitch_tiles"):
# Same CUDA-graph output-ownership contract as _encode_clip:
# _decode collects chunks across _stitch_tiles replays before
# torch.cat, so the pooled canvas must not escape this driver.
# (The streaming _decode_to_pixels path copies each chunk out
# before the next decode and never held stale storage; the
# clone keeps that path correct too at one D2D copy per chunk.)
return self._stitch_tiles(rows, y_overlaps, x_overlaps).clone()
if not self.use_tiling:
return self.decoder(self.post_quant_conv(z))
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
for y_position, y_length in zip(y_indices, y_lengths):
row = []
for x_position, x_length in zip(x_indices, x_lengths):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
row.append(self.decoder(self.post_quant_conv(tile)))
rows.append(row)
return self._stitch_tiles(rows, y_overlaps, x_overlaps)
def _encode(self, x: torch.Tensor) -> torch.Tensor:
clip_length = self.config.clip_length
@@ -826,157 +747,43 @@ class AutoencoderKLMiniMaxH3(nn.Module):
moments = moments[:, :, :-self.config.token_drop]
return moments
def _encode_pixels(self, pixels: torch.Tensor) -> torch.Tensor:
"""Encode unnormalized pixels while keeping full videos off the accelerator."""
clip_length = self.config.clip_length
moments = []
for frame_start in range(0, pixels.shape[2], clip_length):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=self.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = self.normalize_pixels(clip)
moments.append(self._encode_clip(clip))
del clip
encoded = torch.cat(moments, dim=2)
if self.config.token_drop > 0:
encoded = encoded[:, :, :-self.config.token_drop]
return encoded
def _temporal_decode_plan(self, latent_num_frames: int) -> tuple[int, int, int]:
"""Return pad tokens, chunk count, and exact decoded frame count."""
if latent_num_frames <= 0:
raise ValueError(f"MiniMax-H3 latent frame count must be positive, got {latent_num_frames}.")
token_drop = self.config.token_drop
def _decode(self, z: torch.Tensor) -> torch.Tensor:
tokens_chunk_size = self.tokens_chunk_size
token_drop = self.config.token_drop
temporal_ratio = self.temporal_compression_ratio
num_tokens = latent_num_frames + token_drop
chunk_num_frames = tokens_chunk_size * temporal_ratio
num_tokens = z.shape[2] + token_drop
pad_tokens = (-num_tokens) % tokens_chunk_size
num_chunks = (num_tokens + pad_tokens) // tokens_chunk_size - int(token_drop > 0)
if num_chunks < 1:
pad_tokens += tokens_chunk_size
num_chunks = 1
decoded_num_frames = num_chunks * (tokens_chunk_size * temporal_ratio - self.frame_pre_padding)
if token_drop > 0:
decoded_num_frames += self.frame_overlap
if pad_tokens > 0:
intra_tail = self.config.clip_length % temporal_ratio
pad_frames = sum(intra_tail if intra_tail and (latent_num_frames + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded_num_frames -= pad_frames
if decoded_num_frames <= 0:
raise RuntimeError(
f"MiniMax-H3 decode plan produced {decoded_num_frames} frames for {latent_num_frames} latent "
"frames; the clip_length/token_drop configuration is inconsistent.")
return pad_tokens, num_chunks, decoded_num_frames
def _decode_chunks(self, z: torch.Tensor) -> Iterator[torch.Tensor]:
"""Yield finalized temporal chunks in decode order."""
tokens_chunk_size = self.tokens_chunk_size
chunk_num_frames = tokens_chunk_size * self.temporal_compression_ratio
pad_tokens, num_chunks, output_num_frames = self._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
output_frame_start = 0
decoded_chunks = []
overlap = None
for chunk_index in range(num_chunks):
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}"):
start = chunk_index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.0"):
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
for index in range(num_chunks):
start = index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
for overlap_index in range(int(token_drop > 0) + 1):
frame_start = overlap_index * chunk_num_frames
chunk = clip[:, :, frame_start:frame_start + chunk_num_frames]
chunk = chunk[:, :, self.frame_pre_padding:]
if overlap_index == 0:
if overlap is not None:
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
chunk = chunk[:, :, :num_frames]
decoded_chunks.append(chunk)
else:
overlap = chunk
if overlap is not None:
decoded_chunks.append(overlap)
decoded = torch.cat(decoded_chunks, dim=2)
next_overlap = None
if self.config.token_drop > 0:
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.1"):
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
# Yield after the ranges close so consumer-side CPU copies do not inflate decoder timing.
overlap = next_overlap
if num_frames > 0:
output_frame_start += num_frames
yield chunk
if overlap is not None and output_frame_start < output_num_frames:
yield overlap[:, :, :output_num_frames - output_frame_start]
def decoded_pixel_shape(self, latent_shape: torch.Size | tuple[int, ...]) -> tuple[int, int, int, int, int]:
"""Return the exact CPU pixel-buffer shape for a latent tensor shape."""
if len(latent_shape) != 5:
raise ValueError(f"MiniMax-H3 latents must be five-dimensional, got shape {tuple(latent_shape)}.")
batch_size, channels, latent_num_frames, latent_height, latent_width = map(int, latent_shape)
if channels != self.latent_channels:
raise ValueError(f"MiniMax-H3 latents must have {self.latent_channels} channels, got {channels}.")
_, _, decoded_num_frames = self._temporal_decode_plan(latent_num_frames)
return (
batch_size,
int(self.config.out_channels),
decoded_num_frames,
latent_height * self.spatial_compression_ratio,
latent_width * self.spatial_compression_ratio,
)
@staticmethod
def _streams_chunk_copies(z: torch.Tensor, output: torch.Tensor) -> bool:
"""Whether finalized chunks copy to ``output`` asynchronously on the current CUDA stream."""
return z.device.type == "cuda" and output.is_pinned()
@staticmethod
def _copy_chunk_pixels(pixels: torch.Tensor, output: torch.Tensor, frame_start: int, non_blocking: bool) -> None:
"""Copy one finalized fp32 pixel chunk into the CPU ``output`` buffer.
Device-to-host copies run per (batch, channel) plane: the temporal
slice of ``output`` is strided across channels, but each plane is
contiguous on both sides, so every transfer stays a direct memcpy
instead of staging through a pageable CPU temporary. With a pinned
``output`` and ``non_blocking=True`` the copies are additionally
asynchronous on the current CUDA stream; callers synchronize once
before releasing the buffer.
"""
target = output[:, :, frame_start:frame_start + pixels.shape[2]]
if pixels.device.type == "cuda":
pixels = pixels.contiguous()
for batch_index in range(pixels.shape[0]):
for channel_index in range(pixels.shape[1]):
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
non_blocking=non_blocking)
else:
target.copy_(pixels)
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
Each finalized chunk streams through ``_copy_chunk_pixels`` (direct
per-plane memcpys; asynchronous with a pinned ``output``) so the
copies overlap the next chunk's decode; ``decode_to_pixels``
synchronizes once before returning.
"""
non_blocking = self._streams_chunk_copies(z, output)
output_frame_start = 0
for chunk in self._decode_chunks(z):
num_frames = chunk.shape[2]
pixels = self.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._copy_chunk_pixels(pixels, output, output_frame_start, non_blocking)
output_frame_start += num_frames
if output_frame_start != output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {output_frame_start} frames into an output buffer expecting "
f"{output.shape[2]}.")
def _decode(self, z: torch.Tensor) -> torch.Tensor:
return torch.cat(list(self._decode_chunks(z)), dim=2)
if pad_tokens > 0:
intra_tail = self.config.clip_length % temporal_ratio
num_tokens_before_pad = z.shape[2] - pad_tokens
pad_frames = sum(intra_tail if intra_tail and (num_tokens_before_pad + offset) %
tokens_chunk_size == 0 else temporal_ratio for offset in range(pad_tokens))
decoded = decoded[:, :, :-pad_frames]
return decoded
def encode(
self,
@@ -992,34 +799,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
return (posterior, )
return AutoencoderKLOutput(latent_dist=posterior)
def encode_pixels(
self,
pixels: torch.Tensor,
return_dict: bool = True,
) -> AutoencoderKLOutput | tuple[DiagonalGaussianDistribution]:
"""Encode CPU-resident pixels one VAE clip at a time.
``pixels`` stays on CPU as ``uint8`` in ``[0, 255]`` or floating point
in ``[0, 1]``; each clip is moved to the VAE device, normalized, and
encoded so only one clip of pixels is resident on the accelerator.
"""
if pixels.ndim != 5 or pixels.shape[1] != self.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {self.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if self.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([self._encode_pixels(pixel_slice) for pixel_slice in pixels.split(1)])
else:
moments = self._encode_pixels(pixels)
posterior = DiagonalGaussianDistribution(moments)
if not return_dict:
return (posterior, )
return AutoencoderKLOutput(latent_dist=posterior)
def encode_keyframe(
self,
x: torch.Tensor,
@@ -1046,26 +825,6 @@ class AutoencoderKLMiniMaxH3(nn.Module):
return (decoded, )
return DecoderOutput(sample=decoded)
def decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> torch.Tensor:
"""Stream decoded ``[0, 1]`` FP32 pixels into a caller-owned CPU buffer."""
expected_shape = self.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
try:
if self.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
self._decode_to_pixels(z_slice, output[batch_index:batch_index + 1])
else:
self._decode_to_pixels(z, output)
finally:
# Drain async chunk copies before the caller (or an exception
# handler) can read or release the pinned buffer.
if self._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def forward(
self,
sample: torch.Tensor,
@@ -277,7 +277,7 @@ class LTX2Pipeline(LoRAPipeline):
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
component_model_path = self._component_path(module_name, fastvideo_args)
if module_name == "tokenizer" and not os.path.isdir(component_model_path):
gemma_path = os.path.join(self.model_path, "text_encoder", "gemma")
if os.path.isdir(gemma_path):
@@ -49,6 +49,16 @@ class LTX2AudioDecodingStage(PipelineStage):
audio_latents = batch.extra.get("ltx2_audio_latents")
if audio_latents is None:
return batch
# A checkpoint that declares no audio decoder builds these as None
# (``get_module`` returns its default), while the transformer is still
# audio-video and still produces audio latents. Having latents is
# therefore not evidence that anything can decode them -- guard on the
# modules too, and leave the video path unaffected.
if self.audio_decoder is None or self.vocoder is None:
logger.info(
"Skipping audio decoding: this checkpoint declares no audio "
"decoder/vocoder. Video output is unaffected.")
return batch
device = get_local_torch_device()
self.audio_decoder = self.audio_decoder.to(device)
@@ -12,9 +12,9 @@ from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_IMAGE_PAD_TOKEN,
MINIMAX_H3_TEXT_ENCODER_LAYER,
MINIMAX_H3_TEXT_TAG,
MINIMAX_H3_VIDEO_PAD_TOKEN,
MINIMAX_H3_VIDEO_TAG,
@@ -42,6 +42,25 @@ def _token_ids(tokenized: Any) -> list[int]:
return [int(token_id) for token_id in input_ids]
def _create_mm_token_type_ids(processor: Any, token_ids: list[int]) -> list[list[int]]:
"""Build Qwen3-VL modality IDs across old and new Transformers releases."""
create_ids = getattr(processor, "create_mm_token_type_ids", None)
if callable(create_ids):
return create_ids([token_ids])
modality_ids = [0] * len(token_ids)
for modality, modality_type in (("image", 1), ("video", 2), ("audio", 3)):
special_ids = getattr(processor, f"{modality}_token_ids", None)
if special_ids is None:
special_id = getattr(processor, f"{modality}_token_id", None)
special_ids = [] if special_id is None else [special_id]
resolved_ids = {int(special_id) for special_id in special_ids if special_id is not None}
for index, token_id in enumerate(token_ids):
if token_id in resolved_ids:
modality_ids[index] = modality_type
return [modality_ids]
def build_ref2va_presentation(
tokenizer: Any,
prompt: str,
@@ -136,10 +155,20 @@ class MiniMaxH3ConditioningStage(PipelineStage):
device: torch.device,
**vision_inputs: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
hidden_state_index = MINIMAX_H3_TEXT_ENCODER_LAYER
input_ids = torch.tensor([token_ids], dtype=torch.long, device=device)
mm_token_type_ids = torch.as_tensor(
_create_mm_token_type_ids(self.processor, token_ids),
dtype=torch.long,
device=device,
)
dtype = self.conditioner.dtype
prompt_embeds = self.conditioner(
input_ids,
outputs = self.conditioner(
input_ids=input_ids,
attention_mask=torch.ones_like(input_ids),
mm_token_type_ids=mm_token_type_ids,
use_cache=False,
output_hidden_states=True,
**{
name:
None if value is None else value.to(
@@ -149,10 +178,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
for name, value in vision_inputs.items()
},
)
if prompt_embeds.ndim != 2 or prompt_embeds.shape[0] != len(token_ids):
raise ValueError(f"MiniMax-H3 slim text encoder returned unexpected shape={tuple(prompt_embeds.shape)}")
if outputs.hidden_states is None or len(outputs.hidden_states) <= hidden_state_index:
raise ValueError(f"Qwen3-VL did not return `hidden_states[{hidden_state_index}]`.")
return (
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
torch.tensor(token_tags, dtype=torch.long),
)
@@ -257,7 +286,6 @@ class MiniMaxH3ConditioningStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Encode one H3 prompt presentation and attach its packed text features."""
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
@@ -265,13 +293,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
if moved_for_forward:
self.conditioner.to(device)
try:
# Keep both H3 prompt-presentation modes under one text-encoding
# range so Nsight Systems exposes their complete conditioning cost.
with nvtx_range("minimax_h3.text_encoding"):
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
finally:
if moved_for_forward:
self.conditioner.to("cpu")

Some files were not shown because too many files have changed in this diff Show More