Compare commits
49
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a861031c7 | ||
|
|
f08c5ee8af | ||
|
|
755f4a7967 | ||
|
|
d543a67b10 | ||
|
|
741aa8d289 | ||
|
|
ad9cd63122 | ||
|
|
68e6ffca9e | ||
|
|
64cdcf6be4 | ||
|
|
aa0d98a6b8 | ||
|
|
93b03bc14d | ||
|
|
160f0c9ccf | ||
|
|
e5d1110a0f | ||
|
|
622217ff2a | ||
|
|
cbab605eff | ||
|
|
ac98869aa1 | ||
|
|
56d4a6074f | ||
|
|
9713ea1275 | ||
|
|
37aa382cce | ||
|
|
3f00983287 | ||
|
|
2dc57f4070 | ||
|
|
9df19be719 | ||
|
|
089eea3970 | ||
|
|
dd8447ecc5 | ||
|
|
dca423fd31 | ||
|
|
1b43af8e8e | ||
|
|
628591b620 | ||
|
|
0980ca563f | ||
|
|
b158388733 | ||
|
|
c4ad4227c0 | ||
|
|
a63ccce73d | ||
|
|
ac56806aff | ||
|
|
0462e1b0e7 | ||
|
|
907f2100ec | ||
|
|
e0a3db5651 | ||
|
|
fca45bc8e1 | ||
|
|
aa95a4c18e | ||
|
|
942f7db3db | ||
|
|
aadb23f409 | ||
|
|
f56f567042 | ||
|
|
15a164a052 | ||
|
|
aaaa7a14a3 | ||
|
|
86d639c848 | ||
|
|
00338aa9ca | ||
|
|
74b409d7cf | ||
|
|
8537dcd6de | ||
|
|
528cef02c4 | ||
|
|
3c3da4d057 | ||
|
|
0399713e7b | ||
|
|
0af2e9e8ef |
@@ -0,0 +1,149 @@
|
||||
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
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
@@ -16,6 +17,7 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
**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/).
|
||||
@@ -62,6 +63,11 @@ 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
@@ -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=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
|
||||
|
||||
# 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
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
{
|
||||
"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"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
(() => {
|
||||
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);
|
||||
})();
|
||||
@@ -42,6 +42,46 @@ 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;
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,128 @@
|
||||
# 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,6 +76,10 @@ 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."
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
# 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
|
||||
@@ -19,6 +20,40 @@ 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:
|
||||
@@ -536,6 +571,7 @@ 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!")
|
||||
@@ -549,6 +585,7 @@ 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!")
|
||||
|
||||
@@ -65,6 +65,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Quick Start](quick_start.md) - Generate your first video
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
|
||||
|
||||
@@ -49,10 +49,12 @@ brew install ffmpeg
|
||||
|
||||
### Installation
|
||||
|
||||
FastWan's native Apple Silicon runtime requires the `mlx` extra.
|
||||
|
||||
#### With uv (recommended)
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
uv pip install "fastvideo[mlx]"
|
||||
```
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
@@ -60,7 +62,7 @@ uv pip install fastvideo
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
uv pip install "fastvideo[mlx]"
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -76,13 +78,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install -e ".[mlx]"
|
||||
```
|
||||
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install -e ".[mlx]"
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -23,61 +23,21 @@ Also optionally install flash-attn:
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
## Choose a maintained recipe
|
||||
|
||||
### Text-to-Video Generation
|
||||
The cookbook selects complete, checked-in recipes instead of mixing model,
|
||||
parallelism, offload, and attention settings independently.
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
|
||||
|
||||
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()
|
||||
```
|
||||
!!! 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.
|
||||
|
||||
## 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
|
||||
|
||||
@@ -178,6 +178,17 @@ 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
|
||||
|
||||
@@ -33,6 +33,14 @@ 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.
|
||||
|
||||
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.
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# 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()
|
||||
@@ -27,7 +27,8 @@ 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 converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.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,7 +27,8 @@ 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 converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.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,7 +24,8 @@ 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 converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.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.
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,510 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,130 @@
|
||||
"""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
@@ -0,0 +1,365 @@
|
||||
"""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()
|
||||
@@ -0,0 +1,104 @@
|
||||
"""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()
|
||||
@@ -318,6 +318,38 @@ 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}
|
||||
)
|
||||
@@ -333,10 +365,14 @@ 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
|
||||
|
||||
+40
-11
@@ -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 (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| 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 |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
@@ -64,28 +64,49 @@ cd fastvideo-kernel
|
||||
./build.sh --rocm
|
||||
```
|
||||
|
||||
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
|
||||
### Optional: FA4 CuTe block-sparse backend (VSA-128/256 fastpath)
|
||||
|
||||
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
|
||||
The VSA-128/256 fastpaths (tile volume 128 or 256, on NVIDIA Blackwell / sm_100) route 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`,
|
||||
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
|
||||
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
|
||||
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
|
||||
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
|
||||
commit `14c377950125c70b7a9dabf9c561fca53715ac7d`, the revision FastVideo pins as
|
||||
the `flash-attn-4` source in the repo-root `pyproject.toml`; other revisions may
|
||||
have an incompatible `_flash_attn_fwd` signature.
|
||||
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.
|
||||
|
||||
```bash
|
||||
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"
|
||||
pip install torchvision
|
||||
pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@14c377950125c70b7a9dabf9c561fca53715ac7d#subdirectory=flash_attn/cute"
|
||||
```
|
||||
|
||||
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
|
||||
`tests/test_vsa256_forward*.py`.
|
||||
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`.
|
||||
|
||||
## Usage
|
||||
|
||||
@@ -142,6 +163,14 @@ 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
|
||||
|
||||
@@ -2,8 +2,9 @@
|
||||
"""
|
||||
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
|
||||
|
||||
This script benchmarks the autograd-enabled wrapper:
|
||||
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
|
||||
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
|
||||
|
||||
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
|
||||
"""
|
||||
@@ -23,9 +24,6 @@ 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)
|
||||
@@ -41,7 +39,11 @@ 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 /64)")
|
||||
p.add_argument("--q_seq_lens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[49152],
|
||||
help="Q sequence lengths (must be divisible by --block_size)")
|
||||
p.add_argument("--kv_seq_lens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
@@ -51,9 +53,13 @@ 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()
|
||||
|
||||
|
||||
@@ -84,18 +90,38 @@ 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
|
||||
@@ -105,20 +131,22 @@ 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_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
|
||||
print(f"block_size={block_size}")
|
||||
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
|
||||
if args.force_triton:
|
||||
if args.use_cute:
|
||||
print("dispatch: FA4 CuTe")
|
||||
elif 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_M != 0 or kv_len % BLOCK_N != 0:
|
||||
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
|
||||
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}")
|
||||
continue
|
||||
|
||||
num_q_blocks = q_len // BLOCK_M
|
||||
num_kv_blocks = kv_len // BLOCK_N
|
||||
num_q_blocks = q_len // block_size
|
||||
num_kv_blocks = kv_len // block_size
|
||||
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
|
||||
topk = min(topk, num_kv_blocks)
|
||||
|
||||
@@ -129,11 +157,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 blocks (64 tokens per KV block)
|
||||
variable_block_sizes = torch.full((num_kv_blocks, ), BLOCK_N, dtype=torch.int32, device="cuda")
|
||||
# Variable block sizes: default full logical blocks.
|
||||
variable_block_sizes = torch.full((num_kv_blocks, ), block_size, dtype=torch.int32, device="cuda")
|
||||
|
||||
def _fwd():
|
||||
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
|
||||
return attention(q, k, v, block_map, variable_block_sizes)
|
||||
|
||||
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
|
||||
|
||||
@@ -142,7 +170,7 @@ def main() -> None:
|
||||
q_ = q.detach().requires_grad_(True)
|
||||
k_ = k.detach().requires_grad_(True)
|
||||
v_ = v.detach().requires_grad_(True)
|
||||
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
|
||||
o_, _aux_ = attention(q_, k_, v_, block_map, variable_block_sizes)
|
||||
og = torch.randn_like(o_)
|
||||
loss = (o_ * og).sum()
|
||||
|
||||
@@ -156,7 +184,7 @@ def main() -> None:
|
||||
rep=max(5, args.rep // 2),
|
||||
)
|
||||
|
||||
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
|
||||
flops = flops_sparse_attention(bs, h, d, q_len, topk, block_size)
|
||||
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
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
// 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
@@ -0,0 +1,201 @@
|
||||
#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
|
||||
@@ -0,0 +1,114 @@
|
||||
// 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};
|
||||
}
|
||||
@@ -0,0 +1,877 @@
|
||||
// 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,10 +28,31 @@ 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-256 block-sparse attention wrapper.
|
||||
"""VSA-128/256 block-sparse attention wrappers.
|
||||
|
||||
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 KV blocks (this wrapper expands the logical 256-block map /
|
||||
sizes into that physical 128-block representation). The CuTe kernel
|
||||
on 128-token Q/KV blocks (the 256 wrapper expands its logical KV map and
|
||||
sizes into that physical 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 256-block VSA path.
|
||||
"""Pick the backend for the 128/256-block VSA paths.
|
||||
|
||||
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,6 +49,26 @@ 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,
|
||||
@@ -112,6 +132,63 @@ 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 @@
|
||||
"""CuTe-DSL block-sparse attention forward kernel.
|
||||
"""FA4 CuTe-DSL block-sparse attention adapter.
|
||||
|
||||
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.
|
||||
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.
|
||||
|
||||
Both [B, H, S, D] (BHSD) and [B, S, H, D] (BSHD) entrypoints are provided.
|
||||
The BSHD variant is preferred from VSA-256 callers to avoid layout
|
||||
The BSHD variant is preferred from VSA-128/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-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``.
|
||||
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``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -22,13 +22,14 @@ from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
|
||||
_FA4_IMPORT_HINT = ("VSA-128/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 VSA-256 path is Triton. Install the FA4 CuTe "
|
||||
"dependency; the default 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.
|
||||
|
||||
@@ -38,14 +39,39 @@ def _load_fa4_cute():
|
||||
"""
|
||||
try:
|
||||
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
|
||||
from flash_attn.cute.interface import _flash_attn_fwd
|
||||
from flash_attn.cute.interface import (
|
||||
_flash_attn_bwd,
|
||||
_flash_attn_fwd,
|
||||
flash_attn_func,
|
||||
)
|
||||
except ImportError as exc: # pragma: no cover - optional dependency
|
||||
raise ImportError(_FA4_IMPORT_HINT) from exc
|
||||
return BlockSparseTensorsTorch, _flash_attn_fwd
|
||||
return BlockSparseTensorsTorch, flash_attn_func, _flash_attn_fwd, _flash_attn_bwd
|
||||
|
||||
|
||||
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
|
||||
_M_BLOCK_SIZE_DEFAULT = 128
|
||||
# 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)
|
||||
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
@@ -64,12 +90,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, 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.
|
||||
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.
|
||||
major, _ = torch.cuda.get_device_capability()
|
||||
if major >= 10 and q_len > m_block_size:
|
||||
return 2 * m_block_size
|
||||
return m_block_size
|
||||
if major >= 10 and q_len > q_tile_size:
|
||||
return 2 * q_tile_size
|
||||
return q_tile_size
|
||||
|
||||
|
||||
def _aggregate_q_block_map(
|
||||
@@ -134,23 +160,35 @@ def _build_vbs_mask_mod(kv_block_size: int):
|
||||
return _vbs_mask_mod
|
||||
|
||||
|
||||
def _cute_forward(
|
||||
q_bshd: torch.Tensor,
|
||||
k_bshd: torch.Tensor,
|
||||
v_bshd: torch.Tensor,
|
||||
def _build_sparse_tensors(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
*,
|
||||
q_len: int,
|
||||
q_block_size: int,
|
||||
kv_block_size: int,
|
||||
) -> 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,
|
||||
)
|
||||
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")
|
||||
sparse_map = _aggregate_q_block_map(
|
||||
block_map,
|
||||
q_sparse_block_size=q_sparse_block_size,
|
||||
@@ -158,35 +196,166 @@ def _cute_forward(
|
||||
)
|
||||
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
|
||||
|
||||
full_block_idx, full_block_cnt = _map_to_index(full_map)
|
||||
mask_block_idx, mask_block_cnt = _map_to_index(mask_map)
|
||||
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),
|
||||
)
|
||||
|
||||
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),
|
||||
forward_sparse_tensors = from_maps(
|
||||
sparse_map & kv_full,
|
||||
sparse_map & kv_partial,
|
||||
)
|
||||
|
||||
# _flash_attn_fwd returns (out, lse, p, row_max); keep the first two.
|
||||
out, lse = _flash_attn_fwd(
|
||||
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(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
|
||||
mask_mod=_build_vbs_mask_mod(kv_block_size),
|
||||
block_sparse_tensors=sparse_tensors,
|
||||
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,
|
||||
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,
|
||||
@@ -194,34 +363,25 @@ def block_sparse_attn_cute_fwd(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""CuTe forward-only block-sparse attention with [B, H, S, D] inputs."""
|
||||
"""Autograd-enabled CuTe block-sparse attention for [B, H, S, D]."""
|
||||
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_bshd = _cute_forward(
|
||||
out_bshd, lse = _cute_attention(
|
||||
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()
|
||||
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
|
||||
# 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()
|
||||
|
||||
|
||||
def block_sparse_attn_cute_fwd_bshd(
|
||||
@@ -231,27 +391,16 @@ def block_sparse_attn_cute_fwd_bshd(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""CuTe forward-only block-sparse attention with [B, S, H, D] inputs."""
|
||||
"""Autograd-enabled CuTe block-sparse attention for [B, S, H, D]."""
|
||||
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_bshd = _cute_forward(
|
||||
out, lse = _cute_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
q_block_size=q_block_size,
|
||||
kv_block_size=kv_block_size,
|
||||
)
|
||||
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
|
||||
# lse is [B, H, S] regardless of the q/k/v layout; see above.
|
||||
return out, lse.detach()
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
# 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)
|
||||
@@ -2,6 +2,8 @@ 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,
|
||||
)
|
||||
@@ -74,12 +76,13 @@ 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 256-block path.
|
||||
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 128/256-block paths.
|
||||
"""
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
@@ -119,8 +122,9 @@ def video_sparse_attn(
|
||||
# Sparse branch (fused Triton topk mask)
|
||||
mask = fused_topk_mask(scores, topk)
|
||||
|
||||
if block_elements == 256:
|
||||
out_s = block_sparse_attn_256(q, k, v, mask, variable_block_sizes)[0]
|
||||
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]
|
||||
else:
|
||||
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
|
||||
@@ -142,14 +146,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 256-block path; the 64-block path still expects BHSD and is not
|
||||
the CuTe 128/256-block paths; 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 != 256:
|
||||
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=256 "
|
||||
if block_elements not in (128, 256):
|
||||
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=128 or 256 "
|
||||
f"(got {block_elements}); use video_sparse_attn for the 64-block path.")
|
||||
|
||||
batch, q_seq_len, heads, dim = q.shape
|
||||
@@ -171,19 +175,15 @@ 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: 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)
|
||||
|
||||
# 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.
|
||||
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() * 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_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_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,13 +195,15 @@ def video_sparse_attn_bshd(
|
||||
|
||||
# Sparse branch (fused Triton topk mask + CuTe BSHD).
|
||||
mask = fused_topk_mask(scores, topk)
|
||||
out_s, _ = block_sparse_attn_256_bshd(q, k, v, mask, variable_block_sizes)
|
||||
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 = out_s
|
||||
out_view = out.view(batch, q_num_blocks, block_elements, heads, dim)
|
||||
# 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)
|
||||
if compress_attn_weight is not None:
|
||||
gate_view = compress_attn_weight.view(batch, q_num_blocks, block_elements, heads, dim)
|
||||
out_view.add_(out_c_blk.unsqueeze(2) * gate_view)
|
||||
out = out_view + out_c_blk.unsqueeze(2) * gate_view
|
||||
else:
|
||||
out_view.add_(out_c_blk.unsqueeze(2))
|
||||
return out
|
||||
out = out_view + out_c_blk.unsqueeze(2)
|
||||
return out.view(batch, q_seq_len, heads, dim)
|
||||
|
||||
+19
-8
@@ -237,7 +237,12 @@ 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)
|
||||
qkT = tl.dot(k, qT)
|
||||
# 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)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
@@ -268,6 +273,7 @@ def _attn_bwd_dq(
|
||||
do,
|
||||
m,
|
||||
D,
|
||||
sm_scale,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
@@ -315,7 +321,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)
|
||||
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
|
||||
p = tl.math.exp2(qk - m)
|
||||
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
|
||||
mask = offs_in_block < block_size
|
||||
@@ -324,8 +330,7 @@ def _attn_bwd_dq(
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
@@ -453,6 +458,7 @@ def _attn_bwd(
|
||||
do,
|
||||
m,
|
||||
D, #
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
@@ -470,7 +476,7 @@ def _attn_bwd(
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
dq *= sm_scale
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
@@ -591,6 +597,7 @@ def _attn_bwd_dq_kernel(
|
||||
Q,
|
||||
K,
|
||||
V,
|
||||
sm_scale,
|
||||
DO, #
|
||||
DQ,
|
||||
M,
|
||||
@@ -663,6 +670,7 @@ def _attn_bwd_dq_kernel(
|
||||
do,
|
||||
m,
|
||||
D,
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
@@ -680,7 +688,7 @@ def _attn_bwd_dq_kernel(
|
||||
)
|
||||
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq_acc *= LN2
|
||||
dq_acc *= sm_scale
|
||||
tl.store(dq_ptrs, dq_acc)
|
||||
|
||||
|
||||
@@ -748,9 +756,11 @@ 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
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
# 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.)
|
||||
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)
|
||||
@@ -813,6 +823,7 @@ 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,7 +14,9 @@ import math
|
||||
import torch
|
||||
|
||||
VSA_TILE_SIZE = (4, 4, 4)
|
||||
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 256)
|
||||
# 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)
|
||||
|
||||
|
||||
def _canonicalize_device(device: torch.device | str) -> torch.device:
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
"""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)
|
||||
@@ -0,0 +1,224 @@
|
||||
"""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,6 +22,7 @@ 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()
|
||||
@@ -55,6 +56,8 @@ 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
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
"""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}"
|
||||
@@ -23,6 +23,20 @@ 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,
|
||||
|
||||
@@ -19,7 +19,12 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
# 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)
|
||||
|
||||
# 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,13 +5,17 @@ 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 (4,8,8) video tiles]``;
|
||||
prefix tiles never straddle segment boundaries.
|
||||
- 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``).
|
||||
- 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 H3
|
||||
checkpoint does not carry: the loader zero-initializes it, so untrained
|
||||
inference is exactly pure sparse and finetuning can learn the gate.
|
||||
- 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.
|
||||
- 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,
|
||||
@@ -20,22 +24,46 @@ 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.
|
||||
|
||||
Targets sm10.x through the FA4 CuTe 256-tile path
|
||||
At tile 256 this 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.
|
||||
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.
|
||||
"""
|
||||
|
||||
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)
|
||||
@@ -43,51 +71,115 @@ 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
|
||||
|
||||
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
|
||||
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)
|
||||
_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) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
|
||||
tile_elems: int = _TILE_ELEMS) -> 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 = VSA_H3_TILE_SIZE
|
||||
ts_t, ts_h, ts_w = tile_shape
|
||||
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, VSA_H3_TILE_SIZE)
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
|
||||
num_video_tiles = int(video_sizes.numel())
|
||||
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
|
||||
tile_partition_indices = torch.cat([
|
||||
torch.arange(prefix_len, device=device, dtype=torch.long),
|
||||
video_indices,
|
||||
@@ -100,9 +192,11 @@ 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)
|
||||
|
||||
|
||||
@@ -139,6 +233,9 @@ 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
|
||||
@@ -158,24 +255,28 @@ 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, ...] = (),
|
||||
**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, ...] = (),
|
||||
tile_size: int = _TILE_ELEMS,
|
||||
**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)
|
||||
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
|
||||
|
||||
return MiniMaxH3VSAMetadata(
|
||||
current_timestep=current_timestep,
|
||||
@@ -186,13 +287,14 @@ 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) -> torch.Tensor:
|
||||
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
|
||||
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].
|
||||
|
||||
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
|
||||
written), so a plain sum with fp32 accumulation needs no validity mask
|
||||
@@ -200,8 +302,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Te
|
||||
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)
|
||||
|
||||
@@ -232,6 +334,24 @@ 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__(
|
||||
@@ -259,7 +379,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 * _TILE_ELEMS, x.shape[-2], x.shape[-1])
|
||||
target_shape = (x.shape[0], n_tiles * attn_metadata.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
|
||||
@@ -281,7 +401,11 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
gate_compress: torch.Tensor | None,
|
||||
attn_metadata: MiniMaxH3VSAMetadata,
|
||||
) -> torch.Tensor:
|
||||
if block_sparse_attn_256_bshd is None:
|
||||
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:
|
||||
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
|
||||
|
||||
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
|
||||
@@ -291,8 +415,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)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
|
||||
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
|
||||
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)
|
||||
@@ -309,18 +433,75 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
attn_metadata.exempt,
|
||||
)
|
||||
|
||||
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
|
||||
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)
|
||||
|
||||
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)
|
||||
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
|
||||
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, _, heads, dim = out.shape
|
||||
batch, seq_len, heads, dim = out.shape
|
||||
n_tiles = attn_metadata.variable_block_sizes.numel()
|
||||
out.view(batch, n_tiles, _TILE_ELEMS, heads,
|
||||
dim).addcmul_(out_c.unsqueeze(2), gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim))
|
||||
# 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)
|
||||
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)
|
||||
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
|
||||
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)]
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,375 @@
|
||||
# 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()
|
||||
@@ -0,0 +1,826 @@
|
||||
# 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()
|
||||
@@ -62,14 +62,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
hidden_size: int = 5120
|
||||
intermediate_size: int = 25600
|
||||
num_hidden_layers: int = 64
|
||||
# H3 conditions on one intermediate hidden state and reads nothing above it,
|
||||
# so the remaining layers are built, weight-loaded and then discarded: 14
|
||||
# layers, 13.7 GB in bf16. Building exactly this many leaves that hidden
|
||||
# state bit-identical, because the tuple records each layer's *input*, so
|
||||
# entry N is the output of layer N-1. Set to None to keep the full stack.
|
||||
# Must equal MINIMAX_H3_TEXT_ENCODER_LAYER in
|
||||
# fastvideo/pipelines/basic/minimax_h3/packing.py; a test pins them together
|
||||
# rather than importing across the models -> pipelines boundary.
|
||||
output_hidden_state_index: int = 50
|
||||
num_hidden_layers_override: int | None = 50
|
||||
num_attention_heads: int = 64
|
||||
num_key_value_heads: int = 8
|
||||
@@ -116,7 +109,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
vision_initializer_range: float = 0.02
|
||||
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
|
||||
|
||||
output_hidden_states: bool = True
|
||||
output_hidden_states: bool = False
|
||||
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,
|
||||
@@ -127,15 +120,16 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# Runs both at construction and after ``update_model_arch`` merges the
|
||||
# checkpoint's config.json, so it also guards config-file overrides. A
|
||||
# non-positive override would build no decoder layers at all, and a
|
||||
# negative one would additionally make the surplus-key filter drop
|
||||
# every ``language_model.layers.*`` checkpoint key, so the conditioner
|
||||
# would "load" with no transformer stack and only fail at generation.
|
||||
if self.num_hidden_layers_override is not None and self.num_hidden_layers_override < 1:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be a positive layer count "
|
||||
f"or None for the full stack; got {self.num_hidden_layers_override}.")
|
||||
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))
|
||||
|
||||
@@ -808,14 +808,19 @@ 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)
|
||||
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.
|
||||
# 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``.
|
||||
# ``skip_pixel_prealloc`` also gates the slow-path warning.
|
||||
skip_pixel_prealloc = is_latent_output or not needs_samples_buffer
|
||||
needs_samples_out = batch.return_frames
|
||||
skip_pixel_prealloc = is_latent_output or not needs_samples_out
|
||||
if skip_pixel_prealloc:
|
||||
samples = torch.empty(0, device='cpu')
|
||||
else:
|
||||
@@ -835,9 +840,11 @@ class VideoGenerator:
|
||||
"This usually means the executor/pipeline failed earlier.")
|
||||
|
||||
audio_only = bool(output_batch.extra.get("audio_only"))
|
||||
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.
|
||||
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.
|
||||
pass
|
||||
elif audio_only:
|
||||
# Audio-only return-frames requests expose the small placeholder
|
||||
@@ -869,8 +876,13 @@ 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(
|
||||
samples,
|
||||
output_batch.output if needs_frame_output else samples,
|
||||
(target_height, target_width, batch.num_frames),
|
||||
pixel_output=not is_latent_output and not audio_only,
|
||||
)
|
||||
@@ -882,13 +894,26 @@ class VideoGenerator:
|
||||
elif not needs_frame_output:
|
||||
frames = None
|
||||
else:
|
||||
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())
|
||||
# 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
|
||||
]
|
||||
postprocess_time = time.perf_counter() - postprocess_start
|
||||
logger.info("PostDecodeFrameProcessStage completed in %.3f s", postprocess_time)
|
||||
if logging_info is not None:
|
||||
|
||||
@@ -21,12 +21,17 @@ 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
|
||||
@@ -217,10 +222,34 @@ 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":
|
||||
|
||||
@@ -146,6 +146,19 @@ 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).
|
||||
@@ -169,6 +182,7 @@ 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
|
||||
@@ -286,8 +300,27 @@ 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``.
|
||||
|
||||
@@ -631,6 +664,18 @@ 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,
|
||||
@@ -644,6 +689,12 @@ 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(
|
||||
|
||||
+4
-1
@@ -114,7 +114,10 @@ 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):
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
# 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)
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -0,0 +1,273 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,228 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -0,0 +1,980 @@
|
||||
# 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
|
||||
@@ -0,0 +1,154 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -0,0 +1,243 @@
|
||||
# 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.")
|
||||
@@ -0,0 +1,129 @@
|
||||
# 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
|
||||
@@ -0,0 +1,454 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -0,0 +1,284 @@
|
||||
# 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
|
||||
@@ -0,0 +1,690 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,137 @@
|
||||
# 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
|
||||
@@ -0,0 +1,139 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,191 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,286 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,113 @@
|
||||
# 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
|
||||
@@ -0,0 +1,557 @@
|
||||
# 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,
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
# 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)
|
||||
@@ -10,6 +10,7 @@ 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
|
||||
@@ -23,12 +24,50 @@ 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):
|
||||
@@ -62,6 +101,7 @@ 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(
|
||||
@@ -78,11 +118,15 @@ 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)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
hidden_states = hidden_states * F.silu(gate)
|
||||
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, _ = self.fc_out(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
@@ -99,6 +143,7 @@ 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
|
||||
@@ -134,6 +179,7 @@ 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
|
||||
@@ -211,11 +257,18 @@ 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))
|
||||
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)
|
||||
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)
|
||||
|
||||
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
|
||||
# branch cannot express. Apply it above, then pass no RoPE here.
|
||||
@@ -397,6 +450,9 @@ 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)
|
||||
@@ -408,6 +464,7 @@ 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(
|
||||
@@ -415,6 +472,7 @@ 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,
|
||||
@@ -423,6 +481,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
prefix=f"{prefix}.adaln_proj",
|
||||
apply_silu=adaln_apply_silu,
|
||||
)
|
||||
self.fuse_modulate = fuse_modulate
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -435,19 +494,39 @@ 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))
|
||||
|
||||
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)
|
||||
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)
|
||||
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
|
||||
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)
|
||||
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)
|
||||
feed_forward_output = self.ff(norm_hidden_states)
|
||||
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
|
||||
|
||||
class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
@@ -493,6 +572,17 @@ 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 "
|
||||
@@ -546,7 +636,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 "
|
||||
"tools/minimax_h3/fit_adaln_basis.py.")
|
||||
"scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
|
||||
adaln_dim = self.adaln_rank or arch.time_embed_dim
|
||||
self.adaln_basis = ReplicatedLinear(
|
||||
arch.time_embed_dim,
|
||||
@@ -590,6 +680,9 @@ 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(
|
||||
@@ -616,6 +709,20 @@ 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,
|
||||
@@ -734,14 +841,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
|
||||
rotary_emb = (rotary_cos, rotary_sin)
|
||||
|
||||
for block in self.transformer_blocks:
|
||||
packed_hidden_states = block(
|
||||
packed_hidden_states,
|
||||
temb,
|
||||
adaln_indices,
|
||||
rotary_emb,
|
||||
original_seq_len,
|
||||
)
|
||||
# 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,
|
||||
)
|
||||
|
||||
packed_hidden_states = self.norm_out(
|
||||
packed_hidden_states,
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -0,0 +1,302 @@
|
||||
# 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)
|
||||
@@ -0,0 +1,174 @@
|
||||
# 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"]
|
||||
@@ -0,0 +1,104 @@
|
||||
# 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"]
|
||||
@@ -1,6 +1,7 @@
|
||||
# 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
|
||||
@@ -8,11 +9,16 @@ 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__()
|
||||
@@ -23,13 +29,7 @@ class TextEncoder(nn.Module, ABC):
|
||||
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
|
||||
|
||||
@abstractmethod
|
||||
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:
|
||||
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
|
||||
pass
|
||||
|
||||
@property
|
||||
|
||||
@@ -0,0 +1,453 @@
|
||||
# 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,19 +227,13 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
# Build only as far as the consumer reads. The hidden-state tuple records
|
||||
# each layer's input, so stopping after N layers still yields entry N,
|
||||
# the output of layer N-1, unchanged. Everything above it exists only to
|
||||
# feed `last_hidden_state`, which nothing consumes.
|
||||
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))
|
||||
# The final norm sits above the tapped layer, so a truncated stack drops
|
||||
# it. Keeping it would overwrite the tapped entry with a normalised
|
||||
# tensor and change conditioning without raising anything.
|
||||
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
if self.num_layers == config.num_hidden_layers else None)
|
||||
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
|
||||
@@ -249,18 +243,14 @@ 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,
|
||||
) -> BaseEncoderOutput:
|
||||
) -> torch.Tensor:
|
||||
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:
|
||||
@@ -269,13 +259,9 @@ 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 self.norm is not None:
|
||||
hidden_states = self.norm(hidden_states)
|
||||
# Truncated or not, the last entry is appended here, so the tapped index
|
||||
# lands in the same place either way.
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states, )
|
||||
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
|
||||
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}]")
|
||||
|
||||
|
||||
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
|
||||
@@ -513,10 +499,18 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
|
||||
return self.merger(hidden_states), deepstack_features
|
||||
|
||||
|
||||
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
|
||||
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
|
||||
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
|
||||
|
||||
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)
|
||||
@@ -530,15 +524,12 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
|
||||
@property
|
||||
def num_hidden_layers(self) -> int:
|
||||
"""The checkpoint architecture's nominal depth, matching its config.json.
|
||||
|
||||
When ``num_hidden_layers_override`` truncates the stack at the
|
||||
conditioning tap, fewer layers exist; the built count is
|
||||
``self.language_model.num_layers``, and the hidden-state tuple has
|
||||
``num_layers + 1`` entries, not ``num_hidden_layers + 1``.
|
||||
"""
|
||||
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,
|
||||
@@ -631,35 +622,39 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
f"tokens={int(mask.sum())}, features={features.shape[0]}")
|
||||
return mask
|
||||
|
||||
def forward(
|
||||
# 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(
|
||||
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,
|
||||
input_ids: torch.Tensor,
|
||||
*,
|
||||
pixel_values: torch.Tensor | None = None,
|
||||
pixel_values_videos: 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,
|
||||
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")
|
||||
) -> 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)
|
||||
|
||||
image_mask = None
|
||||
video_mask = None
|
||||
image_deepstack = None
|
||||
video_deepstack = None
|
||||
if pixel_values is not None:
|
||||
if input_ids is None or image_grid_thw is None:
|
||||
if 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)
|
||||
@@ -667,7 +662,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
"image")
|
||||
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
|
||||
if pixel_values_videos is not None:
|
||||
if input_ids is None or video_grid_thw is None:
|
||||
if 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)
|
||||
@@ -695,50 +690,34 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
visual_mask = video_mask
|
||||
deepstack_features = video_deepstack
|
||||
|
||||
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(
|
||||
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
|
||||
hidden_states = self.language_model(
|
||||
inputs_embeds,
|
||||
position_ids,
|
||||
attention_mask,
|
||||
output_hidden_states,
|
||||
None,
|
||||
visual_mask,
|
||||
deepstack_features,
|
||||
)
|
||||
outputs.attention_mask = attention_mask
|
||||
return outputs
|
||||
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 _is_above_the_tap(self, name: str) -> bool:
|
||||
"""Whether this checkpoint key belongs to a layer we did not build.
|
||||
|
||||
A truncated language stack still ships every layer in the checkpoint, and
|
||||
the unexpected-key check below is strict on purpose, so the surplus keys
|
||||
have to be dropped here rather than by relaxing it.
|
||||
"""
|
||||
language_model = self.language_model
|
||||
# The final norm is dropped exactly when the stack is truncated, so its
|
||||
# absence is the signal.
|
||||
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]
|
||||
if not index.isdigit():
|
||||
return False
|
||||
# Only drop indexes the full stack would have built. Anything at or
|
||||
# above the checkpoint's own num_hidden_layers is corrupt and must
|
||||
# still raise below, exactly as it does without truncation.
|
||||
return language_model.num_layers <= int(index) < self.config.num_hidden_layers
|
||||
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,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
parameters = dict(self.named_parameters())
|
||||
@@ -748,7 +727,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
if source_name == "lm_head.weight":
|
||||
continue
|
||||
name = source_name[6:] if source_name.startswith("model.") else source_name
|
||||
if self._is_above_the_tap(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}")
|
||||
@@ -758,7 +737,23 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
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"]
|
||||
__all__ = [
|
||||
"MiniMaxH3Qwen3VLConditioner",
|
||||
"MiniMaxH3SerializedFP8Config",
|
||||
]
|
||||
|
||||
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
|
||||
from collections.abc import Generator, Iterable
|
||||
from contextlib import nullcontext
|
||||
from copy import deepcopy
|
||||
from typing import cast
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
@@ -30,9 +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.utils import set_default_torch_dtype
|
||||
from fastvideo.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files,
|
||||
@@ -347,22 +351,46 @@ class TextEncoderLoader(ComponentLoader):
|
||||
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,
|
||||
@@ -381,11 +409,20 @@ 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):
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
safetensors_weights_iterator(
|
||||
[fastvideo_args.override_text_encoder_safetensors],
|
||||
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],
|
||||
to_cpu=use_cpu_offload,
|
||||
)) # type: ignore
|
||||
)
|
||||
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
|
||||
else:
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
self._get_all_weights(
|
||||
@@ -400,6 +437,10 @@ 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)
|
||||
|
||||
@@ -442,7 +483,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:
|
||||
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
|
||||
raise ValueError("Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
|
||||
@@ -1057,7 +1098,12 @@ 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)
|
||||
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
|
||||
# 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)
|
||||
model = maybe_load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
init_params={
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# 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
|
||||
@@ -392,9 +392,16 @@ 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
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -7,6 +7,7 @@ This module intentionally uses only PyTorch and FastVideo configuration types.
|
||||
"""
|
||||
|
||||
import math
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
@@ -14,7 +15,10 @@ 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:
|
||||
@@ -291,6 +295,7 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
|
||||
class MiniMaxH3VideoAttention(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
|
||||
"""Build projections and the selected dense FastVideo attention implementation."""
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
@@ -302,12 +307,34 @@ class MiniMaxH3VideoAttention(nn.Module):
|
||||
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
|
||||
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
|
||||
self.attn_impl = None
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
if current_platform.is_cuda_alike():
|
||||
attention_backend = get_attn_backend(
|
||||
dim_head,
|
||||
# FlashAttention executes the FP32 VAE activations in BF16 and
|
||||
# restores FP32 output, so resolve against the kernel dtype.
|
||||
torch.bfloat16,
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
),
|
||||
)
|
||||
self.attn_impl = attention_backend.get_impl_cls()(
|
||||
num_heads=heads,
|
||||
head_size=dim_head,
|
||||
softmax_scale=dim_head**-0.5,
|
||||
num_kv_heads=heads,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply dense self-attention to one spatial VAE token sequence."""
|
||||
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
|
||||
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
|
||||
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
|
||||
@@ -328,9 +355,17 @@ class MiniMaxH3VideoAttention(nn.Module):
|
||||
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
|
||||
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
|
||||
|
||||
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
|
||||
if self.attn_impl is not None and query.device.type != "cpu":
|
||||
# VAE decoding has no diffusion-step metadata, so call the selected
|
||||
# backend implementation directly with dense BSHD tensors.
|
||||
hidden_states = self.attn_impl.forward(query, key, value, None)
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
else:
|
||||
# Keep CPU construction and execution available without requiring
|
||||
# an accelerator attention backend.
|
||||
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
|
||||
return self.to_out[0](hidden_states)
|
||||
|
||||
|
||||
@@ -433,6 +468,7 @@ 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,
|
||||
@@ -482,6 +518,11 @@ 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."""
|
||||
|
||||
@@ -489,6 +530,7 @@ 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__()
|
||||
@@ -654,12 +696,15 @@ 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 = []
|
||||
@@ -676,6 +721,12 @@ 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))
|
||||
@@ -699,36 +750,64 @@ 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]
|
||||
return self._stitch_tiles(rows, latent_y_overlaps, latent_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()
|
||||
|
||||
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
|
||||
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)
|
||||
"""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()
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
clip_length = self.config.clip_length
|
||||
@@ -747,43 +826,157 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
moments = moments[:, :, :-self.config.token_drop]
|
||||
return moments
|
||||
|
||||
def _decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
tokens_chunk_size = self.tokens_chunk_size
|
||||
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
|
||||
tokens_chunk_size = self.tokens_chunk_size
|
||||
temporal_ratio = self.temporal_compression_ratio
|
||||
chunk_num_frames = tokens_chunk_size * temporal_ratio
|
||||
num_tokens = z.shape[2] + token_drop
|
||||
num_tokens = latent_num_frames + 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)
|
||||
|
||||
decoded_chunks = []
|
||||
output_frame_start = 0
|
||||
overlap = None
|
||||
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:
|
||||
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]
|
||||
if overlap is not None:
|
||||
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
|
||||
decoded_chunks.append(chunk)
|
||||
else:
|
||||
overlap = chunk
|
||||
if overlap is not None:
|
||||
decoded_chunks.append(overlap)
|
||||
decoded = torch.cat(decoded_chunks, dim=2)
|
||||
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
|
||||
chunk = chunk[:, :, :num_frames]
|
||||
|
||||
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
|
||||
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)
|
||||
|
||||
def encode(
|
||||
self,
|
||||
@@ -799,6 +992,34 @@ 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,
|
||||
@@ -825,6 +1046,26 @@ 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,
|
||||
|
||||
@@ -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,25 +42,6 @@ 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,
|
||||
@@ -155,20 +136,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
device: torch.device,
|
||||
**vision_inputs: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
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,
|
||||
)
|
||||
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
|
||||
dtype = self.conditioner.dtype
|
||||
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,
|
||||
prompt_embeds = self.conditioner(
|
||||
input_ids,
|
||||
**{
|
||||
name:
|
||||
None if value is None else value.to(
|
||||
@@ -178,10 +149,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
for name, value in vision_inputs.items()
|
||||
},
|
||||
)
|
||||
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}]`.")
|
||||
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)}")
|
||||
return (
|
||||
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
|
||||
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
|
||||
torch.tensor(token_tags, dtype=torch.long),
|
||||
)
|
||||
|
||||
@@ -286,6 +257,7 @@ 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
|
||||
@@ -293,10 +265,13 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
if moved_for_forward:
|
||||
self.conditioner.to(device)
|
||||
try:
|
||||
if self.ref2va:
|
||||
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
|
||||
else:
|
||||
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
|
||||
# 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)
|
||||
finally:
|
||||
if moved_for_forward:
|
||||
self.conditioner.to("cpu")
|
||||
|
||||
@@ -7,10 +7,13 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import DEFAULT_DECODE_GATHER_STRATEGY, decode_to_pixels_parallel
|
||||
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
|
||||
from fastvideo.profiler import nvtx_range
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MiniMaxH3PackedLayout,
|
||||
unpack_audio_tokens,
|
||||
@@ -21,6 +24,9 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.utils import is_pin_memory_available
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
|
||||
@@ -30,6 +36,23 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
|
||||
return layout
|
||||
|
||||
|
||||
def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]:
|
||||
"""Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages.
|
||||
|
||||
The executors consume rank 0's ForwardBatch and the training validation
|
||||
callback consumes each sequence-parallel group leader's, so the output
|
||||
rank is the SP group's first rank (identical to world rank 0 in the
|
||||
single-group e2e case). ``parallel`` is only true when every group rank
|
||||
will run the decode body — the collectives inside require uniform
|
||||
participation, so no rank-dependent branch may guard them.
|
||||
"""
|
||||
if not model_parallel_is_initialized():
|
||||
return None, True, False
|
||||
sp_group = get_sp_group()
|
||||
parallel = bool(want_parallel) and sp_group.world_size > 1
|
||||
return sp_group, sp_group.is_first_rank, parallel
|
||||
|
||||
|
||||
class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
"""Drop visual condition rows, unpatchify, and decode the target video."""
|
||||
|
||||
@@ -54,6 +77,16 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Decode H3 video latents into normalized CPU pixels."""
|
||||
placeholder = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
|
||||
sp_group, is_output_rank, parallel = _decode_participation(fastvideo_args, fastvideo_args.vae_parallel_decode)
|
||||
if not is_output_rank and not parallel:
|
||||
# Consumers read the output rank's ForwardBatch. Keep a
|
||||
# verifier-compatible placeholder on other ranks and avoid
|
||||
# duplicating the full VAE decode and CPU output buffer.
|
||||
batch.output = placeholder
|
||||
return batch
|
||||
|
||||
layout = _layout(batch)
|
||||
if batch.latents is None or batch.raw_latent_shape is None or len(batch.raw_latent_shape) != 5:
|
||||
raise ValueError("MiniMax-H3 video latents or raw geometry are missing at decode.")
|
||||
@@ -71,13 +104,33 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
try:
|
||||
latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32))
|
||||
if fastvideo_args.output_type == "latent":
|
||||
batch.output = latents.detach().float().cpu()
|
||||
# No collectives on this path, so uniform participation is
|
||||
# trivial: every rank returns here.
|
||||
batch.output = latents.detach().float().cpu() if is_output_rank else placeholder
|
||||
return batch
|
||||
|
||||
# The published decode recipe uses FP16 autocast over FP32 weights.
|
||||
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
|
||||
video = self.vae.decode(latents).sample
|
||||
batch.output = self.vae.denormalize_pixels(video.float()).clamp_(0, 1).cpu()
|
||||
output = None
|
||||
if is_output_rank:
|
||||
output = torch.empty(
|
||||
self.vae.decoded_pixel_shape(latents.shape),
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
|
||||
)
|
||||
# Attribute the streamed decoder computation while retaining
|
||||
# per-chunk device-to-host transfer and pinned-buffer reuse.
|
||||
with (
|
||||
nvtx_range("minimax_h3.vae"),
|
||||
torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"),
|
||||
):
|
||||
if parallel:
|
||||
strategy = fastvideo_args.vae_parallel_decode_strategy or DEFAULT_DECODE_GATHER_STRATEGY
|
||||
logger.info_once(f"MiniMax-H3 VAE decode: sequence-parallel chunks across "
|
||||
f"{sp_group.world_size} ranks ({strategy})")
|
||||
decode_to_pixels_parallel(self.vae, latents, output, sp_group, strategy=strategy)
|
||||
else:
|
||||
self.vae.decode_to_pixels(latents, output)
|
||||
batch.output = output if is_output_rank else placeholder
|
||||
return batch
|
||||
finally:
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
@@ -107,6 +160,15 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Decode H3 audio latents into a stereo CPU waveform."""
|
||||
# Audio decode is sub-second, so it always runs serially on the SP
|
||||
# group's first rank (the rank whose ForwardBatch consumers read).
|
||||
if model_parallel_is_initialized() and not get_sp_group().is_first_rank:
|
||||
batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32)
|
||||
batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate
|
||||
self._clear_runtime(batch)
|
||||
return batch
|
||||
|
||||
layout = _layout(batch)
|
||||
if batch.audio_latents is None:
|
||||
raise ValueError("MiniMax-H3 audio latents are missing at decode.")
|
||||
@@ -124,7 +186,10 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
|
||||
self._clear_runtime(batch)
|
||||
return batch
|
||||
|
||||
decoded = self.audio_vae.decode(latents).sample.float()
|
||||
# The range isolates waveform synthesis from packing and runtime
|
||||
# cleanup so the audio decoder has one stable timeline boundary.
|
||||
with nvtx_range("minimax_h3.audio_vae"):
|
||||
decoded = self.audio_vae.decode(latents).sample.float()
|
||||
if decoded.ndim != 3 or decoded.shape[0] != 2 or decoded.shape[1] != 1:
|
||||
raise ValueError("MiniMax-H3 audio VAE must decode stereo channels as two mono batch items; "
|
||||
f"got {tuple(decoded.shape)}.")
|
||||
|
||||
@@ -11,8 +11,8 @@ from fastvideo.attention.selector import component_attention_backend, get_attn_b
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.profiler import profiler_region
|
||||
from fastvideo.hooks.activation_trace import trace_step
|
||||
from fastvideo.profiler import nvtx_range, profiler_region
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_KEYFRAME_NOISE_AUG,
|
||||
MiniMaxH3PackedLayout,
|
||||
@@ -89,6 +89,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Denoise the packed H3 video and audio streams over one shared schedule."""
|
||||
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
|
||||
if not isinstance(layout, MiniMaxH3PackedLayout):
|
||||
raise ValueError("MiniMax-H3 packed layout is missing before denoising.")
|
||||
@@ -145,9 +146,15 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
vsa_exempt = vsa_mode == "exempt"
|
||||
vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ()))
|
||||
vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0))
|
||||
# Run-level tile geometry (256 default, 64 = native Triton path),
|
||||
# plumbed like the run-level sparsity; the builder validates the
|
||||
# value against VSA_H3_TILE_SHAPES.
|
||||
vsa_tile_size = int(fastvideo_args.VSA_tile_size)
|
||||
|
||||
try:
|
||||
with profiler_region("inference_denoising"):
|
||||
# The stage range groups the complete denoising loop while the
|
||||
# indexed model ranges retain timing detail for every H3 block.
|
||||
with profiler_region("inference_denoising"), nvtx_range("minimax_h3.dit"):
|
||||
for index, (video_timestep,
|
||||
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
|
||||
unique_timesteps, timestep_indices = row_timestep_plan[index]
|
||||
@@ -167,6 +174,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
device=device,
|
||||
exempt=vsa_exempt,
|
||||
dense_layers=vsa_dense_layers,
|
||||
tile_size=vsa_tile_size,
|
||||
)
|
||||
# Under torch.compile(mode="reduce-overhead") each denoising
|
||||
# step must be marked, or cudagraph trees flag cross-step
|
||||
|
||||
@@ -9,8 +9,10 @@ import numpy as np
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import encode_pixels_parallel
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_AUDIO_CHANNELS,
|
||||
MINIMAX_H3_KEYFRAME_ENCODE_SEED,
|
||||
@@ -36,6 +38,8 @@ from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
MINIMAX_H3_LAYOUT_KEY = "minimax_h3_layout"
|
||||
|
||||
|
||||
@@ -105,8 +109,20 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
|
||||
self,
|
||||
references: list[MiniMaxH3PreparedReference],
|
||||
device: torch.device,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> list[torch.Tensor]:
|
||||
patch_size = self.transformer.patch_size
|
||||
# Reference encode runs on every rank (all ranks hold identical
|
||||
# prepared references), so clip-parallel encode keeps participation
|
||||
# uniform by construction: each rank encodes a clip subset and the
|
||||
# all-gather leaves the identical full posterior everywhere.
|
||||
parallel_group = None
|
||||
if fastvideo_args.vae_parallel_encode and model_parallel_is_initialized():
|
||||
sp_group = get_sp_group()
|
||||
if sp_group.world_size > 1:
|
||||
parallel_group = sp_group
|
||||
logger.info_once(f"MiniMax-H3 reference VAE encode: sequence-parallel clips across "
|
||||
f"{sp_group.world_size} ranks")
|
||||
rows: list[torch.Tensor] = []
|
||||
for reference in references:
|
||||
if reference.media_type == "audio":
|
||||
@@ -119,9 +135,11 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
|
||||
if reference.frames is None:
|
||||
raise ValueError("MiniMax-H3 reference video frames are missing.")
|
||||
frames = reference.frames[:trim_reference_num_frames(reference.frames.shape[0])]
|
||||
pixels = torch.from_numpy(frames.copy()).permute(3, 0, 1, 2)[None]
|
||||
pixels = pixels.to(device=device, dtype=torch.float32).div_(255.0)
|
||||
posterior = self.vae.encode(self.vae.normalize_pixels(pixels)).latent_dist
|
||||
pixels = torch.from_numpy(np.ascontiguousarray(frames)).permute(3, 0, 1, 2)[None]
|
||||
if parallel_group is not None:
|
||||
posterior = encode_pixels_parallel(self.vae, pixels, parallel_group).latent_dist
|
||||
else:
|
||||
posterior = self.vae.encode_pixels(pixels).latent_dist
|
||||
latents = self.vae.normalize_latents(_sample_visual_posterior(posterior).to(
|
||||
torch.float16).float()).cpu()
|
||||
reference.num_latent_frames = int(latents.shape[2])
|
||||
@@ -202,7 +220,7 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
|
||||
vae_device = get_local_torch_device()
|
||||
self.vae.to(vae_device)
|
||||
try:
|
||||
video_rows = self._encode_visual_rows(references, vae_device)
|
||||
video_rows = self._encode_visual_rows(references, vae_device, fastvideo_args)
|
||||
finally:
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
@@ -48,7 +48,10 @@ class MpsPlatform(Platform):
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: AttentionBackendEnum | None, head_size: int,
|
||||
dtype: torch.dtype) -> str:
|
||||
# MPS supports SDPA (Scaled Dot-Product Attention) which is the most compatible
|
||||
if selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
raise NotImplementedError("VIDEO_SPARSE_ATTN is not supported on MPS. Unset "
|
||||
"FASTVIDEO_ATTENTION_BACKEND or set it to TORCH_SDPA.")
|
||||
# MPS supports SDPA (Scaled Dot-Product Attention) which is the most compatible.
|
||||
logger.info("Using Torch SDPA backend for MPS.")
|
||||
return "fastvideo.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
|
||||
@@ -38,6 +38,25 @@ logger = init_logger(__name__)
|
||||
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def nvtx_range(name: str):
|
||||
"""Emit one optional NVTX range for an external CUDA profiler.
|
||||
|
||||
``FASTVIDEO_NVTX_PROFILE=1`` enables the marker. The context manager stays
|
||||
a no-op without CUDA so call sites can remain shared with CPU tests.
|
||||
"""
|
||||
enabled = envs.FASTVIDEO_NVTX_PROFILE and torch.cuda.is_available()
|
||||
if not enabled:
|
||||
yield
|
||||
return
|
||||
|
||||
torch.cuda.nvtx.range_push(name)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
torch.cuda.nvtx.range_pop()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProfilerRegion:
|
||||
"""Metadata describing a profiler region."""
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GPU backward checks for the VSA-H3 backend.
|
||||
|
||||
The CuTe backend returns FA4's own output tensor, which FA4's autograd node
|
||||
saved for its backward. Composing the compression branch onto it in place
|
||||
therefore poisons the graph, and the failure only appears once the VSA-256
|
||||
CuTe path has a backward at all. These tests pin the composition.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAImpl, MiniMaxH3VSAMetadataBuilder)
|
||||
|
||||
_SPEC = dict(raw_latent_shape=(16, 16, 24), patch_size=(1, 2, 2), prefix_segments=(64, 32, 16))
|
||||
_HEADS = 2
|
||||
_DIM = 128
|
||||
|
||||
|
||||
def _build_meta(device, sparsity=0.5):
|
||||
return MiniMaxH3VSAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=_SPEC["raw_latent_shape"],
|
||||
patch_size=_SPEC["patch_size"],
|
||||
VSA_sparsity=sparsity,
|
||||
prefix_segments=_SPEC["prefix_segments"],
|
||||
device=device,
|
||||
)
|
||||
|
||||
|
||||
def _select_backend(monkeypatch, backend):
|
||||
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 _forward_backward(impl, meta, gate_compress, device):
|
||||
seq = meta.total_seq_length
|
||||
torch.manual_seed(0)
|
||||
q, k, v = (torch.randn(1, seq, _HEADS, _DIM, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
for _ in range(3))
|
||||
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
|
||||
|
||||
gate = None
|
||||
if gate_compress:
|
||||
gate = torch.randn(1, tq.shape[1], _HEADS, _DIM, device=device, dtype=torch.bfloat16) * 0.1
|
||||
|
||||
out = impl.forward(tq, tk, tv, gate, meta)
|
||||
out = impl.postprocess_output(out, meta)
|
||||
out.float().pow(2).sum().backward()
|
||||
return out, (q, k, v)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("backend", ["triton", "cute"])
|
||||
@pytest.mark.parametrize("gate_compress", [False, True])
|
||||
def test_h3_vsa_backward_runs(monkeypatch, backend: str, gate_compress: bool) -> None:
|
||||
"""Regression: with the CuTe backend and a non-zero gate this used to die
|
||||
with "one of the variables needed for gradient computation has been
|
||||
modified by an inplace operation ... output 0 of FlashAttnFuncBackward".
|
||||
"""
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required")
|
||||
_select_backend(monkeypatch, backend)
|
||||
|
||||
device = torch.device("cuda")
|
||||
meta = _build_meta(device)
|
||||
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
|
||||
|
||||
out, leaves = _forward_backward(impl, meta, gate_compress, device)
|
||||
|
||||
assert torch.isfinite(out).all().item()
|
||||
for name, leaf in zip(("q", "k", "v"), leaves):
|
||||
assert leaf.grad is not None, f"{name} received no gradient"
|
||||
assert torch.isfinite(leaf.grad).all().item(), f"{name}.grad has non-finite values"
|
||||
assert leaf.grad.abs().sum().item() > 0, f"{name}.grad is all zero"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("gate_compress", [False, True])
|
||||
def test_h3_vsa_backward_cute_matches_triton(monkeypatch, gate_compress: bool) -> None:
|
||||
"""CuTe and Triton take different routes to the same math; their gradients
|
||||
should agree to bf16 tolerance."""
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required")
|
||||
|
||||
device = torch.device("cuda")
|
||||
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
|
||||
|
||||
grads = {}
|
||||
for backend in ("triton", "cute"):
|
||||
with monkeypatch.context() as m:
|
||||
_select_backend(m, backend)
|
||||
meta = _build_meta(device)
|
||||
_, leaves = _forward_backward(impl, meta, gate_compress, device)
|
||||
grads[backend] = [leaf.grad.detach().float() for leaf in leaves]
|
||||
|
||||
for name, ref, got in zip(("dq", "dk", "dv"), grads["triton"], grads["cute"]):
|
||||
diff = (ref - got).abs()
|
||||
avg_abs = diff.mean().item()
|
||||
max_rel = (diff.max() / (ref.abs().mean() + 1e-6)).item()
|
||||
print(f"[h3-vsa gate={gate_compress}] {name}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
|
||||
assert avg_abs < 1e-2, f"{name}: avg_abs {avg_abs:.3e}"
|
||||
assert max_rel < 0.5, f"{name}: max_rel {max_rel:.3e}"
|
||||
@@ -5,20 +5,29 @@ reference. The same reference doubles as the GPU kernel parity oracle."""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import (_TILE_ELEMS, MiniMaxH3VSAImpl,
|
||||
MiniMaxH3VSAMetadataBuilder, _build_block_mask,
|
||||
_pool_tiles, token_tile_and_valid)
|
||||
_pool_tiles, _validate_h3_tile_geometry,
|
||||
token_tile_and_valid)
|
||||
|
||||
_720P = dict(raw_latent_shape=(30, 44, 80), patch_size=(1, 2, 2), prefix_segments=(512, 1760, 400))
|
||||
_TINY = dict(raw_latent_shape=(8, 8, 12), patch_size=(1, 2, 2), prefix_segments=(7, 5, 3))
|
||||
# (4,4,4) coverage: dit grid (9, 10, 13) is ragged in all three dims
|
||||
# (t: 4+4+1, h: 4+4+2, w: 4+4+4+1) and every prefix segment leaves a
|
||||
# partial tail tile at 64 (70 -> 64+6, 5 -> 5, 130 -> 64+64+2).
|
||||
_TINY64 = dict(raw_latent_shape=(9, 20, 26), patch_size=(1, 2, 2), prefix_segments=(70, 5, 130))
|
||||
# production-shape request: 768x1344, 124 frames -> latents (37, 48, 84),
|
||||
# patch (1,2,2) -> token grid (37, 24, 42); text 300 + audio 414 rows.
|
||||
_PROD = dict(raw_latent_shape=(37, 48, 84), patch_size=(1, 2, 2), prefix_segments=(300, 0, 414))
|
||||
|
||||
_CPU = torch.device("cpu")
|
||||
|
||||
|
||||
def _build(spec, sparsity=0.0, device=_CPU):
|
||||
def _build(spec, sparsity=0.0, device=_CPU, tile_size=_TILE_ELEMS):
|
||||
return MiniMaxH3VSAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=spec["raw_latent_shape"],
|
||||
@@ -26,6 +35,7 @@ def _build(spec, sparsity=0.0, device=_CPU):
|
||||
VSA_sparsity=sparsity,
|
||||
prefix_segments=spec["prefix_segments"],
|
||||
device=device,
|
||||
tile_size=tile_size,
|
||||
)
|
||||
|
||||
|
||||
@@ -36,7 +46,7 @@ def _impl():
|
||||
def reference_sparse_attention(query, key, value, mask, meta):
|
||||
"""Token-level oracle: SDPA over the padded tile buffer with the block
|
||||
mask expanded to tokens. query/key/value: tiled [B, S_pad, H, D]."""
|
||||
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes)
|
||||
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes, meta.tile_elems)
|
||||
out = torch.empty_like(query)
|
||||
for b in range(query.shape[0]):
|
||||
for h in range(query.shape[2]):
|
||||
@@ -133,9 +143,117 @@ def test_prefix_queries_stay_dense_at_high_sparsity():
|
||||
"video rows should actually be sparse at 75%"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 64-token (4,4,4) tile geometry
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_geometry_tile64_ragged_tails():
|
||||
"""Hand-computed (4,4,4) oracle on a grid ragged in all three dims."""
|
||||
meta = _build(_TINY64, tile_size=64)
|
||||
assert meta.tile_elems == 64
|
||||
t, h, w = 9, 10, 13 # raw latents (9, 20, 26) under patch (1, 2, 2)
|
||||
n_t, n_h, n_w = 3, 3, 4
|
||||
prefix_len = sum(_TINY64["prefix_segments"])
|
||||
seq = prefix_len + t * h * w
|
||||
assert meta.total_seq_length == seq
|
||||
assert meta.num_prefix_tiles == 2 + 1 + 3
|
||||
assert meta.num_video_tiles == n_t * n_h * n_w
|
||||
assert int(meta.variable_block_sizes.sum()) == seq
|
||||
assert int(meta.variable_block_sizes.max()) <= 64
|
||||
assert meta.variable_block_sizes[:meta.num_prefix_tiles].tolist() == [64, 6, 5, 64, 64, 2]
|
||||
|
||||
# per-tile valid sizes: product of the per-dim clamped tails
|
||||
expected = torch.tensor([
|
||||
min(4, t - 4 * tt) * min(4, h - 4 * hh) * min(4, w - 4 * ww) for tt in range(n_t) for hh in range(n_h)
|
||||
for ww in range(n_w)
|
||||
],
|
||||
dtype=torch.long)
|
||||
assert torch.equal(meta.variable_block_sizes[meta.num_prefix_tiles:], expected)
|
||||
assert int(expected.min()) == 1 * 2 * 1 # the (t,h,w) ragged corner
|
||||
|
||||
# every packed video row lands in the 3D tile its (t,h,w) coordinate says
|
||||
idx = meta.untile_combined_index
|
||||
row = torch.arange(t * h * w)
|
||||
row_t, row_h, row_w = row // (h * w), (row // w) % h, row % w
|
||||
expected_tile = meta.num_prefix_tiles + ((row_t // 4) * n_h + row_h // 4) * n_w + row_w // 4
|
||||
assert torch.equal(idx[prefix_len:] // 64, expected_tile)
|
||||
# and in a non-pad slot of that tile
|
||||
assert bool((idx % 64 < meta.variable_block_sizes[idx // 64]).all())
|
||||
|
||||
# untile(tile(x)) == x on the 64-wide padded buffer
|
||||
x = torch.randn(1, seq, 2, 4)
|
||||
buf = _impl().tile(x, meta)
|
||||
assert buf.shape[1] == meta.variable_block_sizes.numel() * 64
|
||||
assert torch.equal(buf[:, idx], x)
|
||||
|
||||
|
||||
def test_geometry_tile64_production_shape():
|
||||
"""Production latents (37, 48, 84): ragged t and w tails at (4,4,4)."""
|
||||
meta64 = _build(_PROD, tile_size=64)
|
||||
assert meta64.num_prefix_tiles == 5 + 7 # 300 -> 4x64+44, 414 -> 6x64+30
|
||||
assert meta64.num_video_tiles == 10 * 6 * 11 # (37, 24, 42) / (4, 4, 4)
|
||||
assert meta64.total_seq_length == 300 + 414 + 37 * 24 * 42
|
||||
assert int(meta64.variable_block_sizes.sum()) == meta64.total_seq_length
|
||||
sizes_vid = meta64.variable_block_sizes[meta64.num_prefix_tiles:]
|
||||
assert int(sizes_vid.max()) == 64 and int(sizes_vid.min()) == 1 * 4 * 2 # (t, w) ragged corner
|
||||
|
||||
# same packed sequence under the default 256 geometry, fewer tiles
|
||||
meta256 = _build(_PROD)
|
||||
assert meta256.tile_elems == _TILE_ELEMS
|
||||
assert meta256.num_prefix_tiles == 2 + 2
|
||||
assert meta256.num_video_tiles == 10 * 3 * 6
|
||||
assert meta256.total_seq_length == meta64.total_seq_length
|
||||
|
||||
x = torch.randn(1, meta64.total_seq_length, 2, 4)
|
||||
buf = _impl().tile(x, meta64)
|
||||
assert torch.equal(buf[:, meta64.untile_combined_index], x)
|
||||
|
||||
|
||||
def test_sparsity_zero_matches_dense_sdpa_tile64():
|
||||
torch.manual_seed(2)
|
||||
meta = _build(_TINY64, tile_size=64)
|
||||
seq = meta.total_seq_length
|
||||
q, k, v = (torch.randn(1, seq, 2, 8) for _ in range(3))
|
||||
impl = _impl()
|
||||
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
|
||||
|
||||
scores = torch.matmul(_pool_tiles(tq, meta.variable_block_sizes, meta.tile_elems),
|
||||
_pool_tiles(tk, meta.variable_block_sizes, meta.tile_elems).transpose(-2, -1))
|
||||
mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, 0.0, exempt=True)
|
||||
sparse_out = impl.postprocess_output(reference_sparse_attention(tq, tk, tv, mask, meta), meta)
|
||||
|
||||
dense_out = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2)
|
||||
assert torch.allclose(sparse_out, dense_out, atol=1e-5), (sparse_out - dense_out).abs().max()
|
||||
|
||||
|
||||
def test_geometry_guard_enforces_tile64_bound():
|
||||
"""A 65-token tile passes the 256 bound but must fail the 64 one."""
|
||||
meta = _build(_TINY64, tile_size=64)
|
||||
prefix = tuple(s for s in _TINY64["prefix_segments"] if s > 0)
|
||||
dit_shape = (9, 10, 13)
|
||||
sizes = meta.variable_block_sizes.clone()
|
||||
sizes[0] = 65
|
||||
with pytest.raises(ValueError, match="tile sizes out of bounds"):
|
||||
_validate_h3_tile_geometry(prefix, dit_shape, sizes, meta.untile_combined_index, 64)
|
||||
# the untampered tile-64 geometry passes its own bound
|
||||
_validate_h3_tile_geometry(prefix, dit_shape, meta.variable_block_sizes, meta.untile_combined_index, 64)
|
||||
|
||||
|
||||
def test_builder_rejects_unknown_tile_size():
|
||||
for bad in (0, 128, 512):
|
||||
with pytest.raises(ValueError, match="tile_size"):
|
||||
_build(_TINY, tile_size=bad)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_geometry_720p()
|
||||
test_mask_policy()
|
||||
test_sparsity_zero_matches_dense_sdpa()
|
||||
test_prefix_queries_stay_dense_at_high_sparsity()
|
||||
test_geometry_tile64_ragged_tails()
|
||||
test_geometry_tile64_production_shape()
|
||||
test_sparsity_zero_matches_dense_sdpa_tile64()
|
||||
test_geometry_guard_enforces_tile64_bound()
|
||||
test_builder_rejects_unknown_tile_size()
|
||||
print("all VSA-H3 CPU checks passed")
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU checks for the VSA-H3 tile-64 sm_100a route selection.
|
||||
|
||||
The opt-in third kernel route (``FASTVIDEO_VSA_SM100A=1``) must (a) stay off by
|
||||
default, (b) engage only when the extension is present, the device qualifies,
|
||||
and the forward carries no grad, and (c) fall back to the Triton-64 entry with
|
||||
one warning when the env is set but a precondition fails. All device/extension
|
||||
probes are monkeypatched; no GPU or kernel install needed.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import fastvideo.attention.backends.video_sparse_attn_h3 as vsa_h3
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import (VSA_SM100A_ENV, MiniMaxH3VSAImpl,
|
||||
MiniMaxH3VSAMetadataBuilder, _sm100a_unavailable_reason)
|
||||
|
||||
# Small tile-64 geometry: 2 prefix segments + a (4,4,8)-token video grid.
|
||||
_SPEC = dict(raw_latent_shape=(4, 8, 16), patch_size=(1, 2, 2), prefix_segments=(70, 30))
|
||||
_HEADS, _DIM = 2, 128
|
||||
|
||||
|
||||
def _build_meta():
|
||||
return MiniMaxH3VSAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=_SPEC["raw_latent_shape"],
|
||||
patch_size=_SPEC["patch_size"],
|
||||
VSA_sparsity=0.0,
|
||||
prefix_segments=_SPEC["prefix_segments"],
|
||||
device=torch.device("cpu"),
|
||||
tile_size=64,
|
||||
)
|
||||
|
||||
|
||||
def _tiled_qkv(meta, requires_grad=False):
|
||||
# bf16 like the real tiled buffers, so forward()'s dtype-cast warning
|
||||
# stays out of the warning assertions below.
|
||||
s_pad = meta.variable_block_sizes.numel() * 64
|
||||
return tuple(
|
||||
torch.randn(1, s_pad, _HEADS, _DIM, dtype=torch.bfloat16, requires_grad=requires_grad) for _ in range(3))
|
||||
|
||||
|
||||
class _FakeSm100a:
|
||||
"""Stands in for fastvideo_kernel.block_sparse_attn_sm100a."""
|
||||
|
||||
def __init__(self, supported=True):
|
||||
self.supported = supported
|
||||
self.calls = []
|
||||
|
||||
def is_supported(self, q, variable_block_sizes):
|
||||
return self.supported
|
||||
|
||||
def block_sparse_attn_sm100a(self, q, k, v, q2k_idx, q2k_num, variable_block_sizes, need_lse=True):
|
||||
self.calls.append(dict(q=q, q2k_idx=q2k_idx, q2k_num=q2k_num, vbs=variable_block_sizes,
|
||||
need_lse=need_lse))
|
||||
return q.clone(), None
|
||||
|
||||
|
||||
def _fake_map_to_index(block_map):
|
||||
"""Pure-torch stand-in for the Triton map_to_index (same contract)."""
|
||||
b, h, t, n = block_map.shape
|
||||
idx = torch.full((b, h, t, n), -1, dtype=torch.int32)
|
||||
num = block_map.sum(dim=-1, dtype=torch.int32)
|
||||
for bi in range(b):
|
||||
for hi in range(h):
|
||||
for ti in range(t):
|
||||
cols = torch.nonzero(block_map[bi, hi, ti], as_tuple=False).flatten()
|
||||
idx[bi, hi, ti, :cols.numel()] = cols.to(torch.int32)
|
||||
return idx, num
|
||||
|
||||
|
||||
class _FakeTriton:
|
||||
def __init__(self):
|
||||
self.calls = 0
|
||||
|
||||
def __call__(self, q, k, v, mask, variable_block_sizes):
|
||||
self.calls += 1
|
||||
return q.clone(), None
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def routed(monkeypatch):
|
||||
"""Backend with both kernel entries faked; returns (fakes, run)."""
|
||||
fake_sm = _FakeSm100a()
|
||||
fake_triton = _FakeTriton()
|
||||
monkeypatch.setattr(vsa_h3, "_sm100a", fake_sm)
|
||||
monkeypatch.setattr(vsa_h3, "block_sparse_attn_64_bhsd", fake_triton)
|
||||
monkeypatch.setattr(vsa_h3, "map_to_index", _fake_map_to_index)
|
||||
meta = _build_meta()
|
||||
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
|
||||
|
||||
def run(requires_grad=False):
|
||||
q, k, v = _tiled_qkv(meta, requires_grad=requires_grad)
|
||||
return impl.forward(q, k, v, None, meta)
|
||||
|
||||
return fake_sm, fake_triton, run, meta
|
||||
|
||||
|
||||
def test_reason_covers_every_precondition():
|
||||
q = torch.randn(1, _HEADS, 128, _DIM)
|
||||
vbs = torch.full((2, ), 64, dtype=torch.long)
|
||||
assert "not installed" in _sm100a_unavailable_reason(None, q, vbs, grad_mode=False)
|
||||
ok = _FakeSm100a(supported=True)
|
||||
assert "forward-only" in _sm100a_unavailable_reason(ok, q, vbs, grad_mode=True)
|
||||
bad = _FakeSm100a(supported=False)
|
||||
assert "is_supported" in _sm100a_unavailable_reason(bad, q, vbs, grad_mode=False)
|
||||
assert _sm100a_unavailable_reason(ok, q, vbs, grad_mode=False) is None
|
||||
|
||||
|
||||
def test_default_off_routes_triton(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
monkeypatch.delenv(VSA_SM100A_ENV, raising=False)
|
||||
run()
|
||||
assert fake_triton.calls == 1
|
||||
assert fake_sm.calls == []
|
||||
|
||||
|
||||
def test_env_on_routes_sm100a_with_index_metadata(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, meta = routed
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
out = run()
|
||||
assert fake_triton.calls == 0
|
||||
assert len(fake_sm.calls) == 1
|
||||
call = fake_sm.calls[0]
|
||||
n_tiles = meta.variable_block_sizes.numel()
|
||||
# sparsity 0 -> all-True mask -> every row's count is n_tiles
|
||||
assert call["q2k_num"].dtype == torch.int32 and (call["q2k_num"] == n_tiles).all()
|
||||
assert call["q2k_idx"].shape[-1] == n_tiles and call["q2k_idx"].dtype == torch.int32
|
||||
assert call["vbs"].dtype == torch.int32
|
||||
assert call["need_lse"] is False
|
||||
# BHSD kernel result comes back in the backend's BSHD layout
|
||||
assert out.shape == (1, n_tiles * 64, _HEADS, _DIM)
|
||||
|
||||
|
||||
def test_env_on_grad_inputs_fall_back_to_triton(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
run(requires_grad=True)
|
||||
assert fake_triton.calls == 1
|
||||
assert fake_sm.calls == []
|
||||
# ...but the same process still routes no-grad forwards to sm_100a
|
||||
run(requires_grad=False)
|
||||
assert len(fake_sm.calls) == 1
|
||||
|
||||
|
||||
def test_env_on_unsupported_warns_once_and_falls_back(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
fake_sm.supported = False
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
warnings = []
|
||||
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
|
||||
run()
|
||||
run()
|
||||
assert fake_triton.calls == 2
|
||||
assert fake_sm.calls == []
|
||||
assert len(warnings) == 2 # warning_once dedups by message; both carry the same one line
|
||||
assert warnings[0] == warnings[1]
|
||||
assert VSA_SM100A_ENV in warnings[0] and "is_supported" in warnings[0]
|
||||
|
||||
|
||||
def test_env_on_missing_module_warns_and_falls_back(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
monkeypatch.setattr(vsa_h3, "_sm100a", None)
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
warnings = []
|
||||
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
|
||||
run()
|
||||
assert fake_triton.calls == 1
|
||||
assert warnings and "not installed" in warnings[0]
|
||||
|
||||
|
||||
def test_env_on_no_grad_context_detaches_route_from_leaf_flags(routed, monkeypatch):
|
||||
"""A requires_grad leaf under torch.no_grad() is still a no-grad forward."""
|
||||
fake_sm, fake_triton, run, meta = routed
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
|
||||
q, k, v = _tiled_qkv(meta, requires_grad=True)
|
||||
with torch.no_grad():
|
||||
impl.forward(q, k, v, None, meta)
|
||||
assert len(fake_sm.calls) == 1
|
||||
assert fake_triton.calls == 0
|
||||
@@ -15,6 +15,12 @@ import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.profiler import nvtx_range
|
||||
|
||||
# Five-window child: ops before any region, inside a region, between regions,
|
||||
# inside a second (short-named) region, after the last region. Exits without
|
||||
@@ -105,3 +111,73 @@ def test_noop_without_profiler_dir(tmp_path):
|
||||
proc = subprocess.run([sys.executable, "-c", child], env=env,
|
||||
capture_output=True, text=True, timeout=300)
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
|
||||
|
||||
def test_nvtx_range_disabled_is_noop(monkeypatch):
|
||||
"""Keep CUDA NVTX untouched when external profiling is disabled."""
|
||||
range_push = Mock()
|
||||
range_pop = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
|
||||
|
||||
with nvtx_range("disabled"):
|
||||
body_executed = True
|
||||
|
||||
assert body_executed is True
|
||||
range_push.assert_not_called()
|
||||
range_pop.assert_not_called()
|
||||
|
||||
|
||||
def test_nvtx_range_without_cuda_is_noop(monkeypatch):
|
||||
"""Keep NVTX untouched when profiling is enabled on a CPU-only process."""
|
||||
range_push = Mock()
|
||||
range_pop = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
|
||||
|
||||
with nvtx_range("cpu-only"):
|
||||
body_executed = True
|
||||
|
||||
assert body_executed is True
|
||||
range_push.assert_not_called()
|
||||
range_pop.assert_not_called()
|
||||
|
||||
|
||||
def test_nvtx_range_enabled_orders_push_body_pop(monkeypatch):
|
||||
"""Place the profiled body between one matching NVTX push and pop."""
|
||||
events = []
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
|
||||
|
||||
with nvtx_range("minimax_h3.test"):
|
||||
events.append(("body", None))
|
||||
|
||||
assert events == [
|
||||
("push", "minimax_h3.test"),
|
||||
("body", None),
|
||||
("pop", None),
|
||||
]
|
||||
|
||||
|
||||
def test_nvtx_range_body_exception_pops_and_propagates(monkeypatch):
|
||||
"""Balance the NVTX stack while preserving a body exception."""
|
||||
events = []
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
|
||||
|
||||
with pytest.raises(RuntimeError, match="profile body failed"):
|
||||
with nvtx_range("minimax_h3.failure"):
|
||||
raise RuntimeError("profile body failed")
|
||||
|
||||
assert events == [
|
||||
("push", "minimax_h3.failure"),
|
||||
("pop", None),
|
||||
]
|
||||
|
||||
@@ -0,0 +1,281 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29514")
|
||||
|
||||
import fastvideo.models.encoders.minimax_h3_checkpoint_fp8 as h3_fp8
|
||||
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
|
||||
from fastvideo.layers.linear import ColumnParallelLinear, UnquantizedLinearMethod
|
||||
from fastvideo.layers.vocab_parallel_embedding import UnquantizedEmbeddingMethod, VocabParallelEmbedding
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import (
|
||||
MiniMaxH3SerializedFP8Config,
|
||||
MiniMaxH3SerializedFP8LinearMethod,
|
||||
)
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
from fastvideo.models.loader.text_encoder_quantization import (
|
||||
_configure_text_encoder_quantization,
|
||||
_process_quantized_text_encoder_weights,
|
||||
_read_text_encoder_checkpoint_quantization_config,
|
||||
)
|
||||
|
||||
|
||||
def _checkpoint_quantization_config(**overrides) -> dict:
|
||||
config = {
|
||||
"quant_method": "fp8",
|
||||
"activation_scheme": "dynamic",
|
||||
"fmt": "e4m3",
|
||||
"weight_block_size": [128, 128],
|
||||
"modules_to_not_convert": ["model.visual", "lm_head"],
|
||||
}
|
||||
config.update(overrides)
|
||||
return config
|
||||
|
||||
|
||||
def test_h3_accepts_only_the_serialized_blockwise_checkpoint_contract() -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
assert config.weight_block_size == (128, 128)
|
||||
assert config.get_supported_act_dtypes() == [torch.bfloat16]
|
||||
|
||||
with pytest.raises(ValueError, match=r"weight_block_size=\[128, 128\]"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(weight_block_size=[1, 128]))
|
||||
with pytest.raises(ValueError, match="dynamic activation"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(activation_scheme="static"))
|
||||
with pytest.raises(ValueError, match="vision stack"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(modules_to_not_convert=["lm_head"]))
|
||||
with pytest.raises(ValueError, match="partially quantized language"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(modules_to_not_convert=["model.visual", "language_model.layers.3"]))
|
||||
|
||||
|
||||
def test_serialized_fp8_allocates_checkpoint_weight_and_scale_without_requantization(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
layer = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=256,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
|
||||
)
|
||||
|
||||
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
|
||||
assert layer.weight.dtype == torch.float8_e4m3fn
|
||||
assert layer.weight.shape == (256, 128)
|
||||
assert layer.weight_scale_inv.dtype == torch.float32
|
||||
assert layer.weight_scale_inv.shape == (2, 1)
|
||||
|
||||
layer.weight.data.zero_()
|
||||
layer.weight_scale_inv.data.fill_(0.25)
|
||||
weight_pointer = layer.weight.data_ptr()
|
||||
scale_pointer = layer.weight_scale_inv.data_ptr()
|
||||
layer.quant_method.process_weights_after_loading(layer)
|
||||
|
||||
assert layer.weight.data_ptr() == weight_pointer
|
||||
assert layer.weight_scale_inv.data_ptr() == scale_pointer
|
||||
assert not hasattr(layer, "_fp8_weight")
|
||||
|
||||
|
||||
def test_serialized_fp8_quantizes_only_language_linears(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
visual_linear = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.visual.blocks.0.attn.proj",
|
||||
)
|
||||
embedding = VocabParallelEmbedding(
|
||||
num_embeddings=128,
|
||||
embedding_dim=128,
|
||||
org_num_embeddings=128,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.embed_tokens",
|
||||
)
|
||||
|
||||
assert isinstance(visual_linear.quant_method, UnquantizedLinearMethod)
|
||||
assert visual_linear.weight.dtype == torch.get_default_dtype()
|
||||
assert isinstance(embedding.quant_method, UnquantizedEmbeddingMethod)
|
||||
assert embedding.weight.dtype == torch.get_default_dtype()
|
||||
|
||||
|
||||
def test_serialized_fp8_cpu_execution_fails_closed(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
layer = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.layers.0.mlp.up_proj",
|
||||
)
|
||||
layer.weight.data.zero_()
|
||||
layer.weight_scale_inv.data.fill_(1.0)
|
||||
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
|
||||
layer.quant_method.process_weights_after_loading(layer)
|
||||
|
||||
with pytest.raises(RuntimeError, match="requires CUDA"):
|
||||
layer(torch.zeros(2, 128, dtype=torch.bfloat16))
|
||||
|
||||
|
||||
def test_runtime_preflight_reports_capability_and_missing_dependencies(monkeypatch) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 0))
|
||||
with pytest.raises(RuntimeError, match="sm100 or newer"):
|
||||
config.validate_runtime(torch.device("cuda"))
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (10, 0))
|
||||
|
||||
def missing_quantizer() -> None:
|
||||
raise RuntimeError("SGLang-compatible Triton quantizer is missing")
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", missing_quantizer)
|
||||
with pytest.raises(RuntimeError, match="Triton quantizer is missing"):
|
||||
config.validate_runtime(torch.device("cuda"))
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", lambda: None)
|
||||
|
||||
def missing_flashinfer():
|
||||
raise RuntimeError("FlashInfer groupwise GEMM is missing")
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", missing_flashinfer)
|
||||
with pytest.raises(RuntimeError, match="FlashInfer groupwise GEMM is missing"):
|
||||
config.validate_runtime(torch.device("cuda"))
|
||||
|
||||
|
||||
def test_loader_detects_and_capability_gates_checkpoint_metadata(tmp_path) -> None:
|
||||
checkpoint_config = _checkpoint_quantization_config()
|
||||
(tmp_path / "config.json").write_text(
|
||||
json.dumps({"quantization_config": checkpoint_config}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
assert _read_text_encoder_checkpoint_quantization_config(str(tmp_path)) == checkpoint_config
|
||||
model_config = MiniMaxH3Qwen3VLConfig()
|
||||
quant_config = _configure_text_encoder_quantization(
|
||||
model_config,
|
||||
MiniMaxH3Qwen3VLConditioner,
|
||||
str(tmp_path),
|
||||
)
|
||||
assert isinstance(quant_config, MiniMaxH3SerializedFP8Config)
|
||||
assert model_config.quant_config is quant_config
|
||||
|
||||
unsupported_config = MiniMaxH3Qwen3VLConfig()
|
||||
with pytest.raises(ValueError, match="does not support serialized 'fp8'"):
|
||||
_configure_text_encoder_quantization(
|
||||
unsupported_config,
|
||||
TextEncoder,
|
||||
str(tmp_path),
|
||||
)
|
||||
|
||||
|
||||
def test_loader_leaves_bf16_checkpoint_path_unchanged(tmp_path) -> None:
|
||||
(tmp_path / "config.json").write_text(json.dumps({"architectures": ["Qwen3VLModel"]}), encoding="utf-8")
|
||||
model_config = MiniMaxH3Qwen3VLConfig()
|
||||
|
||||
quant_config = _configure_text_encoder_quantization(
|
||||
model_config,
|
||||
MiniMaxH3Qwen3VLConditioner,
|
||||
str(tmp_path),
|
||||
)
|
||||
|
||||
assert quant_config is None
|
||||
assert model_config.quant_config is None
|
||||
|
||||
|
||||
def test_post_load_processing_visits_only_serialized_fp8_linears(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
quantized = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
|
||||
)
|
||||
plain = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
prefix="plain",
|
||||
)
|
||||
quantized.weight.data.zero_()
|
||||
quantized.weight_scale_inv.data.fill_(1.0)
|
||||
model = torch.nn.ModuleList([quantized, plain])
|
||||
|
||||
assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 1
|
||||
assert quantized.weight.device.type == "cpu"
|
||||
assert plain.weight.device.type == "cpu"
|
||||
|
||||
|
||||
def test_flashinfer_groupwise_path_pins_output_dtype_and_trtllm_scale_layout(monkeypatch) -> None:
|
||||
input_tensor = torch.zeros(2, 256, dtype=torch.bfloat16)
|
||||
weight = torch.zeros(128, 256, dtype=torch.float8_e4m3fn)
|
||||
weight_scale = torch.ones(1, 2, dtype=torch.float32)
|
||||
quantized_input = torch.zeros_like(input_tensor, dtype=torch.float8_e4m3fn)
|
||||
input_scale = torch.empty(2, 2, dtype=torch.float32).t()
|
||||
input_scale.fill_(1.0)
|
||||
receipt: dict[str, object] = {}
|
||||
|
||||
def fake_quantize(
|
||||
value: torch.Tensor,
|
||||
group_size: int,
|
||||
*,
|
||||
column_major_scales: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
assert value.data_ptr() == input_tensor.data_ptr()
|
||||
assert value.shape == input_tensor.shape
|
||||
assert group_size == 128
|
||||
assert column_major_scales is True
|
||||
return quantized_input, input_scale
|
||||
|
||||
def fake_gemm(
|
||||
activation: torch.Tensor,
|
||||
checkpoint_weight: torch.Tensor,
|
||||
activation_scale: torch.Tensor,
|
||||
checkpoint_scale: torch.Tensor,
|
||||
*,
|
||||
out_dtype: torch.dtype,
|
||||
backend: str,
|
||||
) -> torch.Tensor:
|
||||
receipt.update(
|
||||
activation=activation,
|
||||
checkpoint_weight=checkpoint_weight,
|
||||
activation_scale=activation_scale,
|
||||
checkpoint_scale=checkpoint_scale,
|
||||
out_dtype=out_dtype,
|
||||
backend=backend,
|
||||
)
|
||||
return torch.zeros(activation.shape[0], checkpoint_weight.shape[0], dtype=out_dtype)
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_backend", lambda device: "trtllm")
|
||||
monkeypatch.setattr(h3_fp8, "_sglang_per_token_group_quant_fp8", fake_quantize)
|
||||
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", lambda: fake_gemm)
|
||||
|
||||
previous_default_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(torch.float32)
|
||||
try:
|
||||
output = h3_fp8._flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
|
||||
input_tensor,
|
||||
weight,
|
||||
(128, 128),
|
||||
weight_scale,
|
||||
)
|
||||
assert torch.get_default_dtype() == torch.float32
|
||||
finally:
|
||||
torch.set_default_dtype(previous_default_dtype)
|
||||
|
||||
assert output.dtype == torch.bfloat16
|
||||
assert receipt["out_dtype"] == torch.bfloat16
|
||||
assert receipt["backend"] == "trtllm"
|
||||
assert receipt["activation"] is quantized_input
|
||||
assert receipt["checkpoint_weight"] is weight
|
||||
assert receipt["checkpoint_scale"] is weight_scale
|
||||
assert receipt["activation_scale"] is input_scale
|
||||
@@ -1,27 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""The Qwen3-VL stack is built only as far as MiniMax H3 reads.
|
||||
|
||||
H3 conditions on one intermediate hidden state. The layers above it were built,
|
||||
weight-loaded and then discarded, which is 13.7 GB in bf16 and the difference
|
||||
between fitting and not fitting on a 121 GB unified-memory device.
|
||||
|
||||
The dangerous part is not the truncation, it is getting the tuple index wrong.
|
||||
`hidden_states` records each layer's *input*, so entry N is the output of layer
|
||||
N-1, and the final entry comes from the norm that sits above the whole stack. A
|
||||
truncated stack that still applies that norm puts a normalised tensor where the
|
||||
raw one belongs: the length check in the conditioning stage still passes, and
|
||||
conditioning silently changes. These tests pin the index, the content, and the
|
||||
constant the two sides agree on.
|
||||
"""
|
||||
"""MiniMax-H3 Qwen3-VL layer truncation and slim-forward tests."""
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# Matches the other encoder tests: the module registry these build against wants
|
||||
# a process group, and a single-rank one needs a rendezvous address.
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
|
||||
@@ -34,21 +20,15 @@ from fastvideo.pipelines.basic.minimax_h3.packing import MINIMAX_H3_TEXT_ENCODER
|
||||
|
||||
|
||||
def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
|
||||
"""A stack small enough to run on CPU but shaped like the real one.
|
||||
|
||||
Everything goes through the constructor so ``__post_init__`` validates the
|
||||
small shape the same way it validates the real one.
|
||||
"""
|
||||
kwargs: dict = dict(
|
||||
vocab_size=64,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=8,
|
||||
output_hidden_state_index=5,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
head_dim=8,
|
||||
# __post_init__ reads the sections out of rope_scaling, and they must
|
||||
# cover exactly half of each head.
|
||||
rope_scaling={
|
||||
"mrope_interleaved": True,
|
||||
"mrope_section": [2, 1, 1],
|
||||
@@ -61,25 +41,27 @@ def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
|
||||
|
||||
|
||||
def _small_config(**overrides) -> MiniMaxH3Qwen3VLConfig:
|
||||
"""The outer config, which is what the modules take.
|
||||
|
||||
``ModelConfig.__getattr__`` forwards the architecture fields, so the modules
|
||||
read ``prefix`` off this object and everything else off ``arch_config``.
|
||||
"""
|
||||
config = MiniMaxH3Qwen3VLConfig()
|
||||
config.arch_config = _small_arch(**overrides)
|
||||
return config
|
||||
|
||||
|
||||
def test_default_matches_the_index_the_pipeline_reads() -> None:
|
||||
"""The two sides cannot import each other, so pin them here instead.
|
||||
config = MiniMaxH3Qwen3VLArchConfig()
|
||||
|
||||
`fastvideo/models/` must not import from `fastvideo/pipelines/`, so the tap
|
||||
is written down twice. If they drift, conditioning reads a hidden state that
|
||||
was never built and the run dies with an index error at generation time,
|
||||
after a full model load.
|
||||
"""
|
||||
assert MiniMaxH3Qwen3VLArchConfig().num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
assert config.output_hidden_state_index == MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
assert config.num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
|
||||
|
||||
def test_rejects_build_depth_that_cannot_reach_the_output() -> None:
|
||||
for override in (0, 4):
|
||||
with pytest.raises(ValueError, match="num_hidden_layers_override"):
|
||||
_small_arch(num_hidden_layers_override=override)
|
||||
|
||||
|
||||
def test_rejects_output_index_above_the_checkpoint_depth() -> None:
|
||||
with pytest.raises(ValueError, match="output_hidden_state_index"):
|
||||
_small_arch(output_hidden_state_index=9, num_hidden_layers_override=None)
|
||||
|
||||
|
||||
def test_builds_only_up_to_the_override(distributed_setup) -> None:
|
||||
@@ -87,7 +69,6 @@ def test_builds_only_up_to_the_override(distributed_setup) -> None:
|
||||
|
||||
assert model.num_layers == 5
|
||||
assert len(model.layers) == 5
|
||||
# The norm sits above the tap, so a truncated stack must not keep it.
|
||||
assert model.norm is None
|
||||
|
||||
|
||||
@@ -98,104 +79,94 @@ def test_override_none_keeps_the_full_stack(distributed_setup) -> None:
|
||||
assert model.norm is not None
|
||||
|
||||
|
||||
def test_nominal_and_built_depths_remain_distinct(distributed_setup) -> None:
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
|
||||
|
||||
assert conditioner.num_hidden_layers == 8
|
||||
assert conditioner.num_built_hidden_layers == 5
|
||||
|
||||
|
||||
def test_override_above_the_stack_does_not_over_build(distributed_setup) -> None:
|
||||
# num_hidden_layers comes from the checkpoint's config.json via
|
||||
# update_model_arch, so a smaller variant must clamp rather than ask for
|
||||
# layers that do not exist.
|
||||
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=99))
|
||||
|
||||
assert model.num_layers == 8
|
||||
assert model.norm is not None
|
||||
|
||||
|
||||
def test_override_equal_to_the_stack_keeps_the_norm(distributed_setup) -> None:
|
||||
"""The exact boundary of the clamp: a stack cut at its own depth is full.
|
||||
|
||||
A checkpoint with exactly ``override`` layers taps its final layer, whose
|
||||
tuple entry sits after the norm in the full model, so the norm must stay
|
||||
and nothing may be filtered from the checkpoint.
|
||||
"""
|
||||
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=8))
|
||||
|
||||
assert model.num_layers == 8
|
||||
assert model.norm is not None
|
||||
|
||||
|
||||
def test_non_positive_override_is_rejected() -> None:
|
||||
"""A non-positive override would build no decoder layers at all.
|
||||
|
||||
Worse, a negative one makes ``num_layers`` disagree with the built stack
|
||||
and the surplus-key filter would then drop every layer key, so the
|
||||
conditioner would load "successfully" with no transformer. Reject it at
|
||||
config construction, and again when update_model_arch re-validates.
|
||||
"""
|
||||
for override in (0, -1):
|
||||
with pytest.raises(ValueError, match="num_hidden_layers_override"):
|
||||
_small_arch(num_hidden_layers_override=override)
|
||||
|
||||
config = _small_config()
|
||||
with pytest.raises(ValueError, match="num_hidden_layers_override"):
|
||||
config.update_model_arch({"num_hidden_layers_override": 0})
|
||||
|
||||
|
||||
def test_tapped_hidden_state_is_unchanged_by_truncation(distributed_setup) -> None:
|
||||
"""The whole point: entry `tap` must be bit-identical either way."""
|
||||
"""The slim model returns the raw output at the selected layer."""
|
||||
tap = 5
|
||||
full = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
|
||||
cut = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=tap))
|
||||
|
||||
# These modules allocate uninitialised storage and expect a checkpoint, so
|
||||
# give them finite weights before running anything through them.
|
||||
torch.manual_seed(0)
|
||||
for parameter in full.parameters():
|
||||
parameter.data.normal_(std=0.02)
|
||||
# Then make the shared prefix identical, which is the only part the tapped
|
||||
# hidden state depends on.
|
||||
for (_, a), (_, b) in zip(full.layers[:tap].named_parameters(),
|
||||
cut.layers[:tap].named_parameters(),
|
||||
strict=True):
|
||||
b.data.copy_(a.data)
|
||||
torch.manual_seed(1)
|
||||
inputs_embeds = torch.randn(1, 6, 16)
|
||||
# mRoPE indexes three axes (t, h, w); text tokens share the same position on
|
||||
# all three.
|
||||
position_ids = torch.arange(6).view(1, 1, 6).expand(3, 1, 6)
|
||||
with torch.no_grad():
|
||||
full_out = full(inputs_embeds, position_ids, None, True, None, None)
|
||||
cut_out = cut(inputs_embeds, position_ids, None, True, None, None)
|
||||
expected = inputs_embeds
|
||||
position_embeddings = full.rotary_emb(inputs_embeds, position_ids)
|
||||
for layer in full.layers[:tap]:
|
||||
expected = layer(expected, position_embeddings, None)
|
||||
full_out = full(inputs_embeds, position_ids, None, None, None)
|
||||
cut_out = cut(inputs_embeds, position_ids, None, None, None)
|
||||
|
||||
assert torch.equal(full_out.hidden_states[tap], cut_out.hidden_states[tap])
|
||||
# And the truncated model must not offer states it never computed.
|
||||
assert len(cut_out.hidden_states) == tap + 1
|
||||
# The whole shared prefix must match, not just the tap: this is the same
|
||||
# comparison the production-loader parity gate runs against the official
|
||||
# model, and it is what catches a truncated stack that still applied the
|
||||
# final norm to its last entry.
|
||||
for index, (cut_state, full_state) in enumerate(zip(cut_out.hidden_states, full_out.hidden_states,
|
||||
strict=False)):
|
||||
assert torch.equal(cut_state, full_state), f"hidden state {index} changed under truncation"
|
||||
assert torch.equal(expected, full_out)
|
||||
assert torch.equal(expected, cut_out)
|
||||
|
||||
|
||||
def test_conditioning_stage_adapts_slim_sequence_output() -> None:
|
||||
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage
|
||||
|
||||
class FakeConditioner:
|
||||
|
||||
dtype = torch.float32
|
||||
|
||||
def __call__(self, input_ids: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||
assert input_ids.ndim == 1
|
||||
assert not kwargs
|
||||
return torch.ones(input_ids.shape[0], 4)
|
||||
|
||||
stage = MiniMaxH3ConditioningStage(conditioner=FakeConditioner(), tokenizer=None, processor=None, ref2va=False)
|
||||
embeddings, tags = stage._encode_tokens([1, 2, 3], [0, 0, 0], torch.device("cpu"))
|
||||
|
||||
assert embeddings.shape == (1, 3, 4)
|
||||
assert tags.shape == (3, )
|
||||
|
||||
|
||||
def test_conditioner_exposes_only_the_slim_forward_contract() -> None:
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
assert tuple(inspect.signature(MiniMaxH3Qwen3VLConditioner.forward).parameters) == (
|
||||
"self",
|
||||
"input_ids",
|
||||
"pixel_values",
|
||||
"image_grid_thw",
|
||||
"pixel_values_videos",
|
||||
"video_grid_thw",
|
||||
)
|
||||
|
||||
|
||||
def test_truncated_model_drops_the_surplus_checkpoint_keys(distributed_setup) -> None:
|
||||
"""The unexpected-key check is strict on purpose, so the surplus keys have
|
||||
to be filtered rather than the check relaxed."""
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
|
||||
|
||||
assert conditioner._is_above_the_tap("language_model.layers.5.mlp.gate_proj.weight")
|
||||
assert conditioner._is_above_the_tap("language_model.layers.7.self_attn.q_proj.weight")
|
||||
assert conditioner._is_above_the_tap("language_model.norm.weight")
|
||||
# Kept: layers we built, the embeddings, and the vision tower.
|
||||
assert not conditioner._is_above_the_tap("language_model.layers.4.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_above_the_tap("language_model.embed_tokens.weight")
|
||||
assert not conditioner._is_above_the_tap("visual.blocks.0.attn.qkv.weight")
|
||||
# The filter only drops indexes the full stack would have built. A key at
|
||||
# or above the checkpoint's own num_hidden_layers is corrupt, and it must
|
||||
# keep raising as unexpected exactly as it does without truncation.
|
||||
assert not conditioner._is_above_the_tap("language_model.layers.8.mlp.gate_proj.weight")
|
||||
with pytest.raises(ValueError, match="Unexpected"):
|
||||
conditioner.load_weights([("model.language_model.layers.8.mlp.gate_proj.weight", torch.zeros(1))])
|
||||
assert conditioner._is_omitted_checkpoint_key("language_model.layers.5.mlp.gate_proj.weight")
|
||||
assert conditioner._is_omitted_checkpoint_key("language_model.layers.7.self_attn.q_proj.weight")
|
||||
assert conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.4.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.embed_tokens.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("visual.blocks.0.attn.qkv.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.8.mlp.gate_proj.weight")
|
||||
|
||||
|
||||
def test_full_stack_filters_nothing(distributed_setup) -> None:
|
||||
@@ -203,5 +174,14 @@ def test_full_stack_filters_nothing(distributed_setup) -> None:
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=None))
|
||||
|
||||
assert not conditioner._is_above_the_tap("language_model.layers.7.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_above_the_tap("language_model.norm.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.7.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
|
||||
|
||||
|
||||
def test_corrupt_layer_above_checkpoint_depth_remains_unexpected(distributed_setup) -> None:
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
|
||||
|
||||
with pytest.raises(ValueError, match="Unexpected"):
|
||||
conditioner.load_weights([("language_model.layers.8.mlp.gate_proj.weight", torch.empty(1))])
|
||||
|
||||
@@ -2,8 +2,10 @@ import os
|
||||
from types import SimpleNamespace
|
||||
import warnings
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
import fastvideo.entrypoints.video_generator as video_generator_module
|
||||
from fastvideo.api import (
|
||||
@@ -271,6 +273,70 @@ def test_generate_single_video_return_frames_still_materializes_output(tmp_path)
|
||||
assert result["video_path"] is None
|
||||
|
||||
|
||||
def test_generate_single_video_frames_match_legacy_cpu_loop(tmp_path):
|
||||
"""The on-device quantize path (#1362) must reproduce the legacy
|
||||
per-frame CPU loop (make_grid -> permute -> *255 -> uint8) bit-exactly
|
||||
for in-range fp32 pixels: same uint8 dtype, same HWC grid layout with
|
||||
nrow=6 (batch>1), odd frame count. CPU-only: on CUDA the float->uint8
|
||||
cast may differ by <=1 LSB, but on CPU both orderings run identical
|
||||
fp32 ops, so exact equality is required."""
|
||||
torch.manual_seed(0)
|
||||
output = torch.rand((2, 3, 3, 16, 16), dtype=torch.float32)
|
||||
output_batch = _single_video_output_batch(output)
|
||||
fastvideo_args = _single_video_args()
|
||||
generator = _single_video_generator(output_batch, fastvideo_args)
|
||||
sampling_param = _small_sampling_param(save_video=False, return_frames=True)
|
||||
sampling_param.num_frames = 3
|
||||
sampling_param.num_videos_per_prompt = 2
|
||||
|
||||
result = generator._generate_single_video(
|
||||
prompt="grid parity",
|
||||
sampling_param=sampling_param,
|
||||
fastvideo_args=fastvideo_args,
|
||||
output_path=str(tmp_path / "unused.mp4"),
|
||||
)
|
||||
|
||||
legacy_frames = []
|
||||
for x in rearrange(output, "b c t h w -> t b c h w"):
|
||||
grid = video_generator_module.torchvision.utils.make_grid(x, nrow=6)
|
||||
grid = grid.permute(1, 2, 0).squeeze(-1)
|
||||
legacy_frames.append((grid * 255).to(torch.uint8).contiguous().cpu().numpy())
|
||||
|
||||
torch.testing.assert_close(result["samples"], output)
|
||||
assert len(result["frames"]) == 3
|
||||
for got, want in zip(result["frames"], legacy_frames, strict=True):
|
||||
assert got.dtype == np.uint8
|
||||
assert got.shape == want.shape
|
||||
np.testing.assert_array_equal(got, want)
|
||||
|
||||
|
||||
def test_generate_single_video_frames_clamp_out_of_range_pixels(tmp_path):
|
||||
"""VAE output slightly outside [0, 1] must saturate at 0/255 in the
|
||||
uint8 frames. The pre-#1362 unclamped cast wrapped mod 256 (e.g.
|
||||
1.5 -> 126). CPU-only."""
|
||||
output = torch.full((1, 3, 2, 16, 16), 1.5, dtype=torch.float32)
|
||||
output[:, :, 1] = -0.5
|
||||
output_batch = _single_video_output_batch(output)
|
||||
fastvideo_args = _single_video_args()
|
||||
generator = _single_video_generator(output_batch, fastvideo_args)
|
||||
|
||||
result = generator._generate_single_video(
|
||||
prompt="clamp",
|
||||
sampling_param=_small_sampling_param(save_video=False, return_frames=True),
|
||||
fastvideo_args=fastvideo_args,
|
||||
output_path=str(tmp_path / "unused.mp4"),
|
||||
)
|
||||
|
||||
frames = result["frames"]
|
||||
assert len(frames) == 2
|
||||
# make_grid passes a single image through without grid padding, so
|
||||
# every pixel comes from the (clamped) output tensor.
|
||||
assert frames[0].dtype == np.uint8
|
||||
assert frames[0].shape == (16, 16, 3)
|
||||
assert (frames[0] == 255).all()
|
||||
assert (frames[1] == 0).all()
|
||||
|
||||
|
||||
def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_path):
|
||||
output = torch.ones((1, 3, 2, 16, 16), dtype=torch.float32) * 0.5
|
||||
output_batch = _single_video_output_batch(output)
|
||||
@@ -305,6 +371,37 @@ def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_p
|
||||
}
|
||||
|
||||
|
||||
def test_generate_single_video_save_only_reports_refined_output_size(monkeypatch, tmp_path):
|
||||
"""`GenerationResult.size` must describe the decoded media even when the
|
||||
fp32 `samples` mirror is skipped (`return_frames=False`, the CLI save
|
||||
flow). Refiner pipelines can change the final pixel geometry, so the size
|
||||
has to come from `output_batch.output`, not the base request. CPU-only."""
|
||||
# Refiner-style output: request asks for 2 frames of 16x16, pipeline
|
||||
# produces 5 frames of 32x48.
|
||||
output = torch.full((1, 3, 5, 32, 48), 0.5, dtype=torch.float32)
|
||||
output_batch = _single_video_output_batch(output)
|
||||
fastvideo_args = _single_video_args()
|
||||
generator = _single_video_generator(output_batch, fastvideo_args)
|
||||
saved = {}
|
||||
|
||||
def fake_mimsave(path, frames, *, fps, format):
|
||||
saved["frame_count"] = len(frames)
|
||||
|
||||
monkeypatch.setattr(video_generator_module.imageio, "mimsave", fake_mimsave)
|
||||
|
||||
result = generator._generate_single_video(
|
||||
prompt="refined save",
|
||||
sampling_param=_small_sampling_param(save_video=True, return_frames=False),
|
||||
fastvideo_args=fastvideo_args,
|
||||
output_path=str(tmp_path / "refined.mp4"),
|
||||
)
|
||||
|
||||
assert result["samples"] is None
|
||||
assert result["frames"] is None
|
||||
assert result["size"] == (32, 48, 5)
|
||||
assert saved["frame_count"] == 5
|
||||
|
||||
|
||||
def test_generate_single_video_audio_only_metadata_returns_audio_without_frames(tmp_path):
|
||||
audio = torch.zeros((16, ), dtype=torch.float32)
|
||||
output_batch = _single_video_output_batch(
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user