Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cbc4195c5a | ||
|
|
982ddef6da | ||
|
|
0d9d4ad132 | ||
|
|
7445aeabfb |
@@ -0,0 +1,51 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-1gpu-gb10",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp1",
|
||||
"benchmark_version": 3,
|
||||
"description": "Wan2.1 T2V 1.3B single-GPU inference performance on NVIDIA DGX Spark (GB10). Single-GPU variant of wan-t2v-1.3b (same workload_id for dashboard comparability). Gated to the GB10 via run_config.gpu_types so it does not run on the shared H100/L40S lanes.",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"model_short_name": "Wan2.1-T2V-1.3B"
|
||||
},
|
||||
"init_kwargs": {
|
||||
"num_gpus": 1,
|
||||
"flow_shift": 7.0,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"vae_sp": false,
|
||||
"vae_tiling": true,
|
||||
"text_encoder_precisions": ["fp32"]
|
||||
},
|
||||
"generation_kwargs": {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 3,
|
||||
"embedded_cfg_scale": 6,
|
||||
"seed": 1024,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
},
|
||||
"test_prompts": [
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
],
|
||||
"run_config": {
|
||||
"num_warmup_runs": 2,
|
||||
"num_measurement_runs": 5,
|
||||
"required_gpus": 1,
|
||||
"gpu_types": ["GB10"]
|
||||
},
|
||||
"thresholds": {
|
||||
"GB10": {
|
||||
"max_generation_time_s": 55.0,
|
||||
"max_peak_memory_mb": 12000.0
|
||||
},
|
||||
"default": {
|
||||
"max_generation_time_s": 120.0,
|
||||
"max_peak_memory_mb": 40000.0
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,149 +0,0 @@
|
||||
name: macOS MLX Smoke
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
paths:
|
||||
- ".github/workflows/ci-macos-mlx.yml"
|
||||
- "fastvideo/mlx_runtime/**"
|
||||
- "fastvideo/tests/mlx/**"
|
||||
- "fastvideo/tests/platforms/test_mps_vsa_error.py"
|
||||
- "fastvideo/platforms/mps.py"
|
||||
- "fastvideo/platforms/__init__.py"
|
||||
- "fastvideo/__init__.py"
|
||||
- "examples/inference/basic/mlx_*.py"
|
||||
- "fastvideo/benchmarks/mlx_*.py"
|
||||
- "pyproject.toml"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: macos-mlx-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
mlx-smoke:
|
||||
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
|
||||
runs-on: macos-15
|
||||
timeout-minutes: 25
|
||||
env:
|
||||
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
|
||||
TOKENIZERS_PARALLELISM: "false"
|
||||
MASTER_ADDR: localhost
|
||||
MASTER_PORT: "29513"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install lightweight MLX smoke dependencies
|
||||
run: |
|
||||
uv pip install --system \
|
||||
--index-url https://download.pytorch.org/whl/cpu \
|
||||
torch==2.11.0 torchvision torchaudio
|
||||
uv pip install --system \
|
||||
pytest numpy scipy pillow imageio einops cloudpickle filelock \
|
||||
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru mlx \
|
||||
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
|
||||
|
||||
- name: Show Apple runtime
|
||||
run: |
|
||||
python - <<'PY'
|
||||
import platform
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
print("machine:", platform.machine())
|
||||
print("processor:", platform.processor())
|
||||
print("mlx default device:", mx.default_device())
|
||||
memory_size = mx.metal.device_info().get("memory_size") if mx.metal.is_available() else "metal unavailable"
|
||||
print("mlx memory_size:", memory_size)
|
||||
print("torch:", torch.__version__)
|
||||
print("torch mps available:", torch.backends.mps.is_available())
|
||||
PY
|
||||
|
||||
- name: Run MLX smoke tests
|
||||
run: |
|
||||
python -m pytest \
|
||||
fastvideo/tests/mlx/test_dmd_sampling.py \
|
||||
fastvideo/tests/mlx/test_memory_limits.py \
|
||||
fastvideo/tests/mlx/test_quant_capability.py \
|
||||
fastvideo/tests/mlx/test_mlx_dit_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_compile_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_checkpoint.py \
|
||||
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
|
||||
fastvideo/tests/mlx/test_taehv_decode.py \
|
||||
fastvideo/tests/mlx/test_frame_upsample.py \
|
||||
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
|
||||
fastvideo/tests/mlx/test_mlx_refine.py \
|
||||
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
|
||||
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
|
||||
fastvideo/tests/mlx/test_wan22_sample.py \
|
||||
fastvideo/tests/mlx/test_windowed_attention.py \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
|
||||
fastvideo/tests/platforms/test_mps_vsa_error.py \
|
||||
-q
|
||||
|
||||
# Same tests on MLX's CPU backend. Hosted macOS runners are scarce and
|
||||
# slower to schedule; this Linux job gives fast PR signal on the identical
|
||||
# graph (the parity tests were designed to be backend-agnostic), while the
|
||||
# macOS job above stays the source of truth for Metal behavior.
|
||||
mlx-smoke-linux-cpu:
|
||||
if: github.event_name == 'workflow_dispatch' || github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
env:
|
||||
FASTVIDEO_ATTENTION_BACKEND: TORCH_SDPA
|
||||
TOKENIZERS_PARALLELISM: "false"
|
||||
MASTER_ADDR: localhost
|
||||
MASTER_PORT: "29513"
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
cache: pip
|
||||
|
||||
- uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install lightweight MLX smoke dependencies (CPU backend)
|
||||
run: |
|
||||
uv pip install --system \
|
||||
--index-url https://download.pytorch.org/whl/cpu \
|
||||
torch==2.11.0 torchvision torchaudio
|
||||
uv pip install --system \
|
||||
pytest numpy scipy pillow imageio einops cloudpickle filelock \
|
||||
PyYAML diffusers huggingface_hub remote-pdb safetensors loguru "mlx[cpu]" \
|
||||
"ftfy>=6.3.1" "opencv-python>=4.10.0.84" psutil "transformers>=5.0.0"
|
||||
|
||||
- name: Run MLX smoke tests (CPU backend)
|
||||
run: |
|
||||
python -m pytest \
|
||||
fastvideo/tests/mlx/test_dmd_sampling.py \
|
||||
fastvideo/tests/mlx/test_memory_limits.py \
|
||||
fastvideo/tests/mlx/test_quant_capability.py \
|
||||
fastvideo/tests/mlx/test_mlx_dit_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_compile_parity.py \
|
||||
fastvideo/tests/mlx/test_mlx_checkpoint.py \
|
||||
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
|
||||
fastvideo/tests/mlx/test_taehv_decode.py \
|
||||
fastvideo/tests/mlx/test_frame_upsample.py \
|
||||
fastvideo/tests/mlx/test_mlx_fast_spatial.py \
|
||||
fastvideo/tests/mlx/test_mlx_refine.py \
|
||||
fastvideo/tests/mlx/test_mlx_prompt_to_video_decode.py \
|
||||
fastvideo/tests/mlx/test_mlx_wan22_prompt_cache_fingerprint.py \
|
||||
fastvideo/tests/mlx/test_wan22_sample.py \
|
||||
fastvideo/tests/mlx/test_windowed_attention.py \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_download_unavailable_has_specific_error \
|
||||
fastvideo/tests/mlx/test_mlx_rife_interpolation.py::test_rife_backend_regression_is_not_skip_eligible \
|
||||
fastvideo/tests/platforms/test_mps_vsa_error.py \
|
||||
-q
|
||||
@@ -6,7 +6,6 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
@@ -17,7 +16,6 @@ on:
|
||||
paths:
|
||||
- 'docs/**'
|
||||
- 'examples/**'
|
||||
- 'scripts/inference/**'
|
||||
- 'mkdocs.yml'
|
||||
- 'requirements-mkdocs.in'
|
||||
- 'requirements-mkdocs.txt'
|
||||
|
||||
@@ -23,7 +23,6 @@ Miniconda3-latest-Linux-x86_64.sh
|
||||
*validation/
|
||||
data/
|
||||
outputs/
|
||||
outputs_audio/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
|
||||
@@ -9,7 +9,6 @@
|
||||
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
|
||||
|
||||
## NEWS
|
||||
- `2026/08/19`: FastVideo now supports MLX on Apple Silicon with [FastMetal-QAD](https://huggingface.co/collections/FastVideo/fastmetal), a family of 1.3B, 5B, and 14B models optimized for Mac—follow the [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/) and read the [Blog](https://haoailab.com/blogs/fastmetal/).
|
||||
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
|
||||
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
|
||||
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
|
||||
@@ -63,11 +62,6 @@ UV_TORCH_BACKEND=cu126 uv pip install fastvideo
|
||||
Use `UV_TORCH_BACKEND=cu130` on CUDA 13. Apple silicon users should follow the
|
||||
[MPS installation guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
|
||||
|
||||
> **On an Apple Silicon Mac?** FastVideo runs FastWan text-to-video natively
|
||||
> through an MLX runtime — a 5-second 480p clip generated locally, no cloud,
|
||||
> no discrete GPU. Install with `uv pip install -e '.[mlx]'` and follow the
|
||||
> [Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/) for more detailed installation instructions.
|
||||
|
||||
> **On an NVIDIA DGX Spark (GB10 / ARM64 + CUDA 13)?** There's no prebuilt ARM wheel for the FastVideo CUDA kernel, so it's an editable from-source install (`UV_TORCH_BACKEND=cu130 uv pip install -e .`, which compiles that kernel for you) rather than `UV_TORCH_BACKEND=cu130 uv pip install fastvideo`. A compatible prebuilt ARM64 FlashAttention wheel is available separately. Follow the [DGX Spark install guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/spark/).
|
||||
|
||||
+1
-1
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
|
||||
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
|
||||
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
|
||||
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
|
||||
ARG FA4_CUTE_REF=14c377950125c70b7a9dabf9c561fca53715ac7d
|
||||
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
|
||||
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
|
||||
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
{
|
||||
"recipes": [
|
||||
{
|
||||
"id": "fastwan21-t2v",
|
||||
"task": "Text to video",
|
||||
"label": "FastWan2.1 1.3B (distilled + VSA)",
|
||||
"model": "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
"source": "scripts/inference/inference_wan_VSA_DMD_1_3B.yaml",
|
||||
"command": "FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml"
|
||||
},
|
||||
{
|
||||
"id": "wan22-t2v",
|
||||
"task": "Text to video",
|
||||
"label": "Wan2.2 A14B",
|
||||
"model": "Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_wan2_2.py",
|
||||
"command": "python examples/inference/basic/basic_wan2_2.py"
|
||||
},
|
||||
{
|
||||
"id": "wan21-i2v",
|
||||
"task": "Image to video",
|
||||
"label": "Wan2.1 14B 480P",
|
||||
"model": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
"source": "scripts/inference/inference_wan_i2v.yaml",
|
||||
"command": "fastvideo generate --config scripts/inference/inference_wan_i2v.yaml"
|
||||
},
|
||||
{
|
||||
"id": "turbowan22-i2v",
|
||||
"task": "Image to video",
|
||||
"label": "TurboWan2.2 A14B",
|
||||
"model": "loayrashid/TurboWan2.2-I2V-A14B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_turbodiffusion_i2v.py",
|
||||
"command": "python examples/inference/basic/basic_turbodiffusion_i2v.py"
|
||||
},
|
||||
{
|
||||
"id": "wan22-ti2v",
|
||||
"task": "Text or image to video",
|
||||
"label": "Wan2.2 TI2V 5B",
|
||||
"model": "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
"source": "examples/inference/basic/basic_wan2_2_ti2v.py",
|
||||
"command": "python examples/inference/basic/basic_wan2_2_ti2v.py"
|
||||
},
|
||||
{
|
||||
"id": "matrix-game-2",
|
||||
"task": "Interactive world",
|
||||
"label": "Matrix Game 2.0",
|
||||
"model": "FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"source": "examples/inference/basic/basic_matrixgame2.py",
|
||||
"command": "python examples/inference/basic/basic_matrixgame2.py"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,60 +0,0 @@
|
||||
(() => {
|
||||
let recipesPromise;
|
||||
|
||||
const loadRecipes = (url) => {
|
||||
recipesPromise ||= fetch(url).then((response) => {
|
||||
if (!response.ok) throw new Error(`HTTP ${response.status}`);
|
||||
return response.json();
|
||||
});
|
||||
return recipesPromise;
|
||||
};
|
||||
|
||||
const init = () => {
|
||||
document.querySelectorAll("[data-cookbook]").forEach(async (root) => {
|
||||
if (root.dataset.initialized) return;
|
||||
root.dataset.initialized = "true";
|
||||
|
||||
const select = root.querySelector("[data-cookbook-recipe]");
|
||||
const model = root.querySelector("[data-cookbook-model]");
|
||||
const source = root.querySelector("[data-cookbook-source]");
|
||||
const command = root.querySelector("[data-cookbook-command]");
|
||||
const status = root.querySelector("[data-cookbook-status]");
|
||||
|
||||
try {
|
||||
const { recipes } = await loadRecipes(root.dataset.recipes);
|
||||
const byId = new Map(recipes.map((recipe) => [recipe.id, recipe]));
|
||||
const groups = new Map();
|
||||
|
||||
select.replaceChildren();
|
||||
recipes.forEach((recipe) => {
|
||||
if (!groups.has(recipe.task)) {
|
||||
const group = document.createElement("optgroup");
|
||||
group.label = recipe.task;
|
||||
groups.set(recipe.task, group);
|
||||
select.append(group);
|
||||
}
|
||||
groups.get(recipe.task).append(new Option(recipe.label, recipe.id));
|
||||
});
|
||||
|
||||
const render = () => {
|
||||
const recipe = byId.get(select.value);
|
||||
model.textContent = recipe.model;
|
||||
source.textContent = recipe.source;
|
||||
source.href = `https://github.com/hao-ai-lab/FastVideo/blob/main/${recipe.source}`;
|
||||
command.textContent = recipe.command;
|
||||
status.textContent = `${recipe.label} selected.`;
|
||||
};
|
||||
|
||||
select.addEventListener("change", render);
|
||||
select.disabled = false;
|
||||
render();
|
||||
} catch (error) {
|
||||
status.textContent = "Recipes could not be loaded. Use the examples link below.";
|
||||
console.error("Failed to load FastVideo cookbook recipes", error);
|
||||
}
|
||||
});
|
||||
};
|
||||
|
||||
if (window.document$) window.document$.subscribe(init);
|
||||
else document.addEventListener("DOMContentLoaded", init);
|
||||
})();
|
||||
@@ -42,46 +42,6 @@ img {
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
.cookbook-picker {
|
||||
padding: 1rem;
|
||||
border: 0.05rem solid var(--md-default-fg-color--lightest);
|
||||
border-radius: 0.2rem;
|
||||
}
|
||||
|
||||
.cookbook-picker select {
|
||||
width: 100%;
|
||||
padding: 0.6rem;
|
||||
color: var(--md-default-fg-color);
|
||||
background: var(--md-default-bg-color);
|
||||
border: 0.05rem solid var(--md-default-fg-color--lighter);
|
||||
border-radius: 0.2rem;
|
||||
}
|
||||
|
||||
.cookbook-picker dl {
|
||||
display: grid;
|
||||
grid-template-columns: max-content 1fr;
|
||||
gap: 0.25rem 1rem;
|
||||
}
|
||||
|
||||
.cookbook-picker dt {
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.cookbook-picker dd {
|
||||
margin: 0;
|
||||
min-width: 0;
|
||||
overflow-wrap: anywhere;
|
||||
}
|
||||
|
||||
.cookbook-picker__status {
|
||||
position: absolute;
|
||||
width: 1px;
|
||||
height: 1px;
|
||||
overflow: hidden;
|
||||
clip: rect(0, 0, 0, 0);
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.md-typeset .copy-page-button.md-button {
|
||||
float: right;
|
||||
margin: 0 0 1rem 1rem;
|
||||
|
||||
@@ -58,6 +58,54 @@ records, see `performance_dashboard/README.md`. The dashboard provides a
|
||||
FastAPI API plus a React UI and can be exposed with `ngrok` after building the
|
||||
frontend.
|
||||
|
||||
## DGX Spark (GB10) local benchmarking
|
||||
|
||||
The NVIDIA DGX Spark (GB10) is not available on Modal, so its coverage is
|
||||
**local/manual** rather than automated CI. The GB10 benchmark
|
||||
`wan-t2v-1.3b-1gpu-gb10` is gated to the GB10 via `run_config.gpu_types`
|
||||
(matched as substrings of the CUDA device name), so the shared H100/L40S
|
||||
performance lanes discover it and skip it, while a GB10 owner runs it locally.
|
||||
|
||||
Run just the GB10 benchmark on a DGX Spark:
|
||||
|
||||
```bash
|
||||
pytest 'fastvideo/tests/performance/test_inference_performance.py::test_inference_performance[wan-t2v-1.3b-1gpu-gb10]' -vs
|
||||
```
|
||||
|
||||
To check run-to-run stability (latency, peak memory, throughput), run it a few
|
||||
times from a clean results directory, then normalize:
|
||||
|
||||
```bash
|
||||
rm -f fastvideo/tests/performance/results/perf_*.json
|
||||
for i in 1 2 3 4 5; do
|
||||
pytest 'fastvideo/tests/performance/test_inference_performance.py::test_inference_performance[wan-t2v-1.3b-1gpu-gb10]' -vs
|
||||
done
|
||||
PERF_REPORTS_DIR=/tmp/fastvideo_perf_reports \
|
||||
python fastvideo/tests/performance/compare_baseline.py
|
||||
```
|
||||
|
||||
`compare_baseline.py` reports `CALIBRATION_NEEDED` until a baseline exists for
|
||||
the GB10 identity, and writes one `normalized_perf_*.json` per run. Reference
|
||||
figures on a GB10 (torch 2.12.0+cu130, transformers 5.14.0): generation ~39.3 s,
|
||||
peak ~8.4 GB, throughput ~1.15 fps, stable to ~0.4% across five runs.
|
||||
|
||||
### Seeding the GB10 baseline (follow-up)
|
||||
|
||||
Seeding a baseline-eligible record for the GB10 identity is intentionally **not**
|
||||
done from a local run: `seed_baseline.py` accepts only `scheduled_main`
|
||||
full-suite source artifacts, so ordinary local/manual uploads stay
|
||||
`baseline_eligible=false` (dashboard-visible, but they do not move the rolling
|
||||
baseline). Establishing the GB10 baseline requires either:
|
||||
|
||||
* a scheduled-main performance run on a GB10 CI runner once one is available, or
|
||||
* a carefully scoped, maintainer-approved manual calibration path that preserves
|
||||
the existing exact-identity, batch-consistency, provenance, and
|
||||
explicit-approval safeguards — it must **not** make arbitrary local uploads
|
||||
baseline-eligible.
|
||||
|
||||
The reviewed five-run GB10 artifacts are the stability evidence for that first
|
||||
baseline. Tracked in #1632.
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
|
||||
@@ -1,44 +0,0 @@
|
||||
# Inference Cookbook
|
||||
|
||||
Choose a complete recipe maintained in the FastVideo repository. Each command
|
||||
runs its checked-in source directly, so coupled model, GPU, offload, and
|
||||
attention settings do not drift into unsupported combinations.
|
||||
|
||||
The commands expect a local clone:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git
|
||||
cd FastVideo
|
||||
```
|
||||
|
||||
<div class="cookbook-picker" data-cookbook data-recipes="../assets/cookbook-recipes.json">
|
||||
<label for="cookbook-recipe"><strong>Recipe</strong></label>
|
||||
<select id="cookbook-recipe" data-cookbook-recipe disabled>
|
||||
<option>Loading recipes…</option>
|
||||
</select>
|
||||
<dl>
|
||||
<dt>Model</dt>
|
||||
<dd data-cookbook-model>Loading…</dd>
|
||||
<dt>Source</dt>
|
||||
<dd><a data-cookbook-source href="../inference/examples/basic/">Browse maintained examples</a></dd>
|
||||
</dl>
|
||||
<pre><code class="language-bash" data-cookbook-command>Loading…</code></pre>
|
||||
<p class="cookbook-picker__status" role="status" aria-live="polite" data-cookbook-status></p>
|
||||
<noscript>
|
||||
JavaScript is needed for the recipe picker. Browse the
|
||||
<a href="../inference/examples/examples_inference_index/">inference examples</a>
|
||||
instead.
|
||||
</noscript>
|
||||
</div>
|
||||
|
||||
## Customize a recipe
|
||||
|
||||
Start from the checked-in source, then change only the settings your model
|
||||
supports:
|
||||
|
||||
- [Configuration](../inference/configuration.md) covers the Python and CLI
|
||||
config surfaces.
|
||||
- [Optimizations](../inference/optimizations.md) covers attention backends,
|
||||
compilation, and memory tradeoffs.
|
||||
- [Support matrix](../inference/support_matrix.md) lists supported models and
|
||||
optimizations.
|
||||
@@ -1,128 +0,0 @@
|
||||
# Fast mode (RIFE) — Apple Silicon
|
||||
|
||||
`--fast` makes local generation ~2.7× faster by **generating fewer frames and
|
||||
interpolating the rest** with an Apple-Silicon-native RIFE model, instead of
|
||||
denoising every frame. Video-diffusion denoise is dominated by self-attention,
|
||||
which is O(tokens²); halving the frames cuts the token count ~2× and the denoise
|
||||
compute ~3.7×, so the wall-clock drops far more than 2×. RIFE (which estimates
|
||||
its own optical flow — no motion vectors needed) fills the dropped frames back
|
||||
in for ~1.4 s, and a light unsharp pass counters its softening.
|
||||
|
||||
Measured on the 1.3B INT8 QAD model (fox, 480×832×81, M4): generate 41 + RIFE→81
|
||||
runs in ~35 s of denoise vs ~90 s full, at reconstruction MS-SSIM **0.97**.
|
||||
Reproduce with `python -m fastvideo.benchmarks.eval_metalfx_rife --mode int8`.
|
||||
|
||||
> **Note:** Apple's *MetalFX* frame interpolation is **not** usable here — it
|
||||
> requires game-engine motion vectors + depth, which diffusion output lacks. We
|
||||
> use the video-native **`rife-mlx`** model instead (Metal-backed, torch-free).
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
|
||||
```
|
||||
|
||||
## Use
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/mlx_wan_prompt_to_video.py \
|
||||
--mlx-checkpoint <FastWan2.1-T2V-1.3B-INT8-QAD> \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
|
||||
--num-frames 81 --fast \
|
||||
--output-path video_samples/fox_fast.mp4
|
||||
```
|
||||
|
||||
`--num-frames` stays the *target* length; fast mode generates the smallest
|
||||
VAE-aligned keyframe count that RIFE can interpolate to that target.
|
||||
|
||||
| Flag | Default | Meaning |
|
||||
|---|---|---|
|
||||
| `--fast` / `--no-fast` | off | enable fast mode |
|
||||
| `--fast-factor` | 2 | generate 1/factor of the frames (2 = half) |
|
||||
| `--fast-sharpen` | 0.6 | light unsharp strength to counter RIFE softness (0 disables) |
|
||||
|
||||
Fast mode composes with everything else (`--mlx-quantization int8`,
|
||||
`--mlx-compile`, TAEHV vs `--decode-backend wan-vae`). Keep `--fast-factor` at 2
|
||||
for quality — larger temporal gaps are where RIFE starts inventing motion.
|
||||
|
||||
## Spatial fast mode (`--fast-spatial`)
|
||||
|
||||
The spatial twin of `--fast`: instead of dropping frames, drop pixels. Denoise
|
||||
*and decode* at `height/width // fast-spatial-scale`, then resample the decoded
|
||||
frames up to the requested size. Self-attention is O(tokens²), so halving each
|
||||
spatial axis cuts the token count 4× and the denoise time far more than that —
|
||||
measured on the 1.3B INT8 QAD model at 480×832×81, M4 Max: **86.1 s → 10.3 s**
|
||||
of denoise. It composes with `--fast`; both together run the same clip in
|
||||
**4.5 s** of denoise.
|
||||
|
||||
```bash
|
||||
python examples/inference/basic/mlx_wan_prompt_to_video.py \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour, cinematic" \
|
||||
--height 480 --width 832 --num-frames 81 --fast-spatial \
|
||||
--output-path video_samples/fox_fast_spatial.mp4
|
||||
```
|
||||
|
||||
| Flag | Default | Meaning |
|
||||
|---|---|---|
|
||||
| `--fast-spatial` / `--no-fast-spatial` | off | enable spatial fast mode |
|
||||
| `--fast-spatial-scale` | 2 | denoise at 1/scale of each spatial axis |
|
||||
| `--fast-spatial-upsample-mode` | `lanczos` | pixel interpolation kernel (`lanczos`, `cubic`, `bilinear`, `nearest`) |
|
||||
| `--fast-spatial-sharpen` | 0.4 | light unsharp strength to counter resampling softness (0 disables) |
|
||||
|
||||
### The upsample must happen in pixel space
|
||||
|
||||
This is the one thing to get right. The obvious implementation — bilinearly
|
||||
upsample the finished latents and decode at the target size — **does not work**,
|
||||
and produces a distinctive failure: correct composition and silhouette under a
|
||||
smeared, hazy veil, with ringing along strong edges.
|
||||
|
||||
A Wan latent cell is a *learned code* for an 8×8 (Wan2.1) or 16×16 (Wan2.2)
|
||||
pixel block, not a low-pass sample of the image. The average of two adjacent
|
||||
codes is not the code of the averaged blocks; it is a vector the decoder was
|
||||
never trained on. Measured on Wan2.1-1.3B at 480×832, a 2× bilinear latent
|
||||
upsample destroys **62%** of the latent's high-frequency energy while leaving
|
||||
its overall magnitude intact — exactly the signature of that veil. At Wan2.2-5B
|
||||
the same operation degrades to black or noise.
|
||||
|
||||
Decoded RGB frames have no such problem: an image *is* a sampled 2-D signal, so
|
||||
Lanczos interpolation is the operation it was defined for. The result is soft —
|
||||
it carries stage-1's real detail budget and no more — but clean and coherent.
|
||||
|
||||
`--refine` gets away with a latent-space upsample only because a second DMD pass
|
||||
re-denoises the hand-off; spatial fast mode passes the latent straight to the
|
||||
decoder, so it cannot.
|
||||
|
||||
## Refine (`--refine`) stage-2 timesteps
|
||||
|
||||
`--refine` hands stage 1 to stage 2 as `(1 - sigma) * upsampled + sigma * noise`,
|
||||
where `sigma` comes from the *first* stage-2 timestep. FastWan's DMD grid opens
|
||||
at `t=1000`, which is `sigma == 1` exactly — so a stage-2 grid that starts there
|
||||
weights the stage-1 result at zero and refine silently degrades into a plain
|
||||
full-resolution run at twice the cost.
|
||||
|
||||
Left unset, `--refine-dmd-denoising-steps` now derives the stage-2 grid from the
|
||||
stage-1 one with leading full-noise steps dropped (`1000,757,522` → `757,522`).
|
||||
That keeps the pass on timesteps the distilled student was trained on while
|
||||
letting stage-1 structure through: hand-off `sigma = 0.757`, stage-1 weight
|
||||
`0.243`. Passing a grid that starts at full noise is now an error rather than a
|
||||
silently wasted pass.
|
||||
|
||||
The run prints the resolved hand-off so it is visible:
|
||||
|
||||
```
|
||||
[refine] stage-2 hand-off sigma=0.7568 (stage-1 weight 0.2432)
|
||||
```
|
||||
|
||||
There is a trade-off in choosing that grid. Later start = more of the draft
|
||||
survives, but fewer stage-2 steps. On Wan2.1 the default `757,522` gives weight
|
||||
0.243 with two steps; `--refine-dmd-denoising-steps 522` gives weight 0.478 with
|
||||
one. `--refine-sigma` decouples the noise level from the timestep entirely — it
|
||||
logs a warning, because the DiT is then told a timestep that does not match the
|
||||
noise it receives.
|
||||
|
||||
**Wan2.2-5B has a lower ceiling.** Its warped schedule maps `1000,757,522` to
|
||||
sigmas `1.000, 0.940, 0.845`, so the best available stage-1 weight is **0.060**
|
||||
(vs 0.243 at 1.3B). Un-warped (`--no-warp`) the same grid gives `1.000, 0.757,
|
||||
0.522` and a weight of 0.243 — but warping is what matches the FastVideo
|
||||
sampling schedule, so turning it off changes the timesteps the distilled student
|
||||
sees. Which is better at 5B is unresolved and needs a run on real 5B weights.
|
||||
@@ -76,10 +76,6 @@ surfaces:
|
||||
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
|
||||
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
|
||||
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
|
||||
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
|
||||
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
|
||||
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
|
||||
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
|
||||
|
||||
import itertools
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
@@ -20,40 +19,6 @@ GENERATED_DOC_PREFIXES = (
|
||||
"training/examples/",
|
||||
"distillation/examples/",
|
||||
)
|
||||
COOKBOOK_DATA = ROOT_DIR / "docs/assets/cookbook-recipes.json"
|
||||
COOKBOOK_SOURCE_ROOTS = (
|
||||
ROOT_DIR / "examples/inference",
|
||||
ROOT_DIR / "scripts/inference",
|
||||
)
|
||||
|
||||
|
||||
def validate_cookbook() -> None:
|
||||
"""Keep cookbook entries tied to checked-in runnable sources."""
|
||||
recipes = json.loads(COOKBOOK_DATA.read_text(encoding="utf-8")).get("recipes")
|
||||
if not isinstance(recipes, list) or not recipes:
|
||||
raise ValueError(f"{COOKBOOK_DATA}: recipes must be a non-empty list")
|
||||
|
||||
seen: set[str] = set()
|
||||
for recipe in recipes:
|
||||
required = ("id", "task", "label", "model", "source", "command")
|
||||
missing = {key for key in required if not recipe.get(key)}
|
||||
if missing:
|
||||
raise ValueError(f"Cookbook recipe is missing: {', '.join(sorted(missing))}")
|
||||
if recipe["id"] in seen:
|
||||
raise ValueError(f"Duplicate cookbook recipe id: {recipe['id']}")
|
||||
seen.add(recipe["id"])
|
||||
|
||||
source = (ROOT_DIR / recipe["source"]).resolve()
|
||||
if not any(source.is_relative_to(root.resolve()) for root in COOKBOOK_SOURCE_ROOTS):
|
||||
raise ValueError(f"Cookbook source is outside an approved directory: {recipe['source']}")
|
||||
if not source.is_file():
|
||||
raise ValueError(f"Cookbook source does not exist: {recipe['source']}")
|
||||
|
||||
source_text = source.read_text(encoding="utf-8")
|
||||
if recipe["model"] not in source_text:
|
||||
raise ValueError(f"Cookbook model is not present in {recipe['source']}: {recipe['model']}")
|
||||
if recipe["source"] not in recipe["command"]:
|
||||
raise ValueError(f"Cookbook command does not invoke its source: {recipe['id']}")
|
||||
|
||||
|
||||
def fix_case(text: str) -> str:
|
||||
@@ -571,7 +536,6 @@ def on_pre_build(config, **kwargs):
|
||||
MkDocs hook to generate examples before building the documentation.
|
||||
This function is called automatically by MkDocs' native hook system.
|
||||
"""
|
||||
validate_cookbook()
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
@@ -585,7 +549,6 @@ def on_page_context(context, page, **kwargs):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
validate_cookbook()
|
||||
print("Generating example documentation...")
|
||||
generate_examples(generate_main_index=True)
|
||||
print("Example documentation generated successfully!")
|
||||
|
||||
@@ -65,7 +65,6 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Quick Start](quick_start.md) - Generate your first video
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Quick Start Guide](quick_start.md) - Get started with your first video generation
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore scripts and notebooks
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore example scripts and notebooks
|
||||
|
||||
@@ -49,12 +49,10 @@ brew install ffmpeg
|
||||
|
||||
### Installation
|
||||
|
||||
FastWan's native Apple Silicon runtime requires the `mlx` extra.
|
||||
|
||||
#### With uv (recommended)
|
||||
|
||||
```bash
|
||||
uv pip install "fastvideo[mlx]"
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
@@ -62,7 +60,7 @@ uv pip install "fastvideo[mlx]"
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
uv pip install "fastvideo[mlx]"
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -78,13 +76,13 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[mlx]"
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[mlx]"
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -23,21 +23,61 @@ Also optionally install flash-attn:
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Choose a maintained recipe
|
||||
## Basic Usage
|
||||
|
||||
The cookbook selects complete, checked-in recipes instead of mixing model,
|
||||
parallelism, offload, and attention settings independently.
|
||||
### Text-to-Video Generation
|
||||
|
||||
[Open the inference cookbook](../cookbook/index.md){ .md-button .md-button--primary }
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
!!! tip "Need more control?"
|
||||
Start from a maintained recipe, then use the
|
||||
[configuration](../inference/configuration.md) and
|
||||
[optimization](../inference/optimizations.md) guides for supported changes.
|
||||
def main():
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
### Image-to-Video Generation
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main():
|
||||
# Create the generator
|
||||
model_name = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(model_name, num_gpus=1)
|
||||
|
||||
# Set up parameters with an initial image
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
generator.generate_video(prompt, sampling_param=sampling_param,
|
||||
output_path="my_videos/",
|
||||
save_video=True)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
## Next Steps
|
||||
|
||||
- [Inference Cookbook](../cookbook/index.md) - Choose a maintained recipe
|
||||
- [Installation Guide](installation.md) - Detailed installation instructions
|
||||
- [Configuration](../inference/configuration.md) - Learn about configuration options
|
||||
- [Examples](../inference/examples/examples_inference_index.md) - Explore more
|
||||
|
||||
@@ -64,7 +64,6 @@ column links a runnable script in `examples/inference/basic/` where one exists.
|
||||
| ltx2 | `FastVideo/LTX2-Distilled-Diffusers`<br>`FastVideo/LTX2.3-Distilled-Diffusers`<br>`FastVideo/LTX-2.3-Distilled-Diffusers` | T2V | [basic_ltx2_distilled.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2_distilled.py) |
|
||||
| ltx2 | `Lightricks/LTX-2.3`<br>`FastVideo/LTX2.3-base`<br>`FastVideo/LTX2.3-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
|
||||
| ltx2 | `Lightricks/LTX-2`<br>`FastVideo/LTX2-base`<br>`FastVideo/LTX2-Diffusers` | T2V | [basic_ltx2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_ltx2.py) |
|
||||
| mmaudio | `FastVideo/MMAudio-large-44k-v2-Diffusers` | V2A, T2A | [basic_mmaudio.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mmaudio.py) |
|
||||
| matrixgame | `FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers`<br>`FastVideo/Matrix-Game-2.0-Base-Diffusers`<br>`FastVideo/Matrix-Game-2.0-GTA-Diffusers`<br>`FastVideo/Matrix-Game-2.0-TempleRun-Diffusers`<br>`mignonjia/mg_longtuning_distilled_zelda`<br>`mignonjia/mg_sf_distilled_zelda_1k_steps`<br>`mignonjia/mg_sf_distilled_zelda`<br>`mignonjia/mg_causal_zelda`<br>`mignonjia/mg_bidirectional_zelda` | I2V | [basic_matrixgame2.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame2.py) |
|
||||
| matrixgame | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | I2V | [basic_matrixgame3.py](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_matrixgame3.py) |
|
||||
| minimax_h3 | `MiniMaxAI/MiniMax-H3` | T2V, I2V | [T2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_t2v.py)<br>[FL2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_fl2va.py)<br>[Ref2VA](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_minimax_h3_ref2va.py) |
|
||||
@@ -95,10 +94,6 @@ column links a runnable script in `examples/inference/basic/` where one exists.
|
||||
(`StableAudioT2AConfig` / `StableAudioOpenSmallConfig`); they are registered
|
||||
under the generic T2V workload option in the registry.
|
||||
|
||||
**Note (MMAudio)**: the registered Hugging Face model ID is reserved but not
|
||||
yet public. Follow the [MMAudio inference guide](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/pipelines/basic/mmaudio/README.md)
|
||||
to convert the official weights locally and set `MMAUDIO_MODEL_PATH`.
|
||||
|
||||
**Note (MiniMax H3)**: T2VA, FL2VA, and Ref2VA all generate video with stereo
|
||||
audio. Use the Ref2VA example when passing ordered image, video, or audio
|
||||
references.
|
||||
@@ -178,17 +173,6 @@ optimizations: absence means **untested**, not incompatible.
|
||||
| Matrix Game 3.0 Base Distilled | `FastVideo/Matrix-Game-3.0-Base-Distilled-Diffusers` | 720x1280 | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| GEN3C Cosmos 7B | `FastVideo/GEN3C-Cosmos-7B-Diffusers` | 704px1280p | ❌ | ❌ | ❌ | ⭕ | ⭕ |
|
||||
|
||||
## Apple Silicon native runtime
|
||||
|
||||
| Release path | Model | Mode | Validated hardware | Status |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| MLX FastWan T2V | FastWan-QAD-INT8-1.3B `[release model ID pending]` | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB unified-memory class, MLX 0.31.2 | Release candidate; requires release-owner visual sign-off |
|
||||
|
||||
This is a text-to-video-only source-install release. It is validated on the
|
||||
hardware listed above; MLX allocator caps are not evidence of support for a
|
||||
physical 16 GB Mac. See [Apple Silicon FastWan](../getting_started/installation/mps.md)
|
||||
for the supported command and release gates.
|
||||
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
***Lucy Edit Dev uses a non-commercial model license. FastVideo support is
|
||||
|
||||
@@ -33,14 +33,6 @@ 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!
|
||||
|
||||
@@ -1,180 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
|
||||
|
||||
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
|
||||
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
|
||||
sampler's shift-12 schedule instead of the base model's 50 steps, generating
|
||||
synchronized video and audio in one pipeline call.
|
||||
|
||||
The student was trained with block-sparse video attention (VSA, 64-token
|
||||
tiles) and its checkpoint carries the trained sparse-gate parameters
|
||||
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
|
||||
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
|
||||
dense (every tile is selected); raise the sparsity for additional speedup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
|
||||
# The HF repo is private while the MiniMax H3 Community License review
|
||||
# completes; until it flips public, pass --model-path with a local
|
||||
# snapshot of the release instead (e.g. the team export at
|
||||
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
|
||||
parser.add_argument("--prompt", required=True)
|
||||
parser.add_argument("--output", default="outputs/fasth3")
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1344)
|
||||
parser.add_argument("--num-frames", type=int, default=124)
|
||||
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
|
||||
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
|
||||
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
|
||||
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
|
||||
# default here is 5. Other grids are off-distribution.
|
||||
parser.add_argument("--steps",
|
||||
type=int,
|
||||
default=5,
|
||||
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
|
||||
"forwards. 5 (default) is the distilled 4-forward grid")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=4)
|
||||
parser.add_argument("--vsa-sparsity",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
|
||||
"exactly dense attention; the student was trained at 0.9")
|
||||
# 64 is the trained contract: the student was TRAINED with 64-token
|
||||
# (4,4,4) tiles, and its to_gate_compress gates were learned against
|
||||
# pooling at that granularity — keep 64 unless you are ablating.
|
||||
parser.add_argument("--vsa-tile-size",
|
||||
type=int,
|
||||
choices=(64, 256),
|
||||
default=64,
|
||||
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
|
||||
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
|
||||
"geometry for ablations")
|
||||
parser.add_argument("--vsa-kernel",
|
||||
choices=("triton", "sm100a"),
|
||||
default="triton",
|
||||
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
|
||||
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
|
||||
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
|
||||
"fastvideo-kernel build that carries the extension; if a precondition fails at "
|
||||
"run time the attention layer logs one warning and falls back to Triton. Only "
|
||||
"meaningful with --vsa-tile-size 64")
|
||||
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode",
|
||||
default=None,
|
||||
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
|
||||
parser.add_argument("--repeats",
|
||||
type=int,
|
||||
default=1,
|
||||
help="generate N times; with --torch-compile the first run pays "
|
||||
"compilation, so steady-state is the last repeat")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if args.vsa_kernel == "sm100a":
|
||||
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
|
||||
# before the pipeline boots so spawned GPU workers inherit it. The
|
||||
# kernel is forward-only and inference runs under no-grad, so every
|
||||
# denoising forward qualifies for the CUDA route.
|
||||
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
|
||||
|
||||
# Boot-time run configuration, folded into FastVideoArgs (the same route
|
||||
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
|
||||
# - attention_backend: the checkpoint carries trained to_gate_compress
|
||||
# gates, which only exist under the VSA-H3 backend — a dense-backend
|
||||
# load would reject them as unexpected weights. Layers that do not
|
||||
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
|
||||
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
|
||||
# branch pools per tile, and the gates were trained at 64 tokens/tile.
|
||||
experimental: dict[str, object] = {
|
||||
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
|
||||
"VSA_tile_size": args.vsa_tile_size,
|
||||
}
|
||||
if args.vsa_sparsity > 0.0:
|
||||
experimental["VSA_sparsity"] = args.vsa_sparsity
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
pipeline=PipelineSelection(experimental=experimental),
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=args.num_gpus > 1,
|
||||
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
dit_layerwise=False,
|
||||
text_encoder=True,
|
||||
vae=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=args.torch_compile,
|
||||
mode=args.compile_mode,
|
||||
),
|
||||
),
|
||||
))
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
# the base model is guidance-distilled; the student inherits it
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "fasth3.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
print(f"Output written to: {result.video_path}")
|
||||
if result.generation_time is not None:
|
||||
# machine-readable: benchmark harnesses parse this line to separate
|
||||
# generation from model-load time (last occurrence = steady state)
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
for _ in range(args.repeats - 1):
|
||||
result = generator.generate(request)
|
||||
if result.generation_time is not None:
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -27,8 +27,7 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
|
||||
@@ -27,8 +27,7 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
|
||||
@@ -24,8 +24,7 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("--model-path", default="MiniMaxAI/MiniMax-H3")
|
||||
# Rank-reduced AdaLN checkpoint (-39% params, -23 GiB VRAM): pass
|
||||
# --model-path noctuashap/MiniMax-H3-pruned-r16
|
||||
# (or a local dir produced by
|
||||
# scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py).
|
||||
# (or a local dir converted with tools/minimax_h3/fit_adaln_basis.py).
|
||||
# adaln_rank is read from the checkpoint config; no other flags needed.
|
||||
# Rank-reduced checkpoints are inference-only: training needs the
|
||||
# full-rank release.
|
||||
|
||||
@@ -1,44 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MMAudio large-44k-v2 video-to-audio example."""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--video-path", required=True)
|
||||
parser.add_argument("--output-path", default="outputs_audio/mmaudio.wav")
|
||||
parser.add_argument("--duration-seconds", type=float, default=8.0)
|
||||
parser.add_argument("--prompt", default="")
|
||||
parser.add_argument("--negative-prompt", default="music")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
os.environ.get(
|
||||
"MMAUDIO_MODEL_PATH",
|
||||
"converted_weights/mmaudio/large_44k_v2",
|
||||
),
|
||||
workload_type="v2a",
|
||||
num_gpus=1,
|
||||
)
|
||||
result = generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
video_path=args.video_path,
|
||||
audio_end_in_s=args.duration_seconds,
|
||||
output_path=args.output_path,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
)
|
||||
print(result["video_path"])
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,46 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tiny MLX RIFE frame-interpolation smoke test."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.mlx_runtime.rife_interp import interpolate, load_model
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="MLX RIFE 4.25 frame interpolation smoke test."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--self-test",
|
||||
action="store_true",
|
||||
help="Run a tiny two-frame interpolation test.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
if not args.self_test:
|
||||
raise SystemExit("Nothing to do; pass --self-test")
|
||||
|
||||
frame0 = np.zeros((64, 96, 3), dtype=np.uint8)
|
||||
frame1 = np.zeros((64, 96, 3), dtype=np.uint8)
|
||||
frame1[:, :, 0] = 255
|
||||
start = time.perf_counter()
|
||||
model = load_model()
|
||||
load_s = time.perf_counter() - start
|
||||
start = time.perf_counter()
|
||||
frames = interpolate([frame0, frame1], factor=2, model=model)
|
||||
interp_s = time.perf_counter() - start
|
||||
assert len(frames) == 3
|
||||
assert frames[1].shape == frame0.shape
|
||||
assert frames[1].dtype == np.uint8
|
||||
print(
|
||||
"MLX RIFE self-test passed: "
|
||||
f"load_s={load_s:.3f} interp_s={interp_s:.3f} shape={frames[1].shape}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,510 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""End-to-end Wan2.2-TI2V-5B generation on Apple Silicon (MLX DiT + MLX TAEHV).
|
||||
|
||||
Pipeline: torch/MPS UMT5 encode (shared with 1.3B) → MLXWan22DiT 3-step DMD
|
||||
(warped schedule, flow_shift=5) → MLX TAEHV decode (taew2_2.pth). Fully MLX
|
||||
on the heavy DiT + decode path.
|
||||
|
||||
PYTHONPATH=$PWD python examples/inference/basic/mlx_wan22_generate.py \
|
||||
--prompt "A red fox trotting through a snowy pine forest at golden hour" \
|
||||
--output-path video_samples/demo_5b/fox_5b_mlx.mp4
|
||||
|
||||
Decoder backends: ``taehv`` (default, MLX, ~seconds), ``taehv-torch`` (parity),
|
||||
``wan-vae`` (full AutoencoderKLWan on MPS, slow).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.mlx_runtime.fast_spatial import DEFAULT_FAST_SPATIAL_SHARPEN
|
||||
from fastvideo.mlx_runtime.frame_upsample import DEFAULT_PIXEL_UPSAMPLE_MODE, PIXEL_UPSAMPLE_MODES
|
||||
from fastvideo.mlx_runtime.memory import cleanup_mlx
|
||||
from fastvideo.mlx_runtime.prompt_cache import (
|
||||
fingerprint_digest,
|
||||
load_prompt_cache,
|
||||
save_prompt_cache,
|
||||
text_encoder_fingerprint,
|
||||
)
|
||||
from fastvideo.mlx_runtime.rife_interp import aligned_keyframe_count
|
||||
|
||||
FASTWAN21_MODEL_ID = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
|
||||
FASTWAN22_MODEL_ID = "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers"
|
||||
DEFAULT_HEIGHT = 448
|
||||
DEFAULT_WIDTH = 832
|
||||
DEFAULT_NUM_FRAMES = 121
|
||||
|
||||
def _resolve_model_paths(
|
||||
*,
|
||||
text_encoder_root: Path | None,
|
||||
dit_checkpoint: Path | None,
|
||||
dit_config: Path | None,
|
||||
vae_root: Path | None,
|
||||
mlx_checkpoint: Path | None,
|
||||
decode_backend: str,
|
||||
) -> tuple[Path, Path | None, Path | None, Path | None]:
|
||||
"""Download only the missing assets required by the selected Wan2.2 path."""
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
if text_encoder_root is None:
|
||||
text_encoder_root = Path(snapshot_download(
|
||||
FASTWAN21_MODEL_ID,
|
||||
allow_patterns=["tokenizer/*", "text_encoder/*"],
|
||||
))
|
||||
if mlx_checkpoint is None and (dit_checkpoint is None or dit_config is None):
|
||||
patterns = []
|
||||
if dit_checkpoint is None:
|
||||
patterns.append("transformer/diffusion_pytorch_model.safetensors")
|
||||
if dit_config is None:
|
||||
patterns.append("transformer/config.json")
|
||||
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=patterns))
|
||||
dit_checkpoint = dit_checkpoint or model_root / "transformer/diffusion_pytorch_model.safetensors"
|
||||
dit_config = dit_config or model_root / "transformer/config.json"
|
||||
if decode_backend == "wan-vae" and vae_root is None:
|
||||
model_root = Path(snapshot_download(FASTWAN22_MODEL_ID, allow_patterns=["vae/*"]))
|
||||
vae_root = model_root / "vae"
|
||||
return text_encoder_root, dit_checkpoint, dit_config, vae_root
|
||||
|
||||
|
||||
def _prompt_cache_fingerprint(
|
||||
*,
|
||||
prompt: str,
|
||||
prompt_used: str,
|
||||
enhance_prompt: bool,
|
||||
enhance_prompt_backend: str,
|
||||
text_encoder_root: Path,
|
||||
max_sequence_length: int,
|
||||
dtype: str,
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"prompt": prompt,
|
||||
"prompt_used": prompt_used,
|
||||
"enhance_prompt": enhance_prompt,
|
||||
"enhance_prompt_backend": enhance_prompt_backend,
|
||||
"text_encoder": text_encoder_fingerprint(text_encoder_root),
|
||||
"max_sequence_length": max_sequence_length,
|
||||
"dtype": dtype,
|
||||
}
|
||||
|
||||
|
||||
def _default_prompt_cache_path(fingerprint: dict[str, object]) -> Path:
|
||||
"""Content-addressed default cache file for a prompt fingerprint.
|
||||
|
||||
The Wan2.1 entrypoint caches prompt embeddings by default; this one only
|
||||
did so when handed an explicit ``--prompt-embeds-cache`` path, so every 5B
|
||||
run paid a full UMT5 encode (~45s on an M4 Max) even for a repeat prompt.
|
||||
The fingerprint already covers everything that changes the embedding, so
|
||||
hash it for the filename.
|
||||
"""
|
||||
digest = fingerprint_digest(fingerprint)[:32]
|
||||
return Path.home() / ".cache" / "fastvideo" / "prompt_embeds" / f"wan22_{digest}.npy"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="MLX Wan2.2-5B T2V (encode → DiT DMD → TAEHV/VAE decode)"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="A red fox trotting through a snowy pine forest at golden hour, cinematic",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=Path,
|
||||
default=Path("video_samples/demo_5b/fox_5b_mlx.mp4"),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-root",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Root with text_encoder/ + tokenizer/",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-embeds-cache",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Explicit .npy UMT5 embedding cache file. Overrides the automatic "
|
||||
"content-addressed cache (--prompt-cache).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-cache",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Cache prompt embeddings under ~/.cache/fastvideo/prompt_embeds so "
|
||||
"repeat runs skip the text encoder entirely. Default: on.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-device",
|
||||
choices=("auto", "cpu", "mps"),
|
||||
default="cpu",
|
||||
help="Device for UMT5 encoding. CPU is safest beside the 5B MLX DiT.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enhance-prompt",
|
||||
action="store_true",
|
||||
help="Apply deterministic local cinematic prompt enrichment before UMT5.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enhance-prompt-backend",
|
||||
choices=("template",),
|
||||
default="template",
|
||||
help="Prompt enrichment backend.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-checkpoint",
|
||||
type=Path,
|
||||
default=None,
|
||||
)
|
||||
parser.add_argument("--dit-config", type=Path, default=None)
|
||||
parser.add_argument(
|
||||
"--mlx-checkpoint",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Pre-quantized MLX DiT checkpoint directory. Rewrapped with Wan2.2 per-token conditioning.",
|
||||
)
|
||||
parser.add_argument("--vae-root", type=Path, default=None)
|
||||
parser.add_argument("--height", type=int, default=DEFAULT_HEIGHT)
|
||||
parser.add_argument("--width", type=int, default=DEFAULT_WIDTH)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=DEFAULT_NUM_FRAMES,
|
||||
help="Pixel frames (121 at 24fps = 5.04 seconds)",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=1234)
|
||||
parser.add_argument("--renoise-seed", type=int, default=0)
|
||||
parser.add_argument("--fps", type=int, default=24)
|
||||
parser.add_argument("--flow-shift", type=float, default=5.0)
|
||||
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
|
||||
parser.add_argument(
|
||||
"--no-warp",
|
||||
action="store_true",
|
||||
help="Disable schedule warping (debug only).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fast",
|
||||
action="store_true",
|
||||
help="Generate fewer frames then RIFE-interpolate to --num-frames.",
|
||||
)
|
||||
parser.add_argument("--fast-factor", type=int, default=2)
|
||||
parser.add_argument("--fast-sharpen", type=float, default=0.6)
|
||||
parser.add_argument(
|
||||
"--fast-spatial",
|
||||
action="store_true",
|
||||
help="Denoise and decode at reduced spatial resolution, then resample "
|
||||
"the decoded frames up to the target size.",
|
||||
)
|
||||
parser.add_argument("--fast-spatial-scale", type=int, default=2)
|
||||
parser.add_argument(
|
||||
"--fast-spatial-upsample-mode",
|
||||
choices=PIXEL_UPSAMPLE_MODES,
|
||||
default=DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
)
|
||||
parser.add_argument("--fast-spatial-sharpen", type=float, default=DEFAULT_FAST_SPATIAL_SHARPEN)
|
||||
parser.add_argument(
|
||||
"--refine",
|
||||
action="store_true",
|
||||
help="Two-pass DMD: coarse denoise, upsample/re-noise, full-res denoise.",
|
||||
)
|
||||
parser.add_argument("--refine-scale", type=int, default=2)
|
||||
parser.add_argument(
|
||||
"--refine-upsample-mode",
|
||||
choices=("bilinear", "nearest"),
|
||||
default="bilinear",
|
||||
)
|
||||
parser.add_argument("--no-refine-add-noise", action="store_true")
|
||||
parser.add_argument(
|
||||
"--decode-backend",
|
||||
choices=("taehv", "taehv-torch", "wan-vae"),
|
||||
default="taehv",
|
||||
)
|
||||
parser.add_argument("--save-latents", type=Path, default=None)
|
||||
parser.add_argument("--metrics-json", type=Path, default=None,
|
||||
help="Write measured run metadata as JSON for reports or galleries.")
|
||||
parser.add_argument(
|
||||
"--compile",
|
||||
action="store_true",
|
||||
help="Compile the DiT forward with mx.compile; fallback to eager on failure.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.fast_factor < 2:
|
||||
parser.error("--fast-factor must be at least 2")
|
||||
# --fast-spatial used to be rejected here because it upsampled the completed
|
||||
# 48-channel latent, which is out of distribution for the decoder and gave
|
||||
# black or noisy video. The upsample now runs on decoded frames, so the
|
||||
# latent never leaves the grid it was denoised on and the mode is usable.
|
||||
if args.refine and args.fast_spatial:
|
||||
print("[wan22] --refine takes precedence over --fast-spatial")
|
||||
args.text_encoder_root, args.dit_checkpoint, args.dit_config, args.vae_root = _resolve_model_paths(
|
||||
text_encoder_root=args.text_encoder_root,
|
||||
dit_checkpoint=args.dit_checkpoint,
|
||||
dit_config=args.dit_config,
|
||||
vae_root=args.vae_root,
|
||||
mlx_checkpoint=args.mlx_checkpoint,
|
||||
decode_backend=args.decode_backend,
|
||||
)
|
||||
target_frames = args.num_frames
|
||||
if args.fast:
|
||||
args.num_frames = aligned_keyframe_count(target_frames, args.fast_factor)
|
||||
print(
|
||||
f"[wan22 fast] generating {args.num_frames} frames, "
|
||||
f"RIFE {args.fast_factor}x -> {target_frames}"
|
||||
)
|
||||
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import (
|
||||
_postprocess_video,
|
||||
encode_prompt,
|
||||
make_rotary_embeddings,
|
||||
)
|
||||
from fastvideo.mlx_runtime.fast_spatial import plan_fast_spatial
|
||||
from fastvideo.mlx_runtime.refine import (
|
||||
default_refine_timesteps,
|
||||
plan_refine_resolutions,
|
||||
prepare_refine_latents,
|
||||
)
|
||||
from fastvideo.mlx_runtime.wan22 import (
|
||||
mlx_wan22_dit_from_diffusers_safetensors,
|
||||
mlx_wan22_dit_from_mlx_checkpoint,
|
||||
)
|
||||
from fastvideo.mlx_runtime.wan22_sample import build_wan22_dmd_schedule, sample_wan22_dmd
|
||||
from fastvideo.mlx_runtime.wan_vae import decode_latents_to_video
|
||||
|
||||
if args.mlx_checkpoint is not None:
|
||||
config = json.loads((args.mlx_checkpoint / "mlx_dit.json").read_text())["config"]
|
||||
else:
|
||||
config = json.loads(args.dit_config.read_text())
|
||||
patch_size = tuple(config.get("patch_size", (1, 2, 2)))
|
||||
if args.refine:
|
||||
active_plan = plan_refine_resolutions(
|
||||
height=args.height, width=args.width, num_frames=args.num_frames,
|
||||
spatial_scale=args.refine_scale, vae_spatial_compression=16,
|
||||
vae_temporal_compression=4, patch_size=patch_size, enabled=True,
|
||||
)
|
||||
spatial_mode = "refine"
|
||||
elif args.fast_spatial:
|
||||
fast_spatial_plan = plan_fast_spatial(
|
||||
height=args.height, width=args.width, num_frames=args.num_frames,
|
||||
spatial_scale=args.fast_spatial_scale, vae_spatial_compression=16,
|
||||
vae_temporal_compression=4, patch_size=patch_size,
|
||||
upsample_mode=args.fast_spatial_upsample_mode,
|
||||
sharpen=args.fast_spatial_sharpen, enabled=True,
|
||||
)
|
||||
active_plan = fast_spatial_plan.plan
|
||||
spatial_mode = "fast_spatial"
|
||||
else:
|
||||
active_plan = plan_refine_resolutions(
|
||||
height=args.height, width=args.width, num_frames=args.num_frames,
|
||||
spatial_scale=1, vae_spatial_compression=16, vae_temporal_compression=4,
|
||||
patch_size=patch_size, enabled=False,
|
||||
)
|
||||
spatial_mode = "off"
|
||||
lat_h, lat_w = active_plan.stage1_latent_height, active_plan.stage1_latent_width
|
||||
lat_t = active_plan.latent_frames
|
||||
in_ch = int(config["in_channels"])
|
||||
print(f"[5B] latent {in_ch}x{lat_t}x{lat_h}x{lat_w}", flush=True)
|
||||
|
||||
total_start = time.perf_counter()
|
||||
prompt_for_encode = args.prompt
|
||||
enhance_backend = None
|
||||
enhance_elapsed_s = 0.0
|
||||
if args.enhance_prompt:
|
||||
from fastvideo.mlx_runtime.prompt_enhance import enhance_prompt
|
||||
|
||||
enhancement = enhance_prompt(args.prompt, backend=args.enhance_prompt_backend)
|
||||
prompt_for_encode = enhancement.enhanced
|
||||
enhance_backend = enhancement.backend
|
||||
enhance_elapsed_s = enhancement.elapsed_s
|
||||
print(f"[enhance] backend={enhance_backend} in {enhance_elapsed_s:.2f}s", flush=True)
|
||||
print(f"[enhance] prompt: {prompt_for_encode}", flush=True)
|
||||
|
||||
t0 = time.perf_counter()
|
||||
prompt_cache_fingerprint = _prompt_cache_fingerprint(
|
||||
prompt=args.prompt,
|
||||
prompt_used=prompt_for_encode,
|
||||
enhance_prompt=args.enhance_prompt,
|
||||
enhance_prompt_backend=args.enhance_prompt_backend,
|
||||
text_encoder_root=args.text_encoder_root,
|
||||
max_sequence_length=512,
|
||||
dtype="fp16",
|
||||
)
|
||||
prompt_cache_path = args.prompt_embeds_cache
|
||||
if prompt_cache_path is None and args.prompt_cache:
|
||||
prompt_cache_path = _default_prompt_cache_path(prompt_cache_fingerprint)
|
||||
cached_embeds = load_prompt_cache(
|
||||
prompt_cache_path,
|
||||
prompt_cache_fingerprint,
|
||||
)
|
||||
if cached_embeds is not None:
|
||||
embeds = torch.from_numpy(cached_embeds).contiguous()
|
||||
else:
|
||||
embeds = encode_prompt(
|
||||
model_root=args.text_encoder_root,
|
||||
prompt=prompt_for_encode,
|
||||
max_sequence_length=512,
|
||||
device_arg=args.text_encoder_device,
|
||||
dtype_arg="fp16",
|
||||
)
|
||||
save_prompt_cache(
|
||||
prompt_cache_path,
|
||||
embeds.cpu().numpy(),
|
||||
prompt_cache_fingerprint,
|
||||
)
|
||||
ehs = mx.array(embeds.numpy()).astype(mx.float16)
|
||||
prompt_encode_s = time.perf_counter() - t0
|
||||
print(f"[5B] prompt encoded {tuple(ehs.shape)} in {prompt_encode_s:.1f}s", flush=True)
|
||||
|
||||
t1 = time.perf_counter()
|
||||
if args.mlx_checkpoint is not None:
|
||||
dit = mlx_wan22_dit_from_mlx_checkpoint(
|
||||
args.mlx_checkpoint,
|
||||
compile=args.compile,
|
||||
)
|
||||
else:
|
||||
dit = mlx_wan22_dit_from_diffusers_safetensors(
|
||||
args.dit_checkpoint,
|
||||
args.dit_config,
|
||||
dtype="fp16",
|
||||
compile=args.compile,
|
||||
)
|
||||
dit_load_s = time.perf_counter() - t1
|
||||
print(f"[5B] DiT loaded in {dit_load_s:.1f}s", flush=True)
|
||||
|
||||
freqs = make_rotary_embeddings(config, latent_frames=lat_t, latent_height=lat_h, latent_width=lat_w)
|
||||
gen = torch.Generator().manual_seed(args.seed)
|
||||
noise = mx.array(
|
||||
torch.randn(1, in_ch, lat_t, lat_h, lat_w, generator=gen, dtype=torch.float32).numpy()).astype(mx.float16)
|
||||
|
||||
steps = [int(s) for s in args.dmd_denoising_steps.split(",") if s.strip()]
|
||||
t2 = time.perf_counter()
|
||||
mx.reset_peak_memory()
|
||||
latents = sample_wan22_dmd(
|
||||
dit,
|
||||
ehs,
|
||||
noise,
|
||||
freqs,
|
||||
dmd_denoising_steps=steps,
|
||||
flow_shift=args.flow_shift,
|
||||
warp_denoising_step=not args.no_warp,
|
||||
seed=args.renoise_seed,
|
||||
)
|
||||
if spatial_mode == "refine":
|
||||
schedule, warped_steps = build_wan22_dmd_schedule(
|
||||
steps, flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
|
||||
)
|
||||
# The grid opens at sigma == 1, where the hand-off
|
||||
# `(1 - sigma) * upsampled + sigma * noise` weights stage 1 at zero and
|
||||
# refine silently becomes a plain full-res run. Drop the leading
|
||||
# full-noise steps so stage 1 actually reaches stage 2.
|
||||
stage2_warped = default_refine_timesteps(schedule, warped_steps)
|
||||
stage2_steps = steps[len(warped_steps) - len(stage2_warped):]
|
||||
sigma = schedule.sigma_for(stage2_warped[0])
|
||||
print(f"[5B refine] stage-2 steps={stage2_steps} sigma={sigma:.4f} "
|
||||
f"(stage-1 weight {1.0 - sigma:.4f})", flush=True)
|
||||
latents = prepare_refine_latents(
|
||||
latents, scale=args.refine_scale, sigma=sigma,
|
||||
add_noise_flag=not args.no_refine_add_noise,
|
||||
upsample_mode=args.refine_upsample_mode, seed=args.renoise_seed + 1,
|
||||
)
|
||||
freqs_stage2 = make_rotary_embeddings(
|
||||
config, latent_frames=lat_t,
|
||||
latent_height=active_plan.stage2_latent_height,
|
||||
latent_width=active_plan.stage2_latent_width,
|
||||
)
|
||||
latents = sample_wan22_dmd(
|
||||
dit, ehs, latents, freqs_stage2, dmd_denoising_steps=stage2_steps,
|
||||
flow_shift=args.flow_shift, warp_denoising_step=not args.no_warp,
|
||||
seed=args.renoise_seed + 2,
|
||||
)
|
||||
# spatial_mode == "fast_spatial" leaves the latents on the stage-1 grid;
|
||||
# the resample happens after decode, in _postprocess_video.
|
||||
denoise_s = time.perf_counter() - t2
|
||||
peak = mx.get_peak_memory() / (1024**3)
|
||||
print(f"[5B] denoise {len(steps)} steps in {denoise_s:.1f}s, peak {peak:.2f} GiB", flush=True)
|
||||
|
||||
latents_np = np.array(latents.astype(mx.float32))
|
||||
if args.save_latents is not None:
|
||||
args.save_latents.parent.mkdir(parents=True, exist_ok=True)
|
||||
np.savez(args.save_latents, latents=latents_np, prompt=args.prompt, seed=args.seed)
|
||||
print(f"[5B] wrote latents {args.save_latents}", flush=True)
|
||||
|
||||
if spatial_mode == "refine":
|
||||
del freqs_stage2
|
||||
del dit, latents, ehs, noise, freqs
|
||||
cleanup_mlx()
|
||||
|
||||
metrics = decode_latents_to_video(
|
||||
latents_np,
|
||||
args.output_path,
|
||||
fps=args.fps,
|
||||
backend=args.decode_backend,
|
||||
vae_dir=args.vae_root if args.decode_backend == "wan-vae" else None,
|
||||
z_dim=in_ch,
|
||||
)
|
||||
# One h264 round-trip for both post-decode passes (see _postprocess_video).
|
||||
rife_s = 0.0
|
||||
rife_request = ({
|
||||
"factor": args.fast_factor,
|
||||
"target_frames": target_frames,
|
||||
"sharpen": args.fast_sharpen,
|
||||
} if args.fast else None)
|
||||
spatial_request = fast_spatial_plan if spatial_mode == "fast_spatial" else None
|
||||
if rife_request is not None or spatial_request is not None:
|
||||
rife_start = time.perf_counter()
|
||||
_postprocess_video(
|
||||
video_path=args.output_path, fps=args.fps,
|
||||
rife=rife_request, spatial=spatial_request,
|
||||
)
|
||||
rife_s = time.perf_counter() - rife_start
|
||||
print(f"[5B] decoded via {metrics['backend']} in {metrics['decode_s']:.1f}s → {args.output_path}", flush=True)
|
||||
summary = {
|
||||
"output_path": str(args.output_path.resolve()),
|
||||
"prompt": args.prompt,
|
||||
"prompt_used": prompt_for_encode,
|
||||
"enhance_prompt": args.enhance_prompt,
|
||||
"enhance_backend": enhance_backend,
|
||||
"enhance_elapsed_s": round(enhance_elapsed_s, 3),
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"fps": args.fps,
|
||||
"target_frames": target_frames,
|
||||
"generated_frames": args.num_frames,
|
||||
"seed": args.seed,
|
||||
"renoise_seed": args.renoise_seed,
|
||||
"dmd_denoising_steps": steps,
|
||||
"flow_shift": args.flow_shift,
|
||||
"warp": not args.no_warp,
|
||||
"spatial_mode": spatial_mode,
|
||||
"fast": args.fast,
|
||||
"fast_factor": args.fast_factor if args.fast else None,
|
||||
"fast_spatial_scale": args.fast_spatial_scale if args.fast_spatial else None,
|
||||
"refine_scale": args.refine_scale if args.refine else None,
|
||||
"decode_backend": args.decode_backend,
|
||||
"prompt_encode_s": round(prompt_encode_s, 3),
|
||||
"dit_load_s": round(dit_load_s, 3),
|
||||
"denoise_s": round(denoise_s, 3),
|
||||
"decode_s": round(metrics["decode_s"], 3),
|
||||
"rife_s": round(rife_s, 3),
|
||||
"wall_total_s": round(time.perf_counter() - total_start, 3),
|
||||
"peak_gib": round(peak, 3),
|
||||
"latent_shape": [in_ch, lat_t, lat_h, lat_w],
|
||||
"stage2_latent_shape": [in_ch, lat_t, active_plan.stage2_latent_height, active_plan.stage2_latent_width],
|
||||
"mlx_checkpoint": str(args.mlx_checkpoint.resolve()) if args.mlx_checkpoint else None,
|
||||
}
|
||||
if args.metrics_json is not None:
|
||||
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.metrics_json.write_text(json.dumps(summary, indent=2) + "\n")
|
||||
print(f"[5B] wrote metrics {args.metrics_json}", flush=True)
|
||||
print(json.dumps(summary, indent=2), flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,130 +0,0 @@
|
||||
"""Compare Wan VAE and TAEHV decode on saved FastWan latents."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import DEFAULT_MODEL_ROOT, decode_latents_to_video
|
||||
|
||||
|
||||
def _torch_mps_memory() -> dict[str, int | None]:
|
||||
"""
|
||||
Report current and recommended memory usage for the MPS backend.
|
||||
|
||||
Returns:
|
||||
dict[str, int | None]: Memory metrics in bytes, or `None` values when
|
||||
PyTorch or MPS is unavailable.
|
||||
"""
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
return {
|
||||
"current_allocated_bytes": None,
|
||||
"driver_allocated_bytes": None,
|
||||
"recommended_max_bytes": None,
|
||||
}
|
||||
if not torch.backends.mps.is_available():
|
||||
return {
|
||||
"current_allocated_bytes": None,
|
||||
"driver_allocated_bytes": None,
|
||||
"recommended_max_bytes": None,
|
||||
}
|
||||
return {
|
||||
"current_allocated_bytes": int(torch.mps.current_allocated_memory()),
|
||||
"driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
|
||||
"recommended_max_bytes": int(torch.mps.recommended_max_memory()),
|
||||
}
|
||||
|
||||
|
||||
def _parse_backends(raw: str) -> list[str]:
|
||||
"""
|
||||
Parse and validate a comma-separated list of decoding backends.
|
||||
|
||||
Parameters:
|
||||
raw (str): Comma-separated backend names.
|
||||
|
||||
Returns:
|
||||
list[str]: Trimmed, supported backend names in input order.
|
||||
|
||||
Raises:
|
||||
ValueError: If the input contains an unsupported backend.
|
||||
"""
|
||||
backends = [backend.strip() for backend in raw.split(",") if backend.strip()]
|
||||
allowed = {"wan-vae", "taehv"}
|
||||
unknown = sorted(set(backends) - allowed)
|
||||
if unknown:
|
||||
raise ValueError(f"Unsupported decode backends: {unknown}")
|
||||
return backends
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""
|
||||
Benchmark selected Wan latent decoding backends and record their performance metrics.
|
||||
|
||||
Loads the specified latent array, decodes it with each selected backend, exports the
|
||||
results as MP4 files, and writes per-backend timing and Torch MPS memory metrics to
|
||||
`metrics.json`.
|
||||
"""
|
||||
parser = argparse.ArgumentParser(description="Benchmark decode backends on saved Wan/FastWan latents.")
|
||||
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
|
||||
parser.add_argument("--latents-path", type=Path, required=True)
|
||||
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_decode_benchmark"))
|
||||
parser.add_argument("--backends", default="wan-vae,taehv")
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--torch-device", default="auto")
|
||||
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--taehv-source-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-parallel", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
latents = np.load(args.latents_path)
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
rows = []
|
||||
|
||||
for backend in _parse_backends(args.backends):
|
||||
print(f"=== Decode backend: {backend} ===")
|
||||
before = _torch_mps_memory()
|
||||
start = time.perf_counter()
|
||||
output_path = args.output_dir / f"{args.latents_path.stem}_{backend}.mp4"
|
||||
decode_latents_to_video(
|
||||
model_root=args.model_root,
|
||||
latents_np=latents,
|
||||
output_path=output_path,
|
||||
fps=args.fps,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
backend=backend,
|
||||
taehv_source_path=args.taehv_source_path,
|
||||
taehv_checkpoint_path=args.taehv_checkpoint_path,
|
||||
taehv_parallel=args.taehv_parallel,
|
||||
)
|
||||
elapsed = time.perf_counter() - start
|
||||
after = _torch_mps_memory()
|
||||
metrics = {
|
||||
"backend": backend,
|
||||
"latents_path": str(args.latents_path),
|
||||
"latents_shape": list(latents.shape),
|
||||
"decode_export_s": elapsed,
|
||||
"torch_mps_current_before_bytes": before["current_allocated_bytes"],
|
||||
"torch_mps_current_after_bytes": after["current_allocated_bytes"],
|
||||
"torch_mps_driver_before_bytes": before["driver_allocated_bytes"],
|
||||
"torch_mps_driver_after_bytes": after["driver_allocated_bytes"],
|
||||
"torch_mps_recommended_max_bytes": after["recommended_max_bytes"],
|
||||
"output_path": str(output_path),
|
||||
}
|
||||
rows.append(metrics)
|
||||
print(json.dumps(metrics, indent=2))
|
||||
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(rows, indent=2))
|
||||
print(f"Wrote decode metrics to: {metrics_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,365 +0,0 @@
|
||||
"""Benchmark MLX FastWan quantization modes with one shared prompt encode."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import (
|
||||
DEFAULT_MODEL_ROOT,
|
||||
decode_latents_to_video,
|
||||
encode_prompt,
|
||||
make_rotary_embeddings,
|
||||
)
|
||||
from fastvideo.mlx_runtime.memory import cleanup_mlx
|
||||
|
||||
|
||||
def _parse_modes(raw: str) -> list[str]:
|
||||
"""
|
||||
Parse and validate a comma-separated list of quantization modes.
|
||||
|
||||
Parameters:
|
||||
raw (str): Comma-separated mode names.
|
||||
|
||||
Returns:
|
||||
list[str]: Normalized, whitespace-trimmed mode names.
|
||||
|
||||
Raises:
|
||||
ValueError: If any mode is unsupported.
|
||||
"""
|
||||
modes = [mode.strip() for mode in raw.split(",") if mode.strip()]
|
||||
allowed = {"none", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"}
|
||||
unknown = sorted(set(modes) - allowed)
|
||||
if unknown:
|
||||
raise ValueError(f"Unsupported modes: {unknown}")
|
||||
return modes
|
||||
|
||||
|
||||
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
|
||||
"""
|
||||
Compare candidate and baseline latent arrays using error and signal-quality metrics.
|
||||
|
||||
Parameters:
|
||||
candidate (np.ndarray): Latent array to evaluate.
|
||||
baseline (np.ndarray): Reference latent array for comparison.
|
||||
|
||||
Returns:
|
||||
dict[str, float]: Mean squared error, mean absolute error, maximum absolute
|
||||
error, and signal-to-noise ratio in decibels between the arrays.
|
||||
"""
|
||||
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
|
||||
mse = float(np.mean(np.square(diff)))
|
||||
mae = float(np.mean(np.abs(diff)))
|
||||
max_abs = float(np.max(np.abs(diff)))
|
||||
signal = float(np.mean(np.square(baseline.astype(np.float32))))
|
||||
return {
|
||||
"latent_mse_vs_fp16": mse,
|
||||
"latent_mae_vs_fp16": mae,
|
||||
"latent_max_abs_vs_fp16": max_abs,
|
||||
"latent_snr_db_vs_fp16": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
|
||||
}
|
||||
|
||||
|
||||
def _torch_mps_memory() -> dict[str, int | None]:
|
||||
"""
|
||||
Report PyTorch MPS memory statistics when PyTorch MPS is available.
|
||||
|
||||
Returns:
|
||||
dict[str, int | None]: A mapping of MPS memory metric names to byte counts, or `None` values when PyTorch or MPS is unavailable.
|
||||
"""
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
return {
|
||||
"torch_mps_current_allocated_bytes": None,
|
||||
"torch_mps_driver_allocated_bytes": None,
|
||||
"torch_mps_recommended_max_bytes": None,
|
||||
}
|
||||
if not torch.backends.mps.is_available():
|
||||
return {
|
||||
"torch_mps_current_allocated_bytes": None,
|
||||
"torch_mps_driver_allocated_bytes": None,
|
||||
"torch_mps_recommended_max_bytes": None,
|
||||
}
|
||||
return {
|
||||
"torch_mps_current_allocated_bytes": int(torch.mps.current_allocated_memory()),
|
||||
"torch_mps_driver_allocated_bytes": int(torch.mps.driver_allocated_memory()),
|
||||
"torch_mps_recommended_max_bytes": int(torch.mps.recommended_max_memory()),
|
||||
}
|
||||
|
||||
|
||||
def _decode_with_metrics(*, args, latents: np.ndarray, output_path: Path) -> dict[str, float | int | None | str]:
|
||||
"""
|
||||
Decode latents to a video and collect export timing and PyTorch MPS memory metrics.
|
||||
|
||||
Parameters:
|
||||
args: Configuration values for decoding and video export.
|
||||
latents (np.ndarray): Latent representation to decode.
|
||||
output_path (Path): Destination path for the exported video.
|
||||
|
||||
Returns:
|
||||
dict[str, float | int | None | str]: Video export duration and PyTorch MPS memory measurements.
|
||||
"""
|
||||
before = _torch_mps_memory()
|
||||
decode_start = time.perf_counter()
|
||||
decode_latents_to_video(
|
||||
model_root=args.model_root,
|
||||
latents_np=latents,
|
||||
output_path=output_path,
|
||||
fps=args.fps,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
backend=args.decode_backend,
|
||||
taehv_source_path=args.taehv_source_path,
|
||||
taehv_checkpoint_path=args.taehv_checkpoint_path,
|
||||
taehv_parallel=args.taehv_parallel,
|
||||
)
|
||||
decode_time = time.perf_counter() - decode_start
|
||||
after = _torch_mps_memory()
|
||||
return {
|
||||
"decode_export_s": decode_time,
|
||||
"decode_torch_mps_current_before_bytes": before["torch_mps_current_allocated_bytes"],
|
||||
"decode_torch_mps_current_after_bytes": after["torch_mps_current_allocated_bytes"],
|
||||
"decode_torch_mps_driver_before_bytes": before["torch_mps_driver_allocated_bytes"],
|
||||
"decode_torch_mps_driver_after_bytes": after["torch_mps_driver_allocated_bytes"],
|
||||
"decode_torch_mps_recommended_max_bytes": after["torch_mps_recommended_max_bytes"],
|
||||
}
|
||||
|
||||
|
||||
def _run_one_mode(
|
||||
*,
|
||||
mode: str,
|
||||
args,
|
||||
config: dict,
|
||||
checkpoint_path: Path,
|
||||
config_path: Path,
|
||||
prompt_embeds,
|
||||
freqs_cis,
|
||||
):
|
||||
"""
|
||||
Run denoising for one quantization mode and collect performance and memory metrics.
|
||||
|
||||
Parameters:
|
||||
mode (str): Quantization mode to benchmark.
|
||||
args: Benchmark configuration, including dtype, dimensions, seed, scheduler, and denoising settings.
|
||||
config (dict): Model configuration containing the input channel count.
|
||||
checkpoint_path (Path): Path to the transformer checkpoint.
|
||||
config_path (Path): Path to the transformer configuration.
|
||||
prompt_embeds: Encoded prompt embeddings shared across benchmark modes.
|
||||
freqs_cis: Rotary positional embeddings used during denoising.
|
||||
|
||||
Returns:
|
||||
dict: The mode name, generated latent array, and metrics for model loading,
|
||||
denoising, step timing, and MLX memory usage.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
from fastvideo.benchmarks.mlx_fastwan_bench import denoise_dmd_on_device
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
|
||||
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
|
||||
|
||||
mx_dtype = mx.float16 if args.mlx_dtype == "fp16" else mx.float32
|
||||
quantization = None if mode == "none" else mode
|
||||
latent_frames = (args.num_frames - 1) // 4 + 1
|
||||
latent_height = args.height // 8
|
||||
latent_width = args.width // 8
|
||||
|
||||
load_start = time.perf_counter()
|
||||
mx.clear_cache()
|
||||
mx.reset_peak_memory()
|
||||
dit = mlx_dit_from_diffusers_safetensors(
|
||||
checkpoint_path,
|
||||
config_path,
|
||||
dtype=args.mlx_dtype,
|
||||
quantization=quantization,
|
||||
)
|
||||
load_time = time.perf_counter() - load_start
|
||||
load_peak_memory = mx.get_peak_memory()
|
||||
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
|
||||
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
|
||||
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
|
||||
# Same torch generator sequence as the original host-round-trip loop
|
||||
# (initial latents first, then one re-noise draw per intermediate step),
|
||||
# so every mode still shares identical stochasticity.
|
||||
generator = torch.Generator(device="cpu").manual_seed(args.seed)
|
||||
latents_seed = torch.randn(
|
||||
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
).numpy()
|
||||
renoise_by_step = [
|
||||
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
|
||||
for _ in range(max(0, len(timesteps) - 1))
|
||||
]
|
||||
latents = mx.array(latents_seed).astype(mx_dtype)
|
||||
encoder_hidden_states = mx.array(prompt_embeds.numpy()).astype(mx_dtype)
|
||||
|
||||
denoise_start = time.perf_counter()
|
||||
mx.reset_peak_memory()
|
||||
latents_np, step_times = denoise_dmd_on_device(
|
||||
mx=mx,
|
||||
dit=dit,
|
||||
latents=latents,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
freqs_cis=freqs_cis,
|
||||
timesteps=timesteps,
|
||||
renoise_by_step=renoise_by_step,
|
||||
schedule=schedule,
|
||||
dmd_step=dmd_step,
|
||||
mx_dtype=mx_dtype,
|
||||
)
|
||||
denoise_time = time.perf_counter() - denoise_start
|
||||
denoise_peak_memory = mx.get_peak_memory()
|
||||
active_memory = mx.get_active_memory()
|
||||
return {
|
||||
"mode": mode,
|
||||
"latents": latents_np,
|
||||
"metrics": {
|
||||
"mlx_dit_load_s": load_time,
|
||||
"mlx_denoise_s": denoise_time,
|
||||
"mlx_denoise_first_step_s": step_times[0] if step_times else None,
|
||||
"mlx_load_peak_bytes": int(load_peak_memory),
|
||||
"mlx_denoise_peak_bytes": int(denoise_peak_memory),
|
||||
"mlx_active_after_denoise_bytes": int(active_memory),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""
|
||||
Run the MLX FastWan quantization benchmark for the selected modes and write latency, memory, output, and latent-difference metrics to the output directory.
|
||||
"""
|
||||
parser = argparse.ArgumentParser(description="Benchmark MLX FastWan quantization modes.")
|
||||
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
|
||||
parser.add_argument("--prompt", default="A snow leopard walks across a windy mountain ridge.")
|
||||
parser.add_argument("--height", type=int, default=192)
|
||||
parser.add_argument("--width", type=int, default=320)
|
||||
parser.add_argument("--num-frames", type=int, default=17)
|
||||
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
|
||||
parser.add_argument("--flow-shift", type=float, default=8.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=256)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--torch-device", default="auto")
|
||||
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--mlx-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--modes", default="none,int8,int4,mxfp8,mxfp4,nvfp4")
|
||||
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_quant_benchmark"))
|
||||
parser.add_argument("--decode-backend", choices=("none", "wan-vae", "taehv"), default="taehv")
|
||||
parser.add_argument("--taehv-source-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-parallel", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
mx.random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
config_path = args.model_root / "transformer/config.json"
|
||||
checkpoint_path = args.model_root / "transformer/diffusion_pytorch_model.safetensors"
|
||||
config = json.loads(config_path.read_text())
|
||||
latent_frames = (args.num_frames - 1) // 4 + 1
|
||||
latent_height = args.height // 8
|
||||
latent_width = args.width // 8
|
||||
|
||||
prompt_start = time.perf_counter()
|
||||
prompt_embeds = encode_prompt(
|
||||
model_root=args.model_root,
|
||||
prompt=args.prompt,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
)
|
||||
prompt_time = time.perf_counter() - prompt_start
|
||||
freqs_cis = make_rotary_embeddings(
|
||||
config,
|
||||
latent_frames=latent_frames,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
|
||||
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
|
||||
|
||||
baseline_latents = None
|
||||
rows = []
|
||||
for mode in _parse_modes(args.modes):
|
||||
print(f"=== MLX quant mode: {mode} ===")
|
||||
mode_start = time.perf_counter()
|
||||
try:
|
||||
result = _run_one_mode(
|
||||
mode=mode,
|
||||
args=args,
|
||||
config=config,
|
||||
checkpoint_path=checkpoint_path,
|
||||
config_path=config_path,
|
||||
prompt_embeds=prompt_embeds,
|
||||
freqs_cis=freqs_cis,
|
||||
)
|
||||
except UnsupportedMLXQuantizationError as exc:
|
||||
print(f"skipping mode (unsupported by this MLX build): {exc}")
|
||||
rows.append({"mode": mode, "status": "unsupported_by_mlx", "error": str(exc)})
|
||||
continue
|
||||
cleanup_mlx(mx)
|
||||
latents = result["latents"]
|
||||
if baseline_latents is None:
|
||||
baseline_latents = latents
|
||||
latent_path = args.output_dir / f"latents_{mode}.npy"
|
||||
np.save(latent_path, latents)
|
||||
|
||||
decode_time = 0.0
|
||||
decode_metrics = {}
|
||||
output_path = None
|
||||
if args.decode_backend != "none":
|
||||
output_path = args.output_dir / f"video_{mode}_{args.decode_backend}_{args.height}x{args.width}x{args.num_frames}.mp4"
|
||||
decode_metrics = _decode_with_metrics(args=args, latents=latents, output_path=output_path)
|
||||
decode_time = cast(float, decode_metrics["decode_export_s"])
|
||||
|
||||
mode_total = time.perf_counter() - mode_start
|
||||
mlx_denoise_peak_bytes = int(result["metrics"]["mlx_denoise_peak_bytes"])
|
||||
mlx_active_bytes = int(result["metrics"]["mlx_active_after_denoise_bytes"])
|
||||
metrics = {
|
||||
"mode": mode,
|
||||
"status": "ok",
|
||||
"prompt_encode_shared_s": prompt_time,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": args.num_frames,
|
||||
"decode_backend": args.decode_backend,
|
||||
"decode_export_s": decode_time,
|
||||
"mode_total_excluding_shared_prompt_s": mode_total,
|
||||
"mode_total_including_shared_prompt_s": mode_total + prompt_time,
|
||||
"latents_path": str(latent_path),
|
||||
"output_path": str(output_path) if output_path else None,
|
||||
"mlx_denoise_peak_gib": mlx_denoise_peak_bytes / (1024**3),
|
||||
"mlx_active_after_denoise_gib": mlx_active_bytes / (1024**3),
|
||||
"mlx_dit_peak_under_16gb": mlx_denoise_peak_bytes < 16 * 1024**3,
|
||||
"mlx_dit_active_under_16gb": mlx_active_bytes < 16 * 1024**3,
|
||||
"mac_16gb_status": (
|
||||
"dit_memory_fits_16gb_measured_decode_separately"
|
||||
if mlx_denoise_peak_bytes < 16 * 1024**3 else "dit_memory_exceeds_16gb"
|
||||
),
|
||||
**result["metrics"],
|
||||
**decode_metrics,
|
||||
**_latent_delta_metrics(latents, baseline_latents),
|
||||
}
|
||||
rows.append(metrics)
|
||||
print(json.dumps(metrics, indent=2))
|
||||
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(rows, indent=2))
|
||||
print(f"Wrote benchmark metrics to: {metrics_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,104 +0,0 @@
|
||||
"""Compare generated MP4s against a reference MP4 with simple pixel metrics."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def _read_video(path: Path) -> np.ndarray:
|
||||
"""
|
||||
Read all frames from a video file as an RGB NumPy array.
|
||||
|
||||
Parameters:
|
||||
path (Path): Path to the video file.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Video frames stacked along the first axis.
|
||||
|
||||
Raises:
|
||||
ValueError: If the video contains no readable frames.
|
||||
"""
|
||||
import cv2
|
||||
|
||||
cap = cv2.VideoCapture(str(path))
|
||||
frames = []
|
||||
try:
|
||||
while True:
|
||||
ok, frame_bgr = cap.read()
|
||||
if not ok:
|
||||
break
|
||||
frame_rgb = cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB)
|
||||
frames.append(frame_rgb)
|
||||
finally:
|
||||
cap.release()
|
||||
if not frames:
|
||||
raise ValueError(f"No frames read from {path}")
|
||||
return np.stack(frames, axis=0)
|
||||
|
||||
|
||||
def _metrics(candidate: np.ndarray, reference: np.ndarray) -> dict[str, float | int | list[int]]:
|
||||
"""
|
||||
Compute pixel-level comparison metrics between candidate and reference video frames.
|
||||
|
||||
Parameters:
|
||||
candidate (np.ndarray): Candidate video frames in frame, height, width, and channel order.
|
||||
reference (np.ndarray): Reference video frames with the same shape as the candidate.
|
||||
|
||||
Returns:
|
||||
dict[str, float | int | list[int]]: Frame dimensions and pixel comparison metrics, including MSE, MAE, maximum absolute difference, and PSNR in decibels.
|
||||
|
||||
Raises:
|
||||
ValueError: If the candidate and reference arrays have different shapes.
|
||||
"""
|
||||
if candidate.shape != reference.shape:
|
||||
raise ValueError(f"Shape mismatch: candidate={candidate.shape}, reference={reference.shape}")
|
||||
candidate_f = candidate.astype(np.float32)
|
||||
reference_f = reference.astype(np.float32)
|
||||
diff = candidate_f - reference_f
|
||||
mse = float(np.mean(np.square(diff)))
|
||||
mae = float(np.mean(np.abs(diff)))
|
||||
max_abs = float(np.max(np.abs(diff)))
|
||||
psnr = float(20.0 * np.log10(255.0 / np.sqrt(mse))) if mse > 0 else float("inf")
|
||||
return {
|
||||
"frames": int(candidate.shape[0]),
|
||||
"height": int(candidate.shape[1]),
|
||||
"width": int(candidate.shape[2]),
|
||||
"channels": int(candidate.shape[3]),
|
||||
"mse_vs_reference": mse,
|
||||
"mae_vs_reference": mae,
|
||||
"max_abs_vs_reference": max_abs,
|
||||
"psnr_db_vs_reference": psnr,
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Compare candidate MP4 videos with a reference and write pixel-level metrics to a JSON file."""
|
||||
parser = argparse.ArgumentParser(description="Compare MP4s against a reference MP4.")
|
||||
parser.add_argument("--reference", type=Path, required=True)
|
||||
parser.add_argument("--candidates", type=Path, nargs="+", required=True)
|
||||
parser.add_argument("--metrics-json", type=Path, required=True)
|
||||
args = parser.parse_args()
|
||||
|
||||
reference = _read_video(args.reference)
|
||||
rows = []
|
||||
for candidate_path in args.candidates:
|
||||
candidate = _read_video(candidate_path)
|
||||
row = {
|
||||
"reference_path": str(args.reference),
|
||||
"candidate_path": str(candidate_path),
|
||||
**_metrics(candidate, reference),
|
||||
}
|
||||
rows.append(row)
|
||||
print(json.dumps(row, indent=2))
|
||||
|
||||
args.metrics_json.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.metrics_json.write_text(json.dumps(rows, indent=2))
|
||||
print(f"Wrote video quality metrics to: {args.metrics_json}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -19,13 +19,16 @@ export FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_STACK=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1
|
||||
export FASTVIDEO_TORCH_PROFILER_WAIT_STEPS=2
|
||||
export FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS=1
|
||||
export FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS=1
|
||||
export FASTVIDEO_TORCH_PROFILER_DIR="../profiler_traces/wan_t2v_finetune/"
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_finetune"
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps 1 # Profiler captures every selected region until shutdown.
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 8
|
||||
|
||||
@@ -318,38 +318,6 @@ if(BUILD_CXX_KERNELS)
|
||||
|
||||
# Combined FastVideo Extension
|
||||
# Using name 'fastvideo_kernel_ops' to distinguish from the python package namespace
|
||||
# ---------------------------------------------------------------------------
|
||||
# VSA block-sparse attention forward, Blackwell (sm_100a) only.
|
||||
#
|
||||
# NOTE the "a" suffix: -arch=sm_100a is NOT enough -- it emits a plain sm_100 target and
|
||||
# ptxas rejects every tcgen05 / setmaxnreg instruction. The explicit gencode spelling
|
||||
# below is required, and matches the 10.0a entry in TORCH_CUDA_ARCH_LIST.
|
||||
#
|
||||
# Built for 64-token sparse blocks -- FastVideo's default (4,4,4) tiling, so no
|
||||
# tile-size change and no top-k granularity change is needed. VSA_BHSD
|
||||
# selects [B, H, S, D]. Both are compile-time; the Python is_supported() checks incoming
|
||||
# tensors against them so callers fall back to Triton rather than getting a wrong answer.
|
||||
# Read the ENVIRONMENT as well as the cache variable. When TORCH_CUDA_ARCH_LIST is
|
||||
# exported (build.sh, and `pip install` with it set) the branch above only prints it --
|
||||
# the cmake variable stays empty, so testing that alone silently skips the kernel and
|
||||
# leaves a build that succeeds with the op missing.
|
||||
set(ENABLE_VSA_SM100A OFF)
|
||||
set(_VSA_ARCH_LIST "${TORCH_CUDA_ARCH_LIST}")
|
||||
if(NOT _VSA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
|
||||
set(_VSA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
|
||||
endif()
|
||||
if(_VSA_ARCH_LIST MATCHES "(^|[; ,])(10\\.0a|100a|sm_100a)([; ,]|$)")
|
||||
set(ENABLE_VSA_SM100A ON)
|
||||
endif()
|
||||
if(ENABLE_VSA_SM100A)
|
||||
message(STATUS "fastvideo-kernel: building block_sparse_sm100a (Blackwell, 64- and 128-token blocks)")
|
||||
list(APPEND EXTENSION_SOURCES csrc/attention/block_sparse_sm100a.cu
|
||||
csrc/attention/block_sparse_blk128_sm100a.cu)
|
||||
set_source_files_properties(csrc/attention/block_sparse_sm100a.cu
|
||||
csrc/attention/block_sparse_blk128_sm100a.cu PROPERTIES
|
||||
COMPILE_OPTIONS "-gencode;arch=compute_100a,code=sm_100a;-DVSA_BHSD=true")
|
||||
endif()
|
||||
|
||||
Python_add_library(fastvideo_kernel_ops MODULE USE_SABI ${SKBUILD_SABI_VERSION} WITH_SOABI
|
||||
${EXTENSION_SOURCES}
|
||||
)
|
||||
@@ -365,14 +333,10 @@ if(BUILD_CXX_KERNELS)
|
||||
|
||||
# Build compile definitions list
|
||||
set(COMPILE_DEFS TORCH_EXTENSION_NAME=fastvideo_kernel_ops)
|
||||
if(ENABLE_VSA_SM100A)
|
||||
list(APPEND COMPILE_DEFS TK_COMPILE_BLOCK_SPARSE_VSA_SM100A)
|
||||
endif()
|
||||
if(ENABLE_TK_KERNELS)
|
||||
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
|
||||
endif()
|
||||
|
||||
|
||||
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
|
||||
|
||||
target_compile_options(fastvideo_kernel_ops PRIVATE
|
||||
|
||||
+11
-40
@@ -17,7 +17,7 @@ Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse forward/backward (VSA-128/256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
@@ -64,49 +64,28 @@ cd fastvideo-kernel
|
||||
./build.sh --rocm
|
||||
```
|
||||
|
||||
### Optional: FA4 CuTe block-sparse backend (VSA-128/256 fastpath)
|
||||
### Optional: FA4 CuTe block-sparse backend (VSA-256 fastpath)
|
||||
|
||||
The VSA-128/256 fastpaths (tile volume 128 or 256, on NVIDIA Blackwell / sm_100) route to the
|
||||
The VSA-256 fastpath (tile volume 256, on NVIDIA Blackwell / sm_100) routes to the
|
||||
FlashAttention-4 CuTe-DSL block-sparse kernel exposed as `flash_attn.cute`. This is
|
||||
an **optional** dependency: it is imported lazily, and `video_sparse_attn`
|
||||
transparently falls back to the Triton backend when it is absent (so the package is
|
||||
fully usable without it).
|
||||
|
||||
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`
|
||||
and the public/private forward-backward bridges in `flash_attn.cute.interface`) are provided upstream by
|
||||
The symbols the fastpath needs (`flash_attn.cute.block_sparsity.BlockSparseTensorsTorch`,
|
||||
`flash_attn.cute.interface._flash_attn_fwd`) are provided upstream by
|
||||
[Dao-AILab/flash-attention](https://github.com/Dao-AILab/flash-attention). Pin to
|
||||
commit `14c377950125c70b7a9dabf9c561fca53715ac7d`, the revision FastVideo pins as
|
||||
commit `940cd9680f3315f2f06b43ab5bea2c2cf2d96806`, the revision FastVideo pins as
|
||||
the `flash-attn-4` source in the repo-root `pyproject.toml`; other revisions may
|
||||
have incompatible block-sparse forward/backward interfaces.
|
||||
|
||||
Install it under its distribution name so its own runtime stack resolves with it.
|
||||
Do **not** pre-install `nvidia-cutlass-dsl` by hand: this revision pins
|
||||
`nvidia-cutlass-dsl==4.6.0.dev0` exactly, and a hand-installed 4.5.x floor either
|
||||
gets silently upgraded or, if something else holds it back, leaves the CuTe
|
||||
kernels broken.
|
||||
have an incompatible `_flash_attn_fwd` signature.
|
||||
|
||||
```bash
|
||||
pip install torchvision
|
||||
pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@14c377950125c70b7a9dabf9c561fca53715ac7d#subdirectory=flash_attn/cute"
|
||||
pip install "nvidia-cutlass-dsl>=4.5.0" torchvision
|
||||
pip install "git+https://github.com/Dao-AILab/flash-attention.git@940cd9680f3315f2f06b43ab5bea2c2cf2d96806#subdirectory=flash_attn/cute"
|
||||
```
|
||||
|
||||
That resolves `nvidia-cutlass-dsl` to 4.6.0.dev0 and `quack-kernels` to 0.5.3, a
|
||||
combination this revision works with. A mismatched CuTe DSL only surfaces when the
|
||||
kernel JIT-compiles, so the error points at CuTe internals rather than at the
|
||||
install:
|
||||
|
||||
| Error on first VSA-128/256 CuTe call | Cause |
|
||||
|---|---|
|
||||
| `TypeError: fmax() missing 1 required positional argument: 'b'` | `nvidia-cutlass-dsl` 4.5.x |
|
||||
| `AttributeError: module 'cutlass.cute.core' has no attribute 'ThrMma'` | `quack-kernels` older than 0.5.1 |
|
||||
| `ImportError: cannot import name 'alloc_reserved_mbarrier'` | `quack-kernels` 0.6.2 or newer |
|
||||
|
||||
An environment whose `flash_attn.cute` came from a prebuilt flash-attn wheel rather
|
||||
than from this pin hits the first row; that is what the overlay step in
|
||||
`docker/Dockerfile` works around.
|
||||
|
||||
The CuTe kernels JIT-compile on first use. Forward and backward are verified on
|
||||
Blackwell (sm_100) against `tests/test_vsa128_*.py` and `tests/test_vsa256_*.py`.
|
||||
The CuTe kernel JIT-compiles on first use. Verified on Blackwell (sm_100) against
|
||||
`tests/test_vsa256_forward*.py`.
|
||||
|
||||
## Usage
|
||||
|
||||
@@ -163,14 +142,6 @@ After building/installing `fastvideo-kernel`, run:
|
||||
```bash
|
||||
cd fastvideo-kernel
|
||||
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
|
||||
|
||||
# VSA-256 FA4 CuTe forward/backward on Blackwell
|
||||
python benchmarks/bench_vsa.py --block_size 256 --use_cute \
|
||||
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 20
|
||||
|
||||
# VSA-128 FA4 CuTe forward/backward on Blackwell
|
||||
python benchmarks/bench_vsa.py --block_size 128 --use_cute \
|
||||
--batch_size 1 --num_heads 12 --head_dim 128 --q_seq_lens 39936 --topk 40
|
||||
```
|
||||
|
||||
### TurboDiffusion Kernels
|
||||
|
||||
@@ -2,9 +2,8 @@
|
||||
"""
|
||||
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
|
||||
|
||||
This script benchmarks the autograd-enabled wrappers:
|
||||
- 64-token TK/Triton: fastvideo_kernel.block_sparse_attn.block_sparse_attn
|
||||
- 128/256-token Triton/CuTe: fastvideo_kernel.block_sparse_attn_256
|
||||
This script benchmarks the autograd-enabled wrapper:
|
||||
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
|
||||
|
||||
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
|
||||
"""
|
||||
@@ -24,6 +23,9 @@ try:
|
||||
except Exception as e: # pragma: no cover
|
||||
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
|
||||
def set_seed(seed: int = 42) -> None:
|
||||
random.seed(seed)
|
||||
@@ -39,11 +41,7 @@ def parse_arguments() -> argparse.Namespace:
|
||||
p.add_argument("--num_heads", type=int, default=12)
|
||||
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
|
||||
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
|
||||
p.add_argument("--q_seq_lens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[49152],
|
||||
help="Q sequence lengths (must be divisible by --block_size)")
|
||||
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
|
||||
p.add_argument("--kv_seq_lens",
|
||||
type=int,
|
||||
nargs="+",
|
||||
@@ -53,13 +51,9 @@ def parse_arguments() -> argparse.Namespace:
|
||||
p.add_argument("--rep", type=int, default=20)
|
||||
p.add_argument("--seed", type=int, default=42)
|
||||
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
|
||||
p.add_argument("--block_size", type=int, default=64, choices=[64, 128, 256])
|
||||
p.add_argument("--force_triton",
|
||||
action="store_true",
|
||||
help="Force wrapper to use Triton path (if supported by shapes).")
|
||||
p.add_argument("--use_cute",
|
||||
action="store_true",
|
||||
help="Use the optional FA4 CuTe forward/backward path (requires --block_size 128 or 256).")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
@@ -90,38 +84,18 @@ def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
|
||||
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
|
||||
|
||||
|
||||
def _configure_backend(args: argparse.Namespace) -> None:
|
||||
if args.use_cute and args.block_size not in (128, 256):
|
||||
raise ValueError("--use_cute requires --block_size 128 or 256")
|
||||
if args.use_cute and args.force_triton:
|
||||
raise ValueError("--use_cute and --force_triton are mutually exclusive")
|
||||
|
||||
if args.force_triton:
|
||||
os.environ.pop("FASTVIDEO_VSA_CUTEDSL", None)
|
||||
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
|
||||
elif args.use_cute:
|
||||
os.environ.pop("FASTVIDEO_VSA_TRITON", None)
|
||||
os.environ.pop("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", None)
|
||||
os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_arguments()
|
||||
set_seed(args.seed)
|
||||
_configure_backend(args)
|
||||
|
||||
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
|
||||
|
||||
if args.force_triton:
|
||||
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
|
||||
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128, block_sparse_attn_256
|
||||
|
||||
bs, h, d = args.batch_size, args.num_heads, args.head_dim
|
||||
block_size = args.block_size
|
||||
attention = {
|
||||
64: block_sparse_attn,
|
||||
128: block_sparse_attn_128,
|
||||
256: block_sparse_attn_256,
|
||||
}[block_size]
|
||||
kv_seq_lens = args.kv_seq_lens
|
||||
if kv_seq_lens is None:
|
||||
kv_seq_lens = args.q_seq_lens
|
||||
@@ -131,22 +105,20 @@ def main() -> None:
|
||||
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
|
||||
print(f"device: {torch.cuda.get_device_name(0)}")
|
||||
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
|
||||
print(f"block_size={block_size}")
|
||||
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
|
||||
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
|
||||
if args.use_cute:
|
||||
print("dispatch: FA4 CuTe")
|
||||
elif args.force_triton:
|
||||
if args.force_triton:
|
||||
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
|
||||
else:
|
||||
print("dispatch: SM90 if available, else Triton")
|
||||
|
||||
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
|
||||
if q_len % block_size != 0 or kv_len % block_size != 0:
|
||||
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by {block_size}")
|
||||
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
|
||||
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
|
||||
continue
|
||||
|
||||
num_q_blocks = q_len // block_size
|
||||
num_kv_blocks = kv_len // block_size
|
||||
num_q_blocks = q_len // BLOCK_M
|
||||
num_kv_blocks = kv_len // BLOCK_N
|
||||
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
|
||||
topk = min(topk, num_kv_blocks)
|
||||
|
||||
@@ -157,11 +129,11 @@ def main() -> None:
|
||||
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
|
||||
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
|
||||
|
||||
# Variable block sizes: default full logical blocks.
|
||||
variable_block_sizes = torch.full((num_kv_blocks, ), block_size, dtype=torch.int32, device="cuda")
|
||||
# Variable block sizes: default full blocks (64 tokens per KV block)
|
||||
variable_block_sizes = torch.full((num_kv_blocks, ), BLOCK_N, dtype=torch.int32, device="cuda")
|
||||
|
||||
def _fwd():
|
||||
return attention(q, k, v, block_map, variable_block_sizes)
|
||||
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
|
||||
|
||||
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
|
||||
|
||||
@@ -170,7 +142,7 @@ def main() -> None:
|
||||
q_ = q.detach().requires_grad_(True)
|
||||
k_ = k.detach().requires_grad_(True)
|
||||
v_ = v.detach().requires_grad_(True)
|
||||
o_, _aux_ = attention(q_, k_, v_, block_map, variable_block_sizes)
|
||||
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
|
||||
og = torch.randn_like(o_)
|
||||
loss = (o_ * og).sum()
|
||||
|
||||
@@ -184,7 +156,7 @@ def main() -> None:
|
||||
rep=max(5, args.rep // 2),
|
||||
)
|
||||
|
||||
flops = flops_sparse_attention(bs, h, d, q_len, topk, block_size)
|
||||
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
|
||||
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
|
||||
# Rough backward multiplier (attention backward typically ~2-3x forward)
|
||||
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
// block_sparse_blk128_sm100a.cu -- the 128-token-block instantiation of the torch binding.
|
||||
//
|
||||
// Same source as block_sparse_sm100a.cu with VSA_BLK128 set: the kernel and launch land in
|
||||
// namespace vsa_blk128 (distinct symbols, no ODR clash with the blk64 objects) and the
|
||||
// exported entry point becomes block_sparse_sm100a_blk128_fwd.
|
||||
#define VSA_BLK128 true
|
||||
#include "block_sparse_sm100a.cu"
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,201 +0,0 @@
|
||||
#ifndef BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
|
||||
#define BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
|
||||
|
||||
// Launch surface for the sm_100a VSA block-sparse FMHA forward.
|
||||
//
|
||||
// Everything a caller needs: a POD argument struct, a predicate saying whether this build can
|
||||
// run those arguments, and one launch entry point. The benchmark in
|
||||
// block_sparse_bench_sm100a.cu and the torch binding both go through here, so there is
|
||||
// one tensormap construction and one launch configuration rather than two that can drift.
|
||||
//
|
||||
// Two compile-time knobs select the four builds:
|
||||
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
|
||||
// VSA_BHSD false -> [token][head][dim] (BSHD), true -> [batch][head][token][dim] (BHSD)
|
||||
|
||||
#include "block_sparse_kernel_sm100a.cuh"
|
||||
|
||||
namespace VSA_NAMESPACE {
|
||||
|
||||
struct BlockSparseVsaArgs {
|
||||
const __nv_bfloat16* q;
|
||||
const __nv_bfloat16* k;
|
||||
const __nv_bfloat16* v; // natural layout; only blk128 reads it (blk64 still needs v_t)
|
||||
const __nv_bfloat16* v_t; // unused: kept so the bench's V_T buffer still binds
|
||||
__nv_bfloat16* o;
|
||||
float* lse; // [batch, num_heads, seqlen] fp32, or nullptr
|
||||
|
||||
const int* q2k_idx; // [batch*num_heads*num_blocks, max_kv] int32
|
||||
const int* q2k_num; // [batch*num_heads*num_blocks] int32
|
||||
const int* variable_block_sizes; // [num_blocks] int32, valid tokens per block
|
||||
|
||||
int batch;
|
||||
int num_heads;
|
||||
int seqlen;
|
||||
int head_dim;
|
||||
int num_blocks;
|
||||
int max_kv;
|
||||
float sm_scale;
|
||||
};
|
||||
|
||||
// cudaSuccess iff this build can run `a`. Deliberately conservative: the caller is expected
|
||||
// to fall back to its own implementation rather than get a wrong answer.
|
||||
__host__ inline cudaError_t block_sparse_supported(const BlockSparseVsaArgs& a) {
|
||||
if (a.head_dim != HEAD_DIM) return cudaErrorInvalidValue; // compile-time in the kernel
|
||||
if (a.num_blocks % 2 != 0) return cudaErrorInvalidValue; // a CTA owns an adjacent pair
|
||||
if (a.seqlen != a.num_blocks * BLOCK) return cudaErrorInvalidValue;
|
||||
if (a.max_kv < 1 || a.num_blocks < 1) return cudaErrorInvalidValue;
|
||||
if (a.q == nullptr || a.k == nullptr || a.o == nullptr) return cudaErrorInvalidValue;
|
||||
if (a.q2k_idx == nullptr || a.q2k_num == nullptr) return cudaErrorInvalidValue;
|
||||
// FastVideo always supplies this; without it padded keys would be attended as real zeros.
|
||||
if (a.variable_block_sizes == nullptr) return cudaErrorInvalidValue;
|
||||
// V is read MN-major at BOTH block sizes now, so no pre-transposed V_T is ever needed.
|
||||
if (a.v == nullptr) return cudaErrorInvalidValue;
|
||||
return cudaSuccess;
|
||||
}
|
||||
|
||||
__host__ inline cudaError_t launch_block_sparse_sm100a(const BlockSparseVsaArgs& a,
|
||||
cudaStream_t stream) {
|
||||
const cudaError_t sup = block_sparse_supported(a);
|
||||
if (sup != cudaSuccess) return sup;
|
||||
|
||||
const int B = a.batch, H = a.num_heads, S = a.seqlen, hd = a.head_dim;
|
||||
const int num_blocks = a.num_blocks, max_kv = a.max_kv;
|
||||
const long tq = (long)B * S;
|
||||
const int packed_mtiles_per_seq = num_blocks / 2;
|
||||
const int total_work = B * H * packed_mtiles_per_seq;
|
||||
constexpr bool BHSD = VSA_BHSD;
|
||||
|
||||
CUtensorMap tq_, tk_, tvt_, tv_, to_;
|
||||
{
|
||||
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16, BHSD ? (uint64_t)((long)B * H) : (uint64_t)H,
|
||||
BHSD ? (uint64_t)S : (uint64_t)tq, (uint64_t)Q_SUBTILES };
|
||||
uint64_t gs[3] = { BHSD ? (uint64_t)((long)S * hd) * 2u : (uint64_t)hd * 2u,
|
||||
BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
|
||||
(uint64_t)SUB_COLS_BF16 * 2u };
|
||||
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, 1u, (uint32_t)M_TILE, (uint32_t)Q_SUBTILES };
|
||||
uint32_t es[4] = { 1u, 1u, 1u, 1u };
|
||||
if (cuTensorMapEncodeTiled(&tq_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4,
|
||||
const_cast<__nv_bfloat16*>(a.q), gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
|
||||
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
|
||||
return cudaErrorInvalidValue;
|
||||
if (cuTensorMapEncodeTiled(&to_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 4, a.o, gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
|
||||
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
{
|
||||
uint64_t gd[4] = { (uint64_t)SUB_COLS_BF16,
|
||||
BHSD ? (uint64_t)S : (uint64_t)tq,
|
||||
BHSD ? (uint64_t)(hd / SUB_COLS_BF16)
|
||||
: (uint64_t)((long)H * hd / SUB_COLS_BF16),
|
||||
(uint64_t)((long)B * H) };
|
||||
uint64_t gs[3] = { BHSD ? (uint64_t)hd * 2u : (uint64_t)((long)H * hd) * 2u,
|
||||
(uint64_t)SUB_COLS_BF16 * 2u,
|
||||
(uint64_t)((long)S * hd) * 2u };
|
||||
uint32_t bd[4] = { (uint32_t)SUB_COLS_BF16, (uint32_t)BLOCK,
|
||||
BLK128 ? (uint32_t)K_SUBTILES : 1u, 1u };
|
||||
uint32_t es[4] = { 1u, 1u, 1u, 1u };
|
||||
if (cuTensorMapEncodeTiled(&tk_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
|
||||
const_cast<__nv_bfloat16*>(a.k), gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
|
||||
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
|
||||
return cudaErrorInvalidValue;
|
||||
// V map is byte-for-byte the K map over a.v: MN-major V needs no transpose (blk128).
|
||||
const __nv_bfloat16* vbase = a.v ? a.v : a.k;
|
||||
if (cuTensorMapEncodeTiled(&tv_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, BHSD ? 4 : 3,
|
||||
const_cast<__nv_bfloat16*>(vbase), gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
|
||||
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
// V_T map: blk64 only. Unused at blk128 but must still be a valid tensormap to pass by value.
|
||||
{
|
||||
const __nv_bfloat16* vt = a.v_t ? a.v_t : a.k;
|
||||
if constexpr (BLK128) {
|
||||
uint64_t gd[3] = { (uint64_t)SUB_COLS_BF16, (uint64_t)((long)H * hd),
|
||||
(uint64_t)((long)tq / SUB_COLS_BF16) };
|
||||
uint64_t gs[2] = { (uint64_t)tq * 2u, (uint64_t)SUB_COLS_BF16 * 2u };
|
||||
uint32_t bd[3] = { (uint32_t)SUB_COLS_BF16, (uint32_t)hd, (uint32_t)V_SUBTILES };
|
||||
uint32_t es[3] = { 1u, 1u, 1u };
|
||||
if (cuTensorMapEncodeTiled(&tvt_, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, 3,
|
||||
const_cast<__nv_bfloat16*>(vt), gd, gs, bd, es,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE, CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
|
||||
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) != CUDA_SUCCESS)
|
||||
return cudaErrorInvalidValue;
|
||||
} else {
|
||||
if (make_tma_2d_tiled(&tvt_, const_cast<__nv_bfloat16*>(vt), (long)H * hd, (int)tq, hd,
|
||||
SUB_COLS_BF16, 2, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16,
|
||||
CU_TENSOR_MAP_SWIZZLE_128B) != cudaSuccess)
|
||||
return cudaErrorInvalidValue;
|
||||
}
|
||||
}
|
||||
|
||||
const size_t smem =
|
||||
(size_t)2 * Q_TILE_BYTES + NUM_KV_STAGES * KV_RING_SLOT_BYTES
|
||||
+ (size_t)2 * M_TILE * HEAD_DIM * sizeof(__nv_bfloat16)
|
||||
+ (2 * NUM_KV_STAGES + 22) * 8
|
||||
+ (size_t)CLC_STAGES * (2 * 8 + 16) + 16
|
||||
+ 8
|
||||
+ (size_t)2 * STAT_REGIONS * STATS * sizeof(float)
|
||||
+ 256;
|
||||
|
||||
#ifndef VSA_NAMED_BAR
|
||||
#define VSA_NAMED_BAR false
|
||||
#endif
|
||||
#ifndef VSA_THROTTLE
|
||||
#define VSA_THROTTLE false
|
||||
#endif
|
||||
#ifndef VSA_USE_CLC
|
||||
#define VSA_USE_CLC true
|
||||
#endif
|
||||
constexpr bool FULL_NAMED_BAR = VSA_NAMED_BAR, EX2_EMU = true, SPLIT_P = true,
|
||||
SOFTMAX_THROTTLE = VSA_THROTTLE, USE_CLC = VSA_USE_CLC,
|
||||
Q_RASTER = true, MHA = true;
|
||||
auto kfn = &fmha_context_bf16_gen_kernel<32, FULL_NAMED_BAR, EX2_EMU, SPLIT_P,
|
||||
SOFTMAX_THROTTLE, USE_CLC, Q_RASTER, MHA,
|
||||
/*RESCALE_THRESHOLD=*/8, /*BHSD=*/VSA_BHSD>;
|
||||
cudaError_t e = cudaFuncSetAttribute(kfn, cudaFuncAttributeMaxDynamicSharedMemorySize, (int)smem);
|
||||
if (e != cudaSuccess) return e;
|
||||
|
||||
const unsigned long long magic0 = make_magic((unsigned)(H * packed_mtiles_per_seq));
|
||||
const unsigned long long magic1 = make_magic((unsigned)H);
|
||||
const unsigned long long magic2 = make_magic((unsigned)packed_mtiles_per_seq);
|
||||
const float scale_log2 = a.sm_scale * (float)M_LOG2E;
|
||||
|
||||
int numSM = 0;
|
||||
e = cudaDeviceGetAttribute(&numSM, cudaDevAttrMultiProcessorCount, 0);
|
||||
if (e != cudaSuccess) return e;
|
||||
const int num_ctas = USE_CLC ? total_work : (total_work < numSM ? total_work : numSM);
|
||||
dim3 grid(num_ctas, 1, 1), block(N_WARPS * 32, 1, 1);
|
||||
|
||||
if (USE_CLC) {
|
||||
cudaLaunchConfig_t cfg = {};
|
||||
cfg.gridDim = grid; cfg.blockDim = block; cfg.dynamicSmemBytes = smem; cfg.stream = stream;
|
||||
cudaLaunchAttribute cfgAttr[1];
|
||||
cfgAttr[0].id = cudaLaunchAttributeClusterDimension;
|
||||
cfgAttr[0].val.clusterDim.x = 1; cfgAttr[0].val.clusterDim.y = 1;
|
||||
cfgAttr[0].val.clusterDim.z = 1;
|
||||
cfg.attrs = cfgAttr; cfg.numAttrs = 1;
|
||||
return cudaLaunchKernelEx(&cfg, kfn, tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B,
|
||||
num_blocks, packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
|
||||
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
|
||||
}
|
||||
kfn<<<grid, block, smem, stream>>>(tq_, tk_, tvt_, tv_, to_, S, H, scale_log2, B, num_blocks,
|
||||
packed_mtiles_per_seq, max_kv, magic0, magic1, magic2,
|
||||
a.q2k_idx, a.q2k_num, a.variable_block_sizes, a.lse);
|
||||
return cudaGetLastError();
|
||||
}
|
||||
|
||||
} // namespace VSA_NAMESPACE
|
||||
|
||||
// Callers (the bench, the torch binding) keep using unqualified names; each translation unit
|
||||
// only ever sees the one configuration its VSA_BLK128 selected.
|
||||
using namespace VSA_NAMESPACE;
|
||||
|
||||
#endif // BLOCK_SPARSE_VSA_LAUNCH_SM100A_CUH
|
||||
@@ -1,114 +0,0 @@
|
||||
// block_sparse_sm100a.cu -- torch binding for the sm_100a VSA block-sparse FMHA forward.
|
||||
//
|
||||
// Forward only: returns (out, lse) so FastVideo's existing Triton backward keeps working
|
||||
// unchanged. lse is exactly the M tensor triton_block_sparse_attn_forward writes --
|
||||
// max(qk * qk_scale) + log2(l), [B, H, S] fp32 -- which is what lets
|
||||
// block_sparse_attn_backward_triton run against our forward untouched.
|
||||
//
|
||||
// The build is fixed at compile time by two flags, so one extension carries one configuration:
|
||||
// VSA_BLK128 false -> 64-token sparse blocks, true -> 128-token
|
||||
// VSA_BHSD false -> [B, S, H, D], true -> [B, H, S, D]
|
||||
#include <torch/extension.h>
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include "block_sparse_launch_sm100a.cuh"
|
||||
|
||||
namespace {
|
||||
|
||||
void check_qkv(const torch::Tensor& t, const char* name, int64_t B, int64_t H, int64_t S,
|
||||
int64_t D) {
|
||||
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
|
||||
TORCH_CHECK(t.scalar_type() == at::kBFloat16, name, " must be bfloat16, got ", t.scalar_type());
|
||||
TORCH_CHECK(t.dim() == 4, name, " must be 4-D, got ", t.dim(), " dims");
|
||||
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
|
||||
if (VSA_BHSD) {
|
||||
TORCH_CHECK(t.size(0) == B && t.size(1) == H && t.size(2) == S && t.size(3) == D, name,
|
||||
" has shape ", t.sizes(), ", expected [", B, ",", H, ",", S, ",", D, "]");
|
||||
} else {
|
||||
TORCH_CHECK(t.size(0) == B && t.size(1) == S && t.size(2) == H && t.size(3) == D, name,
|
||||
" has shape ", t.sizes(), ", expected [", B, ",", S, ",", H, ",", D, "]");
|
||||
}
|
||||
}
|
||||
|
||||
void check_index(const torch::Tensor& t, const char* name) {
|
||||
TORCH_CHECK(t.is_cuda(), name, " must be a CUDA tensor");
|
||||
TORCH_CHECK(t.scalar_type() == at::kInt, name, " must be int32, got ", t.scalar_type());
|
||||
TORCH_CHECK(t.is_contiguous(), name, " must be contiguous");
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
// The exported symbol carries the block size: block_sparse_sm100a_fwd is the 64-token build,
|
||||
// block_sparse_sm100a_blk128_fwd the 128-token one (block_sparse_blk128_sm100a.cu re-includes
|
||||
// this file with VSA_BLK128 set). The python backend picks by the metadata's block size.
|
||||
#if VSA_BLK128
|
||||
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_blk128_fwd
|
||||
#else
|
||||
#define BLOCK_SPARSE_SM100A_FWD block_sparse_sm100a_fwd
|
||||
#endif
|
||||
|
||||
// Returns {out} or {out, lse}. Layout of out matches the inputs.
|
||||
std::vector<torch::Tensor> BLOCK_SPARSE_SM100A_FWD(torch::Tensor q, torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
c10::optional<torch::Tensor> v_t,
|
||||
torch::Tensor q2k_idx,
|
||||
torch::Tensor q2k_num,
|
||||
torch::Tensor variable_block_sizes,
|
||||
double sm_scale, bool need_lse) {
|
||||
const at::cuda::OptionalCUDAGuard guard(device_of(q));
|
||||
|
||||
const int64_t B = q.size(0);
|
||||
const int64_t H = VSA_BHSD ? q.size(1) : q.size(2);
|
||||
const int64_t S = VSA_BHSD ? q.size(2) : q.size(1);
|
||||
const int64_t D = q.size(3);
|
||||
|
||||
check_qkv(q, "q", B, H, S, D);
|
||||
check_qkv(k, "k", B, H, S, D);
|
||||
check_qkv(v, "v", B, H, S, D);
|
||||
check_index(q2k_idx, "q2k_idx");
|
||||
check_index(q2k_num, "q2k_num");
|
||||
check_index(variable_block_sizes, "variable_block_sizes");
|
||||
|
||||
const int64_t num_blocks = variable_block_sizes.numel();
|
||||
const int64_t max_kv = q2k_idx.size(-1);
|
||||
TORCH_CHECK(S == num_blocks * BLOCK, "seqlen ", S, " must equal num_blocks (", num_blocks,
|
||||
") * ", BLOCK, "; FastVideo pads the sequence up to whole blocks");
|
||||
|
||||
auto out = torch::empty_like(q);
|
||||
torch::Tensor lse;
|
||||
if (need_lse) lse = torch::empty({B, H, S}, q.options().dtype(torch::kFloat32));
|
||||
|
||||
BlockSparseVsaArgs a{};
|
||||
a.q = reinterpret_cast<const __nv_bfloat16*>(q.data_ptr());
|
||||
a.k = reinterpret_cast<const __nv_bfloat16*>(k.data_ptr());
|
||||
a.v = reinterpret_cast<const __nv_bfloat16*>(v.data_ptr());
|
||||
a.v_t = v_t.has_value() ? reinterpret_cast<const __nv_bfloat16*>(v_t->data_ptr()) : nullptr;
|
||||
a.o = reinterpret_cast<__nv_bfloat16*>(out.data_ptr());
|
||||
a.lse = need_lse ? lse.data_ptr<float>() : nullptr;
|
||||
a.q2k_idx = q2k_idx.data_ptr<int>();
|
||||
a.q2k_num = q2k_num.data_ptr<int>();
|
||||
a.variable_block_sizes = variable_block_sizes.data_ptr<int>();
|
||||
a.batch = (int)B;
|
||||
a.num_heads = (int)H;
|
||||
a.seqlen = (int)S;
|
||||
a.head_dim = (int)D;
|
||||
a.num_blocks = (int)num_blocks;
|
||||
a.max_kv = (int)max_kv;
|
||||
a.sm_scale = (float)sm_scale;
|
||||
|
||||
// Report an unsupported regime loudly rather than returning plausible-looking wrong values.
|
||||
TORCH_CHECK(block_sparse_supported(a) == cudaSuccess,
|
||||
"block_sparse_sm100a: unsupported configuration -- requires head_dim==",
|
||||
HEAD_DIM, ", an even num_blocks, seqlen == num_blocks*", BLOCK,
|
||||
", and a variable_block_sizes tensor. Got head_dim=", D, " num_blocks=",
|
||||
num_blocks, " seqlen=", S);
|
||||
|
||||
const cudaError_t err = launch_block_sparse_sm100a(a, at::cuda::getCurrentCUDAStream());
|
||||
TORCH_CHECK(err == cudaSuccess,
|
||||
"block_sparse_sm100a launch failed: ", cudaGetErrorString(err));
|
||||
|
||||
if (need_lse) return {out, lse};
|
||||
return {out};
|
||||
}
|
||||
@@ -1,877 +0,0 @@
|
||||
// primitives.cuh -- device primitives for the sm_100a VSA block-sparse attention
|
||||
// forward: tcgen05 (alloc / mma / ld / st / commit / wait / fence), TMA load / store /
|
||||
// tensormap, mbarrier, cluster launch control, setmaxnreg, fast math, and the FMHA helpers.
|
||||
//
|
||||
// Generated and pruned to what the kernel reaches -- do not edit by hand.
|
||||
#pragma once
|
||||
#include <cstdint>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <cuda.h>
|
||||
#include <cuda_bf16.h>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_runtime.h>
|
||||
#include <cmath>
|
||||
#include <cassert>
|
||||
#include <cstring>
|
||||
#include <vector_types.h>
|
||||
|
||||
#ifndef CUDA_CHECK
|
||||
#define CUDA_CHECK(stmt) do { \
|
||||
cudaError_t _e = (stmt); \
|
||||
if (_e != cudaSuccess) { \
|
||||
fprintf(stderr, "CUDA error %s:%d: %s -> %s\n", \
|
||||
__FILE__, __LINE__, #stmt, cudaGetErrorString(_e)); \
|
||||
std::exit(1); \
|
||||
} \
|
||||
} while (0)
|
||||
#endif
|
||||
|
||||
__device__ __forceinline__
|
||||
uint64_t mbarrier_arrive(uint32_t mbar_smem) {
|
||||
uint64_t state;
|
||||
asm volatile("mbarrier.arrive.shared::cta.b64 %0, [%1];\n"
|
||||
: "=l"(state) : "r"(mbar_smem) : "memory");
|
||||
return state;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void mbarrier_arrive_nostate(uint32_t mbar_smem) {
|
||||
asm volatile("mbarrier.arrive.shared::cta.b64 _, [%0];\n"
|
||||
:: "r"(mbar_smem) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void mbarrier_arrive_cluster_default(uint32_t cluster_smem_addr) {
|
||||
asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];\n"
|
||||
:: "r"(cluster_smem_addr) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void mbarrier_arrive_expect_tx(uint32_t mbar_smem, uint32_t expected_bytes) {
|
||||
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;\n"
|
||||
:: "r"(mbar_smem), "r"(expected_bytes) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void mbarrier_wait_parity_suspend(uint32_t mbar_smem, uint32_t phase_parity) {
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .pred P1;\n"
|
||||
"LAB_WAIT:\n"
|
||||
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1, 10000000;\n"
|
||||
"@!P1 bra.uni LAB_WAIT;\n"
|
||||
"}\n"
|
||||
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void mbarrier_wait_parity(uint32_t mbar_smem, uint32_t phase_parity) {
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .pred P1;\n"
|
||||
"LAB_WAIT_HOT:\n"
|
||||
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
|
||||
"@!P1 bra.uni LAB_WAIT_HOT;\n"
|
||||
"}\n"
|
||||
:: "r"(mbar_smem), "r"(phase_parity) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void fence_proxy_async_shared_cta() {
|
||||
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void fence_proxy_async_shared() {
|
||||
asm volatile("fence.proxy.async.shared::cta;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void clc_try_cancel_async(
|
||||
uint32_t smem_dst, uint32_t mbar_smem) {
|
||||
asm volatile(
|
||||
"clusterlaunchcontrol.try_cancel.async.shared::cta.mbarrier::complete_tx::bytes.b128"
|
||||
" [%0], [%1];\n"
|
||||
:: "r"(smem_dst), "r"(mbar_smem) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void clc_load_response(
|
||||
uint32_t smem_slot, uint32_t& r0, uint32_t& r1,
|
||||
uint32_t& r2, uint32_t& r3) {
|
||||
asm volatile("ld.shared::cta.v4.b32 {%0, %1, %2, %3}, [%4];\n"
|
||||
: "=r"(r0), "=r"(r1), "=r"(r2), "=r"(r3)
|
||||
: "r"(smem_slot));
|
||||
}
|
||||
|
||||
template <int NUM_STAGES>
|
||||
struct MbarrierPhaseTracker {
|
||||
|
||||
uint32_t phase[NUM_STAGES];
|
||||
int idx;
|
||||
|
||||
__device__ __forceinline__
|
||||
void init() {
|
||||
for (int i = 0; i < NUM_STAGES; ++i) phase[i] = 0;
|
||||
idx = 0;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t current_phase() const { return phase[idx]; }
|
||||
|
||||
__device__ __forceinline__
|
||||
void advance() {
|
||||
phase[idx] ^= 1u;
|
||||
idx = (idx + 1) % NUM_STAGES;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
int stage() const { return idx; }
|
||||
};
|
||||
|
||||
template <int NUM_STAGES>
|
||||
struct PhaseTracker {
|
||||
int stage;
|
||||
uint32_t phase;
|
||||
|
||||
__device__ __forceinline__
|
||||
PhaseTracker() : stage(0), phase(0) {}
|
||||
|
||||
__device__ __forceinline__
|
||||
void advance() {
|
||||
stage++;
|
||||
if (stage == NUM_STAGES) {
|
||||
stage = 0;
|
||||
phase ^= 1;
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
int get_stage() const { return stage; }
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t get_phase() const { return phase; }
|
||||
};
|
||||
|
||||
template <int NUM_STAGES>
|
||||
struct EmptyPhaseTracker {
|
||||
int stage;
|
||||
uint32_t phase;
|
||||
|
||||
__device__ __forceinline__
|
||||
EmptyPhaseTracker() : stage(0), phase(1) {}
|
||||
|
||||
__device__ __forceinline__
|
||||
void advance() {
|
||||
stage++;
|
||||
if (stage == NUM_STAGES) {
|
||||
stage = 0;
|
||||
phase ^= 1;
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
int get_stage() const { return stage; }
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t get_phase() const { return phase; }
|
||||
};
|
||||
|
||||
template <int STAGES>
|
||||
__device__ __forceinline__
|
||||
void advance_stage_phase(int& stage, uint32_t& phase) {
|
||||
++stage;
|
||||
if (stage == STAGES) {
|
||||
stage = 0;
|
||||
phase ^= 1u;
|
||||
}
|
||||
}
|
||||
|
||||
static constexpr uint32_t SM100_CLC_PEER_MASK = 0xFEFFFFFF;
|
||||
|
||||
struct ClcTileInfo {
|
||||
int m_tile;
|
||||
int n_tile;
|
||||
bool valid;
|
||||
};
|
||||
|
||||
enum class ClcRasterOrder { AlongN, AlongM };
|
||||
|
||||
__device__ __forceinline__
|
||||
void clc_arrive_expect_tx_cta(uint32_t clc_full_local_addr, uint32_t tx_bytes) {
|
||||
if ((threadIdx.x & 31) == 0) {
|
||||
mbarrier_arrive_expect_tx(clc_full_local_addr, tx_bytes);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void clc_consumer_release(uint32_t clc_empty_local_addr) {
|
||||
uint32_t peer0_addr = clc_empty_local_addr & SM100_CLC_PEER_MASK;
|
||||
mbarrier_arrive_cluster_default(peer0_addr);
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void clc_consumer_release_cta(uint32_t clc_empty_local_addr) {
|
||||
mbarrier_arrive_nostate(clc_empty_local_addr);
|
||||
}
|
||||
|
||||
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER>
|
||||
__device__ __forceinline__
|
||||
ClcTileInfo clc_parse_response(uint32_t resp_smem_addr) {
|
||||
uint32_t d0, d1, d2, d3;
|
||||
fence_proxy_async_shared_cta();
|
||||
clc_load_response(resp_smem_addr, d0, d1, d2, d3);
|
||||
const int ctaid_x = static_cast<int>(d0);
|
||||
const int ctaid_y = static_cast<int>(d1 & 0xFFFFu);
|
||||
const bool valid = (d2 & 1u) != 0u;
|
||||
(void)d3;
|
||||
|
||||
ClcTileInfo info;
|
||||
info.valid = valid;
|
||||
if constexpr (ORDER == ClcRasterOrder::AlongN) {
|
||||
info.m_tile = ctaid_y / CLUSTER_SHAPE_M;
|
||||
info.n_tile = ctaid_x / CLUSTER_SHAPE_N;
|
||||
} else {
|
||||
info.m_tile = ctaid_x / CLUSTER_SHAPE_M;
|
||||
info.n_tile = ctaid_y / CLUSTER_SHAPE_N;
|
||||
}
|
||||
return info;
|
||||
}
|
||||
|
||||
template <int CLUSTER_SHAPE_M, int CLUSTER_SHAPE_N, ClcRasterOrder ORDER,
|
||||
int CTA_GROUP = 2, bool SUSPEND = false>
|
||||
__device__ __forceinline__
|
||||
ClcTileInfo clc_fetch_next_tile(
|
||||
uint64_t* clc_full_bar, uint64_t* clc_empty_bar, uint32_t* clc_response,
|
||||
int clc_cons_stage, uint32_t clc_cons_phase, bool do_release) {
|
||||
uint32_t full_addr = static_cast<uint32_t>(
|
||||
__cvta_generic_to_shared(&clc_full_bar[clc_cons_stage]));
|
||||
if constexpr (SUSPEND) mbarrier_wait_parity_suspend(full_addr, clc_cons_phase);
|
||||
else mbarrier_wait_parity(full_addr, clc_cons_phase);
|
||||
uint32_t resp_addr = static_cast<uint32_t>(
|
||||
__cvta_generic_to_shared(&clc_response[clc_cons_stage * 4]));
|
||||
ClcTileInfo t = clc_parse_response<
|
||||
CLUSTER_SHAPE_M, CLUSTER_SHAPE_N, ORDER>(resp_addr);
|
||||
if (do_release) {
|
||||
uint32_t empty_local = static_cast<uint32_t>(
|
||||
__cvta_generic_to_shared(&clc_empty_bar[clc_cons_stage]));
|
||||
if constexpr (CTA_GROUP == 1) {
|
||||
clc_consumer_release_cta(empty_local);
|
||||
} else {
|
||||
clc_consumer_release(empty_local);
|
||||
}
|
||||
}
|
||||
return t;
|
||||
}
|
||||
|
||||
template <int STAGES = 2>
|
||||
__device__ __forceinline__
|
||||
void clc_fetch_next_tile_advance(int& clc_cons_stage,
|
||||
uint32_t& clc_cons_phase) {
|
||||
advance_stage_phase<STAGES>(clc_cons_stage, clc_cons_phase);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ unsigned fdiv(unsigned n, unsigned long long pk) {
|
||||
unsigned M = (unsigned)pk;
|
||||
if (M == 0u) return n;
|
||||
return __umulhi(n, M) >> (unsigned)(pk >> 32);
|
||||
}
|
||||
__host__ inline unsigned long long make_magic(unsigned d) {
|
||||
if (d <= 1u) return 0ULL;
|
||||
unsigned l = 0; while ((1u << (l + 1)) <= d) ++l;
|
||||
unsigned p = 31u + l;
|
||||
unsigned long long m = ((1ull << p) + (unsigned long long)d - 1ull) / d;
|
||||
return (m & 0xffffffffULL) | ((unsigned long long)(p - 32u) << 32);
|
||||
}
|
||||
|
||||
template <int BARRIER_ID>
|
||||
__device__ __forceinline__
|
||||
void bar_sync(uint32_t thread_count) {
|
||||
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
|
||||
"bar.sync: BARRIER_ID must be in [0, 15]");
|
||||
asm volatile("bar.sync %0, %1;\n"
|
||||
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
|
||||
}
|
||||
|
||||
template <int BARRIER_ID>
|
||||
__device__ __forceinline__
|
||||
void bar_arrive(uint32_t thread_count) {
|
||||
static_assert(BARRIER_ID >= 0 && BARRIER_ID <= 15,
|
||||
"bar.arrive: BARRIER_ID must be in [0, 15]");
|
||||
asm volatile("bar.arrive %0, %1;\n"
|
||||
:: "n"(BARRIER_ID), "r"(thread_count) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void full_bar_arrive(int m_tile, int band) {
|
||||
switch (1 + m_tile * 4 + band) {
|
||||
case 1: bar_arrive<1>(64); break;
|
||||
case 2: bar_arrive<2>(64); break;
|
||||
case 3: bar_arrive<3>(64); break;
|
||||
case 4: bar_arrive<4>(64); break;
|
||||
case 5: bar_arrive<5>(64); break;
|
||||
case 6: bar_arrive<6>(64); break;
|
||||
case 7: bar_arrive<7>(64); break;
|
||||
case 8: bar_arrive<8>(64); break;
|
||||
}
|
||||
}
|
||||
__device__ __forceinline__ void full_bar_wait(int m_tile, int band) {
|
||||
switch (1 + m_tile * 4 + band) {
|
||||
case 1: bar_sync<1>(64); break;
|
||||
case 2: bar_sync<2>(64); break;
|
||||
case 3: bar_sync<3>(64); break;
|
||||
case 4: bar_sync<4>(64); break;
|
||||
case 5: bar_sync<5>(64); break;
|
||||
case 6: bar_sync<6>(64); break;
|
||||
case 7: bar_sync<7>(64); break;
|
||||
case 8: bar_sync<8>(64); break;
|
||||
}
|
||||
}
|
||||
|
||||
template <bool IS_CAUSAL, int K_TILE>
|
||||
__device__ __forceinline__ void mask_s_row_r2p(float* scores, int k_offset, int q_pos, int seqlen_k) {
|
||||
int n_keep = seqlen_k - k_offset;
|
||||
if constexpr (IS_CAUSAL) {
|
||||
const int causal = q_pos - k_offset + 1;
|
||||
n_keep = n_keep < causal ? n_keep : causal;
|
||||
}
|
||||
#pragma unroll
|
||||
for (int s = 0; s < K_TILE / 32; ++s) {
|
||||
int m = (s + 1) * 32 - n_keep;
|
||||
m = m < 0 ? 0 : (m > 32 ? 32 : m);
|
||||
const uint32_t keep = (m >= 32) ? 0u : (0xFFFFFFFFu >> m);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 32; ++i)
|
||||
if (!(keep & (1u << i))) scores[s * 32 + i] = -INFINITY;
|
||||
}
|
||||
}
|
||||
|
||||
template <int CTA_GROUP>
|
||||
__device__ __forceinline__ void tcgen05_alloc(uint32_t smem_dst_ptr,
|
||||
uint32_t n_cols) {
|
||||
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
|
||||
"tcgen05_alloc: CTA_GROUP must be 1 or 2");
|
||||
if constexpr (CTA_GROUP == 1) {
|
||||
asm volatile(
|
||||
"tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;\n"
|
||||
:: "r"(smem_dst_ptr), "r"(n_cols));
|
||||
} else {
|
||||
asm volatile(
|
||||
"tcgen05.alloc.cta_group::2.sync.aligned.shared::cta.b32 [%0], %1;\n"
|
||||
:: "r"(smem_dst_ptr), "r"(n_cols));
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_st_32x32b_x16(
|
||||
uint32_t tmem_addr, const uint32_t (&r)[16]) {
|
||||
asm volatile("tcgen05.st.sync.aligned.32x32b.x16.b32 "
|
||||
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
|
||||
"%11,%12,%13,%14,%15,%16};\n"
|
||||
:: "r"(tmem_addr),
|
||||
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
|
||||
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
|
||||
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
|
||||
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_st_32x32b_x32(
|
||||
uint32_t tmem_addr, const uint32_t (&r)[32]) {
|
||||
asm volatile("tcgen05.st.sync.aligned.32x32b.x32.b32 "
|
||||
"[%0], {%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,"
|
||||
"%11,%12,%13,%14,%15,%16,%17,%18,%19,%20,"
|
||||
"%21,%22,%23,%24,%25,%26,%27,%28,%29,%30,"
|
||||
"%31,%32};\n"
|
||||
:: "r"(tmem_addr),
|
||||
"r"(r[0]),"r"(r[1]),"r"(r[2]),"r"(r[3]),
|
||||
"r"(r[4]),"r"(r[5]),"r"(r[6]),"r"(r[7]),
|
||||
"r"(r[8]),"r"(r[9]),"r"(r[10]),"r"(r[11]),
|
||||
"r"(r[12]),"r"(r[13]),"r"(r[14]),"r"(r[15]),
|
||||
"r"(r[16]),"r"(r[17]),"r"(r[18]),"r"(r[19]),
|
||||
"r"(r[20]),"r"(r[21]),"r"(r[22]),"r"(r[23]),
|
||||
"r"(r[24]),"r"(r[25]),"r"(r[26]),"r"(r[27]),
|
||||
"r"(r[28]),"r"(r[29]),"r"(r[30]),"r"(r[31]));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_commit1_lead(uint32_t lead, uint32_t mbar_smem_addr) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred q;\n\t"
|
||||
"setp.ne.b32 q, %0, 0;\n\t"
|
||||
"@q tcgen05.commit.cta_group::1.mbarrier::arrive::one.b64 [%1];\n\t"
|
||||
"}\n"
|
||||
:: "r"(lead), "r"(mbar_smem_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_wait_st() {
|
||||
asm volatile("tcgen05.wait::st.sync.aligned;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_fence_before_thread_sync() {
|
||||
asm volatile("tcgen05.fence::before_thread_sync;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void tma_load_2d(uint32_t smem_dst, const void* tensormap_ptr,
|
||||
uint32_t mbar_smem, int coord_x, int coord_y) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.tensor.2d.shared::cluster.global.tile"
|
||||
".mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4}], [%2];\n"
|
||||
:: "r"(smem_dst), "l"(tensormap_ptr),
|
||||
"r"(mbar_smem), "r"(coord_x), "r"(coord_y)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void tma_load_3d(uint32_t smem_dst, const void* tensormap_ptr,
|
||||
uint32_t mbar_smem, int c0, int c1, int c2) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.tensor.3d.shared::cluster.global.tile"
|
||||
".mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4, %5}], [%2];\n"
|
||||
:: "r"(smem_dst), "l"(tensormap_ptr),
|
||||
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void tma_load_4d(uint32_t smem_dst, const void* tensormap_ptr,
|
||||
uint32_t mbar_smem, int c0, int c1, int c2, int c3) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.tensor.4d.shared::cluster.global.tile"
|
||||
".mbarrier::complete_tx::bytes"
|
||||
" [%0], [%1, {%3, %4, %5, %6}], [%2];\n"
|
||||
:: "r"(smem_dst), "l"(tensormap_ptr),
|
||||
"r"(mbar_smem), "r"(c0), "r"(c1), "r"(c2), "r"(c3)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
template <int CTA_GROUP>
|
||||
__device__ __forceinline__ void tcgen05_dealloc(uint32_t tmem_addr,
|
||||
uint32_t n_cols) {
|
||||
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
|
||||
"tcgen05_dealloc: CTA_GROUP must be 1 or 2");
|
||||
if constexpr (CTA_GROUP == 1) {
|
||||
asm volatile(
|
||||
"tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\n"
|
||||
:: "r"(tmem_addr), "r"(n_cols));
|
||||
} else {
|
||||
asm volatile(
|
||||
"tcgen05.dealloc.cta_group::2.sync.aligned.b32 %0, %1;\n"
|
||||
:: "r"(tmem_addr), "r"(n_cols));
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void tma_store_2d(const void* tensormap_ptr, int coord_x, int coord_y,
|
||||
uint32_t smem_src) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.tensor.2d.global.shared::cta.tile.bulk_group"
|
||||
" [%0, {%1, %2}], [%3];\n"
|
||||
:: "l"(tensormap_ptr), "r"(coord_x), "r"(coord_y),
|
||||
"r"(smem_src)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void tma_store_3d(const void* tensormap_ptr, int c0, int c1, int c2,
|
||||
uint32_t smem_src) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.tensor.3d.global.shared::cta.tile.bulk_group"
|
||||
" [%0, {%1, %2, %3}], [%4];\n"
|
||||
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(smem_src)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void tma_store_4d(const void* tensormap_ptr, int c0, int c1, int c2, int c3,
|
||||
uint32_t smem_src) {
|
||||
asm volatile(
|
||||
"cp.async.bulk.tensor.4d.global.shared::cta.tile.bulk_group"
|
||||
" [%0, {%1, %2, %3, %4}], [%5];\n"
|
||||
:: "l"(tensormap_ptr), "r"(c0), "r"(c1), "r"(c2), "r"(c3), "r"(smem_src)
|
||||
: "memory");
|
||||
}
|
||||
|
||||
inline cudaError_t make_tma_2d_tiled(
|
||||
CUtensorMap* out,
|
||||
const void* ptr, int rows, int cols, int box_rows, int box_cols,
|
||||
int elem_bytes, CUtensorMapDataType dtype,
|
||||
CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B,
|
||||
CUtensorMapL2promotion l2 = CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
|
||||
CUtensorMapFloatOOBfill oob = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE) {
|
||||
uint64_t globalDim[2] = { (uint64_t)cols, (uint64_t)rows };
|
||||
uint64_t globalStrides[1] = { (uint64_t)cols * (uint64_t)elem_bytes };
|
||||
uint32_t boxDim[2] = { (uint32_t)box_cols, (uint32_t)box_rows };
|
||||
uint32_t elemStrides[2] = { 1u, 1u };
|
||||
|
||||
CUresult r = cuTensorMapEncodeTiled(
|
||||
out, dtype, 2,
|
||||
const_cast<void*>(ptr), globalDim, globalStrides,
|
||||
boxDim, elemStrides,
|
||||
CU_TENSOR_MAP_INTERLEAVE_NONE,
|
||||
swizzle, l2, oob);
|
||||
return (r == CUDA_SUCCESS) ? cudaSuccess : cudaErrorInvalidValue;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void cp_async_bulk_commit_group() {
|
||||
asm volatile("cp.async.bulk.commit_group;\n" ::: "memory");
|
||||
}
|
||||
|
||||
template <int N>
|
||||
__device__ __forceinline__
|
||||
void cp_async_bulk_wait_group_read() {
|
||||
asm volatile("cp.async.bulk.wait_group.read %0;\n" :: "n"(N) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void mbarrier_init(uint32_t mbar_smem, uint32_t arrive_count) {
|
||||
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;\n"
|
||||
:: "r"(mbar_smem), "r"(arrive_count) : "memory");
|
||||
}
|
||||
|
||||
template <int CTA_GROUP>
|
||||
__device__ __forceinline__ void tcgen05_relinquish_alloc_permit() {
|
||||
static_assert(CTA_GROUP == 1 || CTA_GROUP == 2,
|
||||
"tcgen05_relinquish_alloc_permit: CTA_GROUP must be 1 or 2");
|
||||
if constexpr (CTA_GROUP == 1) {
|
||||
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;\n" ::);
|
||||
} else {
|
||||
asm volatile("tcgen05.relinquish_alloc_permit.cta_group::2.sync.aligned;\n" ::);
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void fence_mbarrier_init_release_cluster() {
|
||||
asm volatile("fence.mbarrier_init.release.cluster;\n" ::: "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_mma_f16_ss_lead(uint32_t lead,
|
||||
uint32_t tmem_c, uint64_t desc_a, uint64_t desc_b, uint32_t idesc,
|
||||
bool enable_input_d) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p, q;\n\t"
|
||||
"setp.ne.b32 q, %0, 0;\n\t"
|
||||
"setp.ne.b32 p, %5, 0;\n\t"
|
||||
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], %2, %3, %4, {%6, %7, %8, %9}, p;\n\t"
|
||||
"}\n"
|
||||
:: "r"(lead), "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(idesc),
|
||||
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_mma_f16_ts_1sm_lead(uint32_t lead,
|
||||
uint32_t tmem_c, uint32_t tmem_a, uint64_t desc_b, uint32_t idesc,
|
||||
bool enable_input_d) {
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p, q;\n\t"
|
||||
"setp.ne.b32 q, %0, 0;\n\t"
|
||||
"setp.ne.b32 p, %5, 0;\n\t"
|
||||
"@q tcgen05.mma.cta_group::1.kind::f16 [%1], [%2], %3, %4, {%6, %7, %8, %9}, p;\n\t"
|
||||
"}\n"
|
||||
:: "r"(lead), "r"(tmem_c), "r"(tmem_a), "l"(desc_b), "r"(idesc),
|
||||
"r"(enable_input_d ? 1u : 0u), "r"(0u), "r"(0u), "r"(0u), "r"(0u));
|
||||
}
|
||||
|
||||
enum class SmemSwizzleBlackwell : uint32_t {
|
||||
None = 0,
|
||||
B128_32atom = 1,
|
||||
B128 = 2,
|
||||
B64 = 4,
|
||||
B32 = 6,
|
||||
};
|
||||
|
||||
__device__ __host__ __forceinline__ uint64_t build_smem_desc_blackwell(
|
||||
uint32_t smem_addr,
|
||||
uint32_t stride_byte_offset,
|
||||
uint32_t leading_byte_offset,
|
||||
SmemSwizzleBlackwell swizzle = SmemSwizzleBlackwell::B128,
|
||||
uint32_t base_offset = 0) {
|
||||
uint64_t d = 0;
|
||||
d |= static_cast<uint64_t>((smem_addr >> 4) & 0x3FFF);
|
||||
d |= static_cast<uint64_t>((leading_byte_offset >> 4) & 0x3FFF) << 16;
|
||||
d |= static_cast<uint64_t>((stride_byte_offset >> 4) & 0x3FFF) << 32;
|
||||
d |= static_cast<uint64_t>(1) << 46;
|
||||
d |= (static_cast<uint64_t>(base_offset) & 0x7) << 49;
|
||||
d |= static_cast<uint64_t>(static_cast<uint32_t>(swizzle) & 0x7) << 61;
|
||||
return d;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t elect_one_sync() {
|
||||
uint32_t elected;
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"elect.sync %0|p, 0xffffffff;\n\t"
|
||||
"selp.b32 %0, 1, 0, p;\n\t"
|
||||
"}\n"
|
||||
: "=r"(elected));
|
||||
return elected;
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t elect_one_sync(uint32_t membermask) {
|
||||
uint32_t elected;
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .pred p;\n\t"
|
||||
"elect.sync %0|p, %1;\n\t"
|
||||
"selp.b32 %0, 1, 0, p;\n\t"
|
||||
"}\n"
|
||||
: "=r"(elected) : "r"(membermask));
|
||||
return elected;
|
||||
}
|
||||
|
||||
template <int N>
|
||||
__device__ __forceinline__
|
||||
void setmaxnreg_dec() {
|
||||
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
|
||||
"setmaxnreg_dec: N must be in [24, 256] and a multiple of 8");
|
||||
asm volatile("setmaxnreg.dec.sync.aligned.u32 %0;\n"
|
||||
:: "n"(N) : "memory");
|
||||
}
|
||||
|
||||
template <int N>
|
||||
__device__ __forceinline__
|
||||
void setmaxnreg_inc() {
|
||||
static_assert(N >= 24 && N <= 256 && N % 8 == 0,
|
||||
"setmaxnreg_inc: N must be in [24, 256] and a multiple of 8");
|
||||
asm volatile("setmaxnreg.inc.sync.aligned.u32 %0;\n"
|
||||
:: "n"(N) : "memory");
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t cvt_f32x2_to_bf16x2(float a, float b) {
|
||||
uint32_t r;
|
||||
asm volatile("cvt.rn.bf16x2.f32 %0, %2, %1;\n"
|
||||
: "=r"(r) : "f"(a), "f"(b));
|
||||
return r;
|
||||
}
|
||||
|
||||
namespace {
|
||||
__device__ __forceinline__ uint64_t f32x2_bits(float2 v) {
|
||||
uint64_t b; __builtin_memcpy(&b, &v, 8); return b;
|
||||
}
|
||||
__device__ __forceinline__ float2 f32x2_make(uint64_t b) {
|
||||
float2 v; __builtin_memcpy(&v, &b, 8); return v;
|
||||
}
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 fmul2(float2 a, float2 b) {
|
||||
uint64_t d;
|
||||
asm volatile("mul.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
|
||||
return f32x2_make(d);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 fadd2(float2 a, float2 b) {
|
||||
uint64_t d;
|
||||
asm volatile("add.f32x2 %0, %1, %2;\n" : "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)));
|
||||
return f32x2_make(d);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 ffma2(float2 a, float2 b, float2 c) {
|
||||
uint64_t d;
|
||||
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
|
||||
: "=l"(d) : "l"(f32x2_bits(a)), "l"(f32x2_bits(b)), "l"(f32x2_bits(c)));
|
||||
return f32x2_make(d);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 f32x2_splat(float s) { return make_float2(s, s); }
|
||||
|
||||
__device__ __forceinline__ float ex2_approx_f32(float z) {
|
||||
float d;
|
||||
asm volatile("ex2.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(z));
|
||||
return d;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float2 ex2_emu_f32x2(float x, float y) {
|
||||
uint32_t ox, oy;
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
".reg .f32 f1,f2,f3,f4,f5,f6,f7;\n\t"
|
||||
".reg .b64 l1,l2,l3,l4,l5,l6,l7,l8,l9,l10;\n\t"
|
||||
".reg .s32 r1,r2,r3,r4,r5,r6,r7,r8;\n\t"
|
||||
"max.f32 f1, %2, 0fC2FE0000;\n\t"
|
||||
"max.f32 f2, %3, 0fC2FE0000;\n\t"
|
||||
"mov.b64 l1, {f1, f2};\n\t"
|
||||
"mov.f32 f3, 0f4B400000;\n\t"
|
||||
"mov.b64 l2, {f3, f3};\n\t"
|
||||
"add.rm.f32x2 l7, l1, l2;\n\t"
|
||||
"sub.rn.f32x2 l8, l7, l2;\n\t"
|
||||
"sub.rn.f32x2 l9, l1, l8;\n\t"
|
||||
"mov.f32 f7, 0f3D9DF09D;\n\t"
|
||||
"mov.b64 l6, {f7, f7};\n\t"
|
||||
"mov.f32 f6, 0f3E6906A4;\n\t"
|
||||
"mov.b64 l5, {f6, f6};\n\t"
|
||||
"mov.f32 f5, 0f3F31F519;\n\t"
|
||||
"mov.b64 l4, {f5, f5};\n\t"
|
||||
"mov.f32 f4, 0f3F800000;\n\t"
|
||||
"mov.b64 l3, {f4, f4};\n\t"
|
||||
"fma.rn.f32x2 l10, l9, l6, l5;\n\t"
|
||||
"fma.rn.f32x2 l10, l10, l9, l4;\n\t"
|
||||
"fma.rn.f32x2 l10, l10, l9, l3;\n\t"
|
||||
"mov.b64 {r1, r2}, l7;\n\t"
|
||||
"mov.b64 {r3, r4}, l10;\n\t"
|
||||
"shl.b32 r5, r1, 23;\n\t"
|
||||
"add.s32 r7, r5, r3;\n\t"
|
||||
"shl.b32 r6, r2, 23;\n\t"
|
||||
"add.s32 r8, r6, r4;\n\t"
|
||||
"mov.b32 %0, r7;\n\t"
|
||||
"mov.b32 %1, r8;\n\t"
|
||||
"}\n"
|
||||
: "=r"(ox), "=r"(oy) : "f"(x), "f"(y));
|
||||
float2 r; __builtin_memcpy(&r.x, &ox, 4); __builtin_memcpy(&r.y, &oy, 4);
|
||||
return r;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ float rcp_approx_ftz_f32(float x) {
|
||||
float d;
|
||||
asm("rcp.approx.ftz.f32 %0, %1;\n" : "=f"(d) : "f"(x));
|
||||
return d;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ uint32_t make_idesc_table44(
|
||||
int M, int N,
|
||||
uint32_t dtype, uint32_t atype, uint32_t btype,
|
||||
bool transpose_a = false, bool transpose_b = false,
|
||||
bool negate_a = false, bool negate_b = false) {
|
||||
uint32_t idesc = 0;
|
||||
idesc |= (dtype & 0x3) << 4;
|
||||
idesc |= (atype & 0x7) << 7;
|
||||
idesc |= (btype & 0x7) << 10;
|
||||
idesc |= (negate_a ? 1u : 0u) << 13;
|
||||
idesc |= (negate_b ? 1u : 0u) << 14;
|
||||
idesc |= (transpose_a ? 1u : 0u) << 15;
|
||||
idesc |= (transpose_b ? 1u : 0u) << 16;
|
||||
idesc |= ((static_cast<uint32_t>(N) >> 3) & 0x3F) << 17;
|
||||
idesc |= ((static_cast<uint32_t>(M) >> 4) & 0x1F) << 24;
|
||||
return idesc;
|
||||
}
|
||||
|
||||
__device__ __forceinline__ uint32_t make_idesc_bf16_f32(
|
||||
int M, int N, bool ta = false, bool tb = false) {
|
||||
return make_idesc_table44(M, N, 1,
|
||||
1, 1, ta, tb);
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_ld_32x32b_x16(
|
||||
uint32_t tmem_addr, uint32_t (&r)[16]) {
|
||||
asm volatile("tcgen05.ld.sync.aligned.32x32b.x16.b32 "
|
||||
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
|
||||
"%10,%11,%12,%13,%14,%15}, [%16];\n"
|
||||
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
|
||||
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
|
||||
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
|
||||
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15])
|
||||
: "r"(tmem_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_ld_32x32b_x32(
|
||||
uint32_t tmem_addr, uint32_t (&r)[32]) {
|
||||
asm volatile("tcgen05.ld.sync.aligned.32x32b.x32.b32 "
|
||||
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
|
||||
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
|
||||
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
|
||||
"%30,%31}, [%32];\n"
|
||||
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
|
||||
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
|
||||
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
|
||||
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
|
||||
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
|
||||
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
|
||||
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
|
||||
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31])
|
||||
: "r"(tmem_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_ld_32x32b_x64(
|
||||
uint32_t tmem_addr, uint32_t (&r)[64]) {
|
||||
asm volatile("tcgen05.ld.sync.aligned.32x32b.x64.b32 "
|
||||
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
|
||||
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
|
||||
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
|
||||
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
|
||||
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
|
||||
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
|
||||
"%60,%61,%62,%63}, [%64];\n"
|
||||
: "=r"(r[0]),"=r"(r[1]),"=r"(r[2]),"=r"(r[3]),
|
||||
"=r"(r[4]),"=r"(r[5]),"=r"(r[6]),"=r"(r[7]),
|
||||
"=r"(r[8]),"=r"(r[9]),"=r"(r[10]),"=r"(r[11]),
|
||||
"=r"(r[12]),"=r"(r[13]),"=r"(r[14]),"=r"(r[15]),
|
||||
"=r"(r[16]),"=r"(r[17]),"=r"(r[18]),"=r"(r[19]),
|
||||
"=r"(r[20]),"=r"(r[21]),"=r"(r[22]),"=r"(r[23]),
|
||||
"=r"(r[24]),"=r"(r[25]),"=r"(r[26]),"=r"(r[27]),
|
||||
"=r"(r[28]),"=r"(r[29]),"=r"(r[30]),"=r"(r[31]),
|
||||
"=r"(r[32]),"=r"(r[33]),"=r"(r[34]),"=r"(r[35]),
|
||||
"=r"(r[36]),"=r"(r[37]),"=r"(r[38]),"=r"(r[39]),
|
||||
"=r"(r[40]),"=r"(r[41]),"=r"(r[42]),"=r"(r[43]),
|
||||
"=r"(r[44]),"=r"(r[45]),"=r"(r[46]),"=r"(r[47]),
|
||||
"=r"(r[48]),"=r"(r[49]),"=r"(r[50]),"=r"(r[51]),
|
||||
"=r"(r[52]),"=r"(r[53]),"=r"(r[54]),"=r"(r[55]),
|
||||
"=r"(r[56]),"=r"(r[57]),"=r"(r[58]),"=r"(r[59]),
|
||||
"=r"(r[60]),"=r"(r[61]),"=r"(r[62]),"=r"(r[63])
|
||||
: "r"(tmem_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__ void tcgen05_ld_32x32b_x128(
|
||||
uint32_t tmem_addr, uint32_t (&r)[128]) {
|
||||
asm volatile("tcgen05.ld.sync.aligned.32x32b.x128.b32 "
|
||||
"{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,"
|
||||
"%10,%11,%12,%13,%14,%15,%16,%17,%18,%19,"
|
||||
"%20,%21,%22,%23,%24,%25,%26,%27,%28,%29,"
|
||||
"%30,%31,%32,%33,%34,%35,%36,%37,%38,%39,"
|
||||
"%40,%41,%42,%43,%44,%45,%46,%47,%48,%49,"
|
||||
"%50,%51,%52,%53,%54,%55,%56,%57,%58,%59,"
|
||||
"%60,%61,%62,%63,%64,%65,%66,%67,%68,%69,"
|
||||
"%70,%71,%72,%73,%74,%75,%76,%77,%78,%79,"
|
||||
"%80,%81,%82,%83,%84,%85,%86,%87,%88,%89,"
|
||||
"%90,%91,%92,%93,%94,%95,%96,%97,%98,%99,"
|
||||
"%100,%101,%102,%103,%104,%105,%106,%107,%108,%109,"
|
||||
"%110,%111,%112,%113,%114,%115,%116,%117,%118,%119,"
|
||||
"%120,%121,%122,%123,%124,%125,%126,%127}, [%128];\n"
|
||||
: "=r"(r[ 0]),"=r"(r[ 1]),"=r"(r[ 2]),"=r"(r[ 3]),
|
||||
"=r"(r[ 4]),"=r"(r[ 5]),"=r"(r[ 6]),"=r"(r[ 7]),
|
||||
"=r"(r[ 8]),"=r"(r[ 9]),"=r"(r[ 10]),"=r"(r[ 11]),
|
||||
"=r"(r[ 12]),"=r"(r[ 13]),"=r"(r[ 14]),"=r"(r[ 15]),
|
||||
"=r"(r[ 16]),"=r"(r[ 17]),"=r"(r[ 18]),"=r"(r[ 19]),
|
||||
"=r"(r[ 20]),"=r"(r[ 21]),"=r"(r[ 22]),"=r"(r[ 23]),
|
||||
"=r"(r[ 24]),"=r"(r[ 25]),"=r"(r[ 26]),"=r"(r[ 27]),
|
||||
"=r"(r[ 28]),"=r"(r[ 29]),"=r"(r[ 30]),"=r"(r[ 31]),
|
||||
"=r"(r[ 32]),"=r"(r[ 33]),"=r"(r[ 34]),"=r"(r[ 35]),
|
||||
"=r"(r[ 36]),"=r"(r[ 37]),"=r"(r[ 38]),"=r"(r[ 39]),
|
||||
"=r"(r[ 40]),"=r"(r[ 41]),"=r"(r[ 42]),"=r"(r[ 43]),
|
||||
"=r"(r[ 44]),"=r"(r[ 45]),"=r"(r[ 46]),"=r"(r[ 47]),
|
||||
"=r"(r[ 48]),"=r"(r[ 49]),"=r"(r[ 50]),"=r"(r[ 51]),
|
||||
"=r"(r[ 52]),"=r"(r[ 53]),"=r"(r[ 54]),"=r"(r[ 55]),
|
||||
"=r"(r[ 56]),"=r"(r[ 57]),"=r"(r[ 58]),"=r"(r[ 59]),
|
||||
"=r"(r[ 60]),"=r"(r[ 61]),"=r"(r[ 62]),"=r"(r[ 63]),
|
||||
"=r"(r[ 64]),"=r"(r[ 65]),"=r"(r[ 66]),"=r"(r[ 67]),
|
||||
"=r"(r[ 68]),"=r"(r[ 69]),"=r"(r[ 70]),"=r"(r[ 71]),
|
||||
"=r"(r[ 72]),"=r"(r[ 73]),"=r"(r[ 74]),"=r"(r[ 75]),
|
||||
"=r"(r[ 76]),"=r"(r[ 77]),"=r"(r[ 78]),"=r"(r[ 79]),
|
||||
"=r"(r[ 80]),"=r"(r[ 81]),"=r"(r[ 82]),"=r"(r[ 83]),
|
||||
"=r"(r[ 84]),"=r"(r[ 85]),"=r"(r[ 86]),"=r"(r[ 87]),
|
||||
"=r"(r[ 88]),"=r"(r[ 89]),"=r"(r[ 90]),"=r"(r[ 91]),
|
||||
"=r"(r[ 92]),"=r"(r[ 93]),"=r"(r[ 94]),"=r"(r[ 95]),
|
||||
"=r"(r[ 96]),"=r"(r[ 97]),"=r"(r[ 98]),"=r"(r[ 99]),
|
||||
"=r"(r[100]),"=r"(r[101]),"=r"(r[102]),"=r"(r[103]),
|
||||
"=r"(r[104]),"=r"(r[105]),"=r"(r[106]),"=r"(r[107]),
|
||||
"=r"(r[108]),"=r"(r[109]),"=r"(r[110]),"=r"(r[111]),
|
||||
"=r"(r[112]),"=r"(r[113]),"=r"(r[114]),"=r"(r[115]),
|
||||
"=r"(r[116]),"=r"(r[117]),"=r"(r[118]),"=r"(r[119]),
|
||||
"=r"(r[120]),"=r"(r[121]),"=r"(r[122]),"=r"(r[123]),
|
||||
"=r"(r[124]),"=r"(r[125]),"=r"(r[126]),"=r"(r[127])
|
||||
: "r"(tmem_addr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
uint32_t smem_ptr_u32(const void* ptr) {
|
||||
return static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
|
||||
}
|
||||
|
||||
__device__ __forceinline__
|
||||
void sts_f32(uint32_t smem_addr, float val) {
|
||||
asm volatile("st.shared.f32 [%0], %1;" :: "r"(smem_addr), "f"(val) : "memory");
|
||||
}
|
||||
|
||||
@@ -28,31 +28,10 @@ void register_rms_norm(pybind11::module_ &);
|
||||
void register_layer_norm(pybind11::module_ &);
|
||||
void register_gemm(pybind11::module_ &);
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
|
||||
extern std::vector<torch::Tensor> block_sparse_sm100a_fwd(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
|
||||
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
|
||||
double sm_scale, bool need_lse);
|
||||
extern std::vector<torch::Tensor> block_sparse_sm100a_blk128_fwd(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, c10::optional<torch::Tensor> v_t,
|
||||
torch::Tensor q2k_idx, torch::Tensor q2k_num, torch::Tensor variable_block_sizes,
|
||||
double sm_scale, bool need_lse);
|
||||
#endif
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "FastVideo CUDA Kernels";
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE_VSA_SM100A
|
||||
m.def("block_sparse_sm100a_fwd",
|
||||
torch::wrap_pybind_function(block_sparse_sm100a_fwd),
|
||||
"VSA block-sparse attention forward, 64-token blocks (Blackwell sm100a)");
|
||||
m.def("block_sparse_sm100a_blk128_fwd",
|
||||
torch::wrap_pybind_function(block_sparse_sm100a_blk128_fwd),
|
||||
"VSA block-sparse attention forward, 128-token blocks (Blackwell sm100a)");
|
||||
#endif
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention (Hopper)");
|
||||
#endif
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""VSA-128/256 block-sparse attention wrappers.
|
||||
"""VSA-256 block-sparse attention wrapper.
|
||||
|
||||
The default 256-block path is Triton: it expands the logical 256-block map
|
||||
to the existing 64-block Triton kernel via a dense 4x4 expansion per logical
|
||||
@@ -7,8 +7,8 @@ edge ("route A"), and requires no optional dependencies.
|
||||
The FA4 CuTe block-sparse fastpath (intended for Blackwell sm_100+) is
|
||||
*opt-in* via ``FASTVIDEO_VSA_CUTEDSL=1``. It routes to
|
||||
:mod:`fastvideo_kernel.block_sparse_attn_cute_fwd`, which natively operates
|
||||
on 128-token Q/KV blocks (the 256 wrapper expands its logical KV map and
|
||||
sizes into that physical representation). The CuTe kernel
|
||||
on 128-token KV blocks (this wrapper expands the logical 256-block map /
|
||||
sizes into that physical 128-block representation). The CuTe kernel
|
||||
(``flash_attn.cute`` with block-sparsity) is an optional dependency,
|
||||
imported lazily only when this fastpath is selected.
|
||||
|
||||
@@ -35,7 +35,7 @@ _KV_BLOCK_TRITON = 64 # Existing Triton path uses 64-token KV blocks.
|
||||
|
||||
|
||||
def _resolve_backend() -> str:
|
||||
"""Pick the backend for the 128/256-block VSA paths.
|
||||
"""Pick the backend for the 256-block VSA path.
|
||||
|
||||
Default is Triton (no optional deps). The FA4 CuTe fastpath is opt-in
|
||||
via ``FASTVIDEO_VSA_CUTEDSL=1`` and requires the optional FA4 CuTe
|
||||
@@ -49,26 +49,6 @@ def _resolve_backend() -> str:
|
||||
return "triton"
|
||||
|
||||
|
||||
def _expand_mask_and_sizes_128_to_64(
|
||||
logical_mask_128: torch.Tensor,
|
||||
logical_kv_sizes_128: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Expand a [B, H, Qb128, KVb128] map to 64-token Triton tiles."""
|
||||
expanded_mask = logical_mask_128.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
|
||||
sizes_i32 = logical_kv_sizes_128.to(torch.int32)
|
||||
offsets = torch.tensor(
|
||||
[0, _KV_BLOCK_TRITON],
|
||||
dtype=torch.int32,
|
||||
device=sizes_i32.device,
|
||||
)
|
||||
expanded_sizes = torch.clamp(
|
||||
sizes_i32[:, None] - offsets[None, :],
|
||||
min=0,
|
||||
max=_KV_BLOCK_TRITON,
|
||||
).reshape(-1)
|
||||
return expanded_mask, expanded_sizes
|
||||
|
||||
|
||||
def _expand_mask_and_sizes_256_to_128(
|
||||
logical_mask_256: torch.Tensor,
|
||||
logical_kv_sizes_256: torch.Tensor,
|
||||
@@ -132,63 +112,6 @@ def _triton_via_route_a(
|
||||
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
|
||||
|
||||
|
||||
def _triton_via_route_a_128(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
logical_mask_128: torch.Tensor,
|
||||
logical_kv_sizes_128: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from .triton_kernels.index import map_to_index as triton_map_to_index
|
||||
|
||||
mask_64, sizes_64 = _expand_mask_and_sizes_128_to_64(logical_mask_128, logical_kv_sizes_128)
|
||||
q2k_idx, q2k_num = triton_map_to_index(mask_64.to(torch.bool))
|
||||
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, sizes_64)
|
||||
|
||||
|
||||
def block_sparse_attn_128(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
logical_block_map_128: torch.Tensor,
|
||||
logical_variable_block_sizes_128: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""VSA-128 sparse-branch entrypoint for [B, H, S, D] inputs."""
|
||||
if logical_block_map_128.dim() == 3:
|
||||
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
|
||||
|
||||
if _resolve_backend() == "triton":
|
||||
return _triton_via_route_a_128(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
|
||||
|
||||
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd
|
||||
return block_sparse_attn_cute_fwd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
|
||||
|
||||
|
||||
def block_sparse_attn_128_bshd(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
logical_block_map_128: torch.Tensor,
|
||||
logical_variable_block_sizes_128: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""VSA-128 sparse-branch entrypoint for [B, S, H, D] inputs."""
|
||||
if logical_block_map_128.dim() == 3:
|
||||
logical_block_map_128 = logical_block_map_128.unsqueeze(0)
|
||||
|
||||
if _resolve_backend() == "triton":
|
||||
out_bhsd, aux = _triton_via_route_a_128(
|
||||
q.transpose(1, 2).contiguous(),
|
||||
k.transpose(1, 2).contiguous(),
|
||||
v.transpose(1, 2).contiguous(),
|
||||
logical_block_map_128,
|
||||
logical_variable_block_sizes_128,
|
||||
)
|
||||
return out_bhsd.transpose(1, 2).contiguous(), aux
|
||||
|
||||
from .block_sparse_attn_cute_fwd import block_sparse_attn_cute_fwd_bshd
|
||||
return block_sparse_attn_cute_fwd_bshd(q, k, v, logical_block_map_128, logical_variable_block_sizes_128)
|
||||
|
||||
|
||||
def block_sparse_attn_256(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
|
||||
@@ -1,18 +1,18 @@
|
||||
"""FA4 CuTe-DSL block-sparse attention adapter.
|
||||
"""CuTe-DSL block-sparse attention forward kernel.
|
||||
|
||||
This module adapts VSA's ``(block_map, variable_block_sizes)`` inputs into
|
||||
FA4's forward and backward ``BlockSparseTensorsTorch`` representations.
|
||||
FA4's public ``flash_attn_func`` owns the forward/backward autograd bridge.
|
||||
Thin wrapper around `flash_attn.cute.interface._flash_attn_fwd` that adapts
|
||||
VSA's `(block_map, variable_block_sizes)` inputs into FA4's
|
||||
`BlockSparseTensorsTorch` representation and the per-KV-block validity mask.
|
||||
|
||||
Both [B, H, S, D] (BHSD) and [B, S, H, D] (BSHD) entrypoints are provided.
|
||||
The BSHD variant is preferred from VSA-128/256 callers to avoid layout
|
||||
The BSHD variant is preferred from VSA-256 callers to avoid layout
|
||||
round-trips on the hot path.
|
||||
|
||||
The FA4 CuTe block-sparse kernel (``flash_attn.cute`` with
|
||||
``block_sparsity``) is an *optional* dependency: it is imported lazily and
|
||||
only exercised when the VSA-128/256 CuTe fastpath is explicitly selected
|
||||
(``FASTVIDEO_VSA_CUTEDSL=1``). The default path is Triton and does not require
|
||||
it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
|
||||
only exercised when the VSA-256 CuTe fastpath is explicitly selected
|
||||
(``FASTVIDEO_VSA_CUTEDSL=1``). The default VSA-256 path is Triton and does
|
||||
not require it. Also needs ``nvidia-cutlass-dsl`` and ``quack-kernels``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -22,14 +22,13 @@ from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
_FA4_IMPORT_HINT = ("VSA-128/256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
|
||||
_FA4_IMPORT_HINT = ("VSA-256 CuTe fastpath requires a FlashAttention-4 CuTe build that "
|
||||
"provides `flash_attn.cute` with block-sparsity support (plus "
|
||||
"`nvidia-cutlass-dsl` and `quack-kernels`). This is an optional "
|
||||
"dependency; the default path is Triton. Install the FA4 CuTe "
|
||||
"dependency; the default VSA-256 path is Triton. Install the FA4 CuTe "
|
||||
"build and set FASTVIDEO_VSA_CUTEDSL=1 to enable the CuTe fastpath.")
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _load_fa4_cute():
|
||||
"""Lazily import the optional FA4 CuTe block-sparse symbols.
|
||||
|
||||
@@ -39,39 +38,14 @@ def _load_fa4_cute():
|
||||
"""
|
||||
try:
|
||||
from flash_attn.cute.block_sparsity import BlockSparseTensorsTorch
|
||||
from flash_attn.cute.interface import (
|
||||
_flash_attn_bwd,
|
||||
_flash_attn_fwd,
|
||||
flash_attn_func,
|
||||
)
|
||||
from flash_attn.cute.interface import _flash_attn_fwd
|
||||
except ImportError as exc: # pragma: no cover - optional dependency
|
||||
raise ImportError(_FA4_IMPORT_HINT) from exc
|
||||
return BlockSparseTensorsTorch, flash_attn_func, _flash_attn_fwd, _flash_attn_bwd
|
||||
return BlockSparseTensorsTorch, _flash_attn_fwd
|
||||
|
||||
|
||||
# FA4's physical Q tile size; KV block size comes from the VSA caller.
|
||||
_FA4_Q_BLOCK_SIZE = 128
|
||||
|
||||
|
||||
class _SingleQStageLength(int):
|
||||
"""Keep the real length while selecting FA4's one-stage Q128 path.
|
||||
|
||||
On sm_100 FA4 derives ``q_stage`` from ``max_seqlen_q > tile_m``. Its
|
||||
kernel supports one 128-token Q stage, but the fixed-length public wrapper
|
||||
does not expose that choice. VSA-128 must select it explicitly; otherwise
|
||||
adjacent logical Q blocks are merged into a 256-token sparse block.
|
||||
"""
|
||||
|
||||
def __mul__(self, other):
|
||||
return type(self)(int(self) * int(other))
|
||||
|
||||
def __rmul__(self, other):
|
||||
return type(self)(int(other) * int(self))
|
||||
|
||||
def __gt__(self, other):
|
||||
if int(other) == _FA4_Q_BLOCK_SIZE:
|
||||
return False
|
||||
return int(self) > int(other)
|
||||
# Q-side tile size; kv_block_size comes from the caller's VSA logical KV block.
|
||||
_M_BLOCK_SIZE_DEFAULT = 128
|
||||
|
||||
|
||||
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
@@ -90,12 +64,12 @@ def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
return triton_map_to_index(block_map)
|
||||
|
||||
|
||||
def _choose_q_sparse_block_size(q_len: int, q_tile_size: int = _FA4_Q_BLOCK_SIZE) -> int:
|
||||
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > q_tile_size.
|
||||
def _choose_q_sparse_block_size(q_len: int, m_block_size: int = _M_BLOCK_SIZE_DEFAULT) -> int:
|
||||
# FA4 supports a doubled Q sparsity granularity on sm_100+ when q_len > m_block_size.
|
||||
major, _ = torch.cuda.get_device_capability()
|
||||
if major >= 10 and q_len > q_tile_size:
|
||||
return 2 * q_tile_size
|
||||
return q_tile_size
|
||||
if major >= 10 and q_len > m_block_size:
|
||||
return 2 * m_block_size
|
||||
return m_block_size
|
||||
|
||||
|
||||
def _aggregate_q_block_map(
|
||||
@@ -160,35 +134,23 @@ def _build_vbs_mask_mod(kv_block_size: int):
|
||||
return _vbs_mask_mod
|
||||
|
||||
|
||||
def _build_sparse_tensors(
|
||||
def _cute_forward(
|
||||
q_bshd: torch.Tensor,
|
||||
k_bshd: torch.Tensor,
|
||||
v_bshd: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
*,
|
||||
q_len: int,
|
||||
q_block_size: int,
|
||||
kv_block_size: int,
|
||||
need_backward: bool,
|
||||
force_q_sparse_block_size: int | None = None,
|
||||
) -> Tuple[object, object | None]:
|
||||
"""Build the Q-owned forward and KV-owned backward sparse metadata.
|
||||
|
||||
``need_backward`` is False on inference-only calls: the backward metadata
|
||||
is a pair of dense ``[B, H, kv_blocks, q_blocks]`` int32 index tensors that
|
||||
FA4 keeps alive on its autograd ctx until backward runs, so building it
|
||||
when nothing requires grad is pure overhead (~80 MiB per call at Wan-14B
|
||||
720p shape).
|
||||
"""
|
||||
BlockSparseTensorsTorch, _, _, _ = _load_fa4_cute()
|
||||
if force_q_sparse_block_size is None:
|
||||
q_sparse_candidate = _choose_q_sparse_block_size(q_len)
|
||||
q_sparse_block_size = max(
|
||||
q_block_size,
|
||||
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
|
||||
)
|
||||
else:
|
||||
q_sparse_block_size = force_q_sparse_block_size
|
||||
if q_sparse_block_size < q_block_size or q_sparse_block_size % q_block_size != 0:
|
||||
raise ValueError("force_q_sparse_block_size must be a positive multiple of q_block_size")
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Internal: FA4 CuTe BSA fwd with BSHD inputs."""
|
||||
BlockSparseTensorsTorch, _flash_attn_fwd = _load_fa4_cute()
|
||||
q_sparse_candidate = _choose_q_sparse_block_size(q_bshd.shape[1])
|
||||
q_sparse_block_size = max(
|
||||
q_block_size,
|
||||
((q_sparse_candidate + q_block_size - 1) // q_block_size) * q_block_size,
|
||||
)
|
||||
sparse_map = _aggregate_q_block_map(
|
||||
block_map,
|
||||
q_sparse_block_size=q_sparse_block_size,
|
||||
@@ -196,166 +158,35 @@ def _build_sparse_tensors(
|
||||
)
|
||||
kv_full = (variable_block_sizes == kv_block_size).view(1, 1, 1, -1)
|
||||
kv_partial = ((variable_block_sizes > 0) & (variable_block_sizes < kv_block_size)).view(1, 1, 1, -1)
|
||||
full_map = sparse_map & kv_full
|
||||
mask_map = sparse_map & kv_partial
|
||||
|
||||
def from_maps(full_map: torch.Tensor, mask_map: torch.Tensor) -> object:
|
||||
full_block_idx, full_block_cnt = _map_to_index(full_map.contiguous())
|
||||
mask_block_idx, mask_block_cnt = _map_to_index(mask_map.contiguous())
|
||||
return BlockSparseTensorsTorch(
|
||||
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
|
||||
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
|
||||
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
|
||||
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
|
||||
block_size=(q_sparse_block_size, kv_block_size),
|
||||
)
|
||||
full_block_idx, full_block_cnt = _map_to_index(full_map)
|
||||
mask_block_idx, mask_block_cnt = _map_to_index(mask_map)
|
||||
|
||||
forward_sparse_tensors = from_maps(
|
||||
sparse_map & kv_full,
|
||||
sparse_map & kv_partial,
|
||||
sparse_tensors = BlockSparseTensorsTorch(
|
||||
full_block_cnt=full_block_cnt.to(torch.int32).contiguous(),
|
||||
full_block_idx=full_block_idx.to(torch.int32).contiguous(),
|
||||
mask_block_cnt=mask_block_cnt.to(torch.int32).contiguous(),
|
||||
mask_block_idx=mask_block_idx.to(torch.int32).contiguous(),
|
||||
block_size=(q_sparse_block_size, kv_block_size),
|
||||
)
|
||||
|
||||
if not need_backward:
|
||||
return forward_sparse_tensors, None
|
||||
|
||||
# FA4 backward is KV-owned: for each physical KV tile, list the sparse
|
||||
# query tiles that selected it. Full and partial KV tiles stay separate
|
||||
# so the token-level validity mask only runs for padded tiles.
|
||||
backward_sparse_tensors = from_maps(
|
||||
(sparse_map & kv_full).transpose(2, 3),
|
||||
(sparse_map & kv_partial).transpose(2, 3),
|
||||
)
|
||||
return forward_sparse_tensors, backward_sparse_tensors
|
||||
|
||||
|
||||
def _cute_attention_q128_forward(
|
||||
q_bshd: torch.Tensor,
|
||||
k_bshd: torch.Tensor,
|
||||
v_bshd: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
*,
|
||||
need_backward: bool,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, object | None]:
|
||||
"""Run FA4 with one physical Q stage per logical VSA-128 block."""
|
||||
_, _, flash_attn_fwd, _ = _load_fa4_cute()
|
||||
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
q_len=q_bshd.shape[1],
|
||||
q_block_size=_FA4_Q_BLOCK_SIZE,
|
||||
kv_block_size=_FA4_Q_BLOCK_SIZE,
|
||||
need_backward=need_backward,
|
||||
force_q_sparse_block_size=_FA4_Q_BLOCK_SIZE,
|
||||
)
|
||||
out, lse = flash_attn_fwd(
|
||||
# _flash_attn_fwd returns (out, lse, p, row_max); keep the first two.
|
||||
out, lse = _flash_attn_fwd(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
tile_mn=(_FA4_Q_BLOCK_SIZE, _FA4_Q_BLOCK_SIZE),
|
||||
max_seqlen_q=_SingleQStageLength(q_bshd.shape[1]),
|
||||
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
|
||||
block_sparse_tensors=forward_sparse_tensors,
|
||||
tile_mn=(_M_BLOCK_SIZE_DEFAULT, kv_block_size),
|
||||
mask_mod=_build_vbs_mask_mod(kv_block_size),
|
||||
block_sparse_tensors=sparse_tensors,
|
||||
aux_tensors=[variable_block_sizes],
|
||||
causal=False,
|
||||
return_lse=True,
|
||||
)[:2]
|
||||
return out, lse, backward_sparse_tensors
|
||||
|
||||
|
||||
class _CuteAttentionQ128(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes):
|
||||
out, lse, backward_sparse_tensors = _cute_attention_q128_forward(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
need_backward=True,
|
||||
)
|
||||
ctx.save_for_backward(q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes)
|
||||
ctx.backward_sparse_tensors = backward_sparse_tensors
|
||||
ctx.mark_non_differentiable(lse)
|
||||
ctx.set_materialize_grads(False)
|
||||
return out, lse
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_out, grad_lse):
|
||||
del grad_lse
|
||||
q_bshd, k_bshd, v_bshd, out, lse, variable_block_sizes = ctx.saved_tensors
|
||||
if grad_out is None:
|
||||
grad_out = torch.zeros_like(out)
|
||||
_, _, _, flash_attn_bwd = _load_fa4_cute()
|
||||
dq, dk, dv = flash_attn_bwd(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
out,
|
||||
grad_out.contiguous(),
|
||||
lse,
|
||||
softmax_scale=q_bshd.shape[-1]**-0.5,
|
||||
mask_mod=_build_vbs_mask_mod(_FA4_Q_BLOCK_SIZE),
|
||||
aux_tensors=[variable_block_sizes],
|
||||
block_sparse_tensors=ctx.backward_sparse_tensors,
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def _cute_attention_q128(
|
||||
q_bshd: torch.Tensor,
|
||||
k_bshd: torch.Tensor,
|
||||
v_bshd: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
|
||||
if need_backward:
|
||||
return _CuteAttentionQ128.apply(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
|
||||
out, lse, _ = _cute_attention_q128_forward(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
need_backward=False,
|
||||
)
|
||||
return out, lse
|
||||
|
||||
|
||||
def _cute_attention(
|
||||
q_bshd: torch.Tensor,
|
||||
k_bshd: torch.Tensor,
|
||||
v_bshd: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Run FA4's autograd-enabled block-sparse attention with BSHD inputs."""
|
||||
_, flash_attn_func, _, _ = _load_fa4_cute()
|
||||
q_block_size = q_bshd.shape[1] // block_map.shape[2]
|
||||
kv_block_size = k_bshd.shape[1] // block_map.shape[3]
|
||||
if q_block_size == kv_block_size == _FA4_Q_BLOCK_SIZE:
|
||||
return _cute_attention_q128(q_bshd, k_bshd, v_bshd, block_map, variable_block_sizes)
|
||||
need_backward = torch.is_grad_enabled() and any(t.requires_grad for t in (q_bshd, k_bshd, v_bshd))
|
||||
forward_sparse_tensors, backward_sparse_tensors = _build_sparse_tensors(
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
q_len=q_bshd.shape[1],
|
||||
q_block_size=q_block_size,
|
||||
kv_block_size=kv_block_size,
|
||||
need_backward=need_backward,
|
||||
)
|
||||
return flash_attn_func(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
mask_mod=_build_vbs_mask_mod(kv_block_size),
|
||||
aux_tensors=[variable_block_sizes],
|
||||
block_sparse_tensors=forward_sparse_tensors,
|
||||
block_sparse_tensors_bwd=backward_sparse_tensors,
|
||||
return_lse=True,
|
||||
)
|
||||
|
||||
|
||||
def block_sparse_attn_cute_fwd(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -363,25 +194,34 @@ def block_sparse_attn_cute_fwd(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Autograd-enabled CuTe block-sparse attention for [B, H, S, D]."""
|
||||
"""CuTe forward-only block-sparse attention with [B, H, S, D] inputs."""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
q_block_size = q.shape[2] // block_map.shape[2]
|
||||
kv_block_size = k.shape[2] // block_map.shape[3]
|
||||
|
||||
q_bshd = q.transpose(1, 2).contiguous()
|
||||
k_bshd = k.transpose(1, 2).contiguous()
|
||||
v_bshd = v.transpose(1, 2).contiguous()
|
||||
out_bshd, lse = _cute_attention(
|
||||
out_bshd, lse_bshd = _cute_forward(
|
||||
q_bshd,
|
||||
k_bshd,
|
||||
v_bshd,
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
q_block_size=q_block_size,
|
||||
kv_block_size=kv_block_size,
|
||||
)
|
||||
out = out_bshd.transpose(1, 2).contiguous()
|
||||
# FA4 already returns lse as [B, H, S], matching the Triton path's aux
|
||||
# contract, so it needs no transpose. Detach before any further op: the
|
||||
# value is informational and callers never backprop through it.
|
||||
return out, lse.detach()
|
||||
if lse_bshd is None:
|
||||
lse = torch.empty(
|
||||
(q.shape[0], q.shape[1], q.shape[2]),
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
else:
|
||||
lse = lse_bshd.transpose(1, 2).contiguous()
|
||||
return out, lse
|
||||
|
||||
|
||||
def block_sparse_attn_cute_fwd_bshd(
|
||||
@@ -391,16 +231,27 @@ def block_sparse_attn_cute_fwd_bshd(
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Autograd-enabled CuTe block-sparse attention for [B, S, H, D]."""
|
||||
"""CuTe forward-only block-sparse attention with [B, S, H, D] inputs."""
|
||||
if block_map.dim() == 3:
|
||||
block_map = block_map.unsqueeze(0)
|
||||
q_block_size = q.shape[1] // block_map.shape[2]
|
||||
kv_block_size = k.shape[1] // block_map.shape[3]
|
||||
|
||||
out, lse = _cute_attention(
|
||||
out, lse_bshd = _cute_forward(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
block_map,
|
||||
variable_block_sizes,
|
||||
q_block_size=q_block_size,
|
||||
kv_block_size=kv_block_size,
|
||||
)
|
||||
# lse is [B, H, S] regardless of the q/k/v layout; see above.
|
||||
return out, lse.detach()
|
||||
if lse_bshd is None:
|
||||
lse = torch.empty(
|
||||
(q.shape[0], q.shape[2], q.shape[1]),
|
||||
dtype=torch.float32,
|
||||
device=q.device,
|
||||
)
|
||||
else:
|
||||
lse = lse_bshd.transpose(1, 2).contiguous()
|
||||
return out, lse
|
||||
|
||||
@@ -1,104 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""sm_100a (Blackwell) CUDA block-sparse VSA forward.
|
||||
|
||||
A third backend behind the same VSA op as the Triton and CuTe-DSL paths. Forward only: it
|
||||
returns ``(out, lse)`` with ``lse`` in exactly the form ``triton_block_sparse_attn_forward``
|
||||
writes -- ``max(qk * qk_scale) + log2(l)``, ``[B, H, S]`` fp32 -- so
|
||||
``block_sparse_attn_backward_triton`` runs against it unchanged.
|
||||
|
||||
The extension carries TWO instantiations of the kernel, for 64- and 128-token sparse blocks
|
||||
(tile volumes 64 and 128 in ``build_vsa_metadata``); the block size is inferred from the
|
||||
tensors and picks the op. Anything else falls back to Triton via ``is_supported``.
|
||||
"""
|
||||
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
# The pybind symbols live on fastvideo_kernel_ops, NOT on the _C package that contains it.
|
||||
# `import fastvideo_kernel._C as _C` resolves to the namespace package, whose __init__ is
|
||||
# empty, so hasattr() fails on a wheel install and the caller silently falls back with the
|
||||
# kernel built and present.
|
||||
from fastvideo_kernel._C import fastvideo_kernel_ops as _C
|
||||
_FWD_BY_BLOCK = {
|
||||
64: getattr(_C, "block_sparse_sm100a_fwd", None),
|
||||
128: getattr(_C, "block_sparse_sm100a_blk128_fwd", None),
|
||||
}
|
||||
_HAS_VSA_SM100A = any(_FWD_BY_BLOCK.values())
|
||||
except ImportError: # pragma: no cover - extension not built
|
||||
_C = None
|
||||
_FWD_BY_BLOCK = {}
|
||||
_HAS_VSA_SM100A = False
|
||||
|
||||
_SM100 = (10, 0)
|
||||
HEAD_DIM = 128
|
||||
# Must match the -DVSA_BHSD the extension was compiled with (see CMakeLists).
|
||||
BHSD = True
|
||||
|
||||
|
||||
def _block_size(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> int:
|
||||
num_blocks = variable_block_sizes.numel()
|
||||
seqlen = q.shape[2] if BHSD else q.shape[1]
|
||||
return 0 if num_blocks == 0 or seqlen % num_blocks else seqlen // num_blocks
|
||||
|
||||
|
||||
def is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
|
||||
"""True iff this build can run these tensors; otherwise the caller uses Triton.
|
||||
|
||||
Static facts only -- shapes, dtypes, arch, layout. Deliberately NO reads of tensor
|
||||
contents: the previous ``int(variable_block_sizes.min())`` was a GPU->CPU sync on every
|
||||
call, and the kernel no longer needs it (see below). This predicate must stay cheap
|
||||
enough to sit on a per-layer dispatch path.
|
||||
|
||||
What the kernel accepts (and is tested to handle):
|
||||
* q/k/v: contiguous 4-D bf16 CUDA tensors on an sm_100 device, head_dim 128, laid out
|
||||
as compiled (BHSD here); seqlen == num_blocks * block with an EVEN num_blocks (a CTA
|
||||
owns an adjacent pair of query blocks) and a 64- or 128-token build present.
|
||||
* q2k_num: any per-row counts in [0, max_kv], NON-uniform across rows included. Rows
|
||||
with count 0 produce exactly-zero output rows (and a finite LSE sentinel) rather
|
||||
than attending anywhere -- so no ``.min()`` floor is required of the caller.
|
||||
* q2k_idx: rows only need valid entries (in [0, num_blocks)) BELOW that row's count;
|
||||
padding past the count (e.g. map_to_index's -1 fill) is never dereferenced. max_kv
|
||||
(= q2k_idx.shape[-1]) must be >= 1, which the host launcher re-checks.
|
||||
* variable_block_sizes: per-KV-block valid-token counts in [0, block]; keys at or past
|
||||
a block's count are masked. Integer metadata is converted to int32/contiguous by
|
||||
``block_sparse_attn_sm100a`` itself, so int64 inputs merely cost a cast.
|
||||
"""
|
||||
if not _HAS_VSA_SM100A or not q.is_cuda:
|
||||
return False
|
||||
if torch.cuda.get_device_capability(q.device) != _SM100:
|
||||
return False
|
||||
if q.dtype != torch.bfloat16 or q.dim() != 4 or q.shape[-1] != HEAD_DIM:
|
||||
return False
|
||||
if not q.is_contiguous():
|
||||
return False
|
||||
if _FWD_BY_BLOCK.get(_block_size(q, variable_block_sizes)) is None:
|
||||
return False
|
||||
# A CTA owns an adjacent pair of query blocks.
|
||||
if variable_block_sizes.numel() % 2 != 0:
|
||||
return False
|
||||
# Metadata must be integer-typed so the wrapper's int32 conversion is value-preserving.
|
||||
if not variable_block_sizes.is_cuda or variable_block_sizes.dtype not in (torch.int32, torch.int64):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def block_sparse_attn_sm100a(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
q2k_idx: torch.Tensor,
|
||||
q2k_num: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
need_lse: bool = True,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Forward pass. Returns ``(out, lse)``; ``out`` has q's layout."""
|
||||
fwd = _FWD_BY_BLOCK[_block_size(q, variable_block_sizes)]
|
||||
idx = q2k_idx.to(torch.int32).contiguous()
|
||||
num = q2k_num.to(torch.int32).contiguous()
|
||||
vbs = variable_block_sizes.to(torch.int32).contiguous()
|
||||
sm_scale = 1.0 / (q.shape[-1]**0.5)
|
||||
res = fwd(q.contiguous(), k.contiguous(), v.contiguous(), None,
|
||||
idx, num, vbs, sm_scale, need_lse)
|
||||
return (res[0], res[1]) if need_lse else (res[0], None)
|
||||
@@ -2,8 +2,6 @@ import math
|
||||
import torch
|
||||
from .block_sparse_attn import block_sparse_attn
|
||||
from .block_sparse_attn_256 import (
|
||||
block_sparse_attn_128,
|
||||
block_sparse_attn_128_bshd,
|
||||
block_sparse_attn_256,
|
||||
block_sparse_attn_256_bshd,
|
||||
)
|
||||
@@ -76,13 +74,12 @@ def video_sparse_attn(
|
||||
|
||||
Dispatches the sparse branch by ``block_elements = prod(block_size)``:
|
||||
- 64 -> existing TK/Triton path (see ``block_sparse_attn_from_indices``).
|
||||
- 128 -> Triton fallback or CuTe FA4 block-sparse attention.
|
||||
- 256 -> CuTe FA4 block-sparse attention (see ``block_sparse_attn_256``).
|
||||
|
||||
Backend overrides:
|
||||
- ``FASTVIDEO_VSA_TRITON=1`` forces Triton in either path.
|
||||
- ``FASTVIDEO_VSA_TK=1`` prefers the sm_90 TK kernel in the 64-block path.
|
||||
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 128/256-block paths.
|
||||
- ``FASTVIDEO_VSA_CUTEDSL=1`` prefers CuTe in the 256-block path.
|
||||
"""
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
@@ -122,9 +119,8 @@ def video_sparse_attn(
|
||||
# Sparse branch (fused Triton topk mask)
|
||||
mask = fused_topk_mask(scores, topk)
|
||||
|
||||
if block_elements in (128, 256):
|
||||
attention = block_sparse_attn_128 if block_elements == 128 else block_sparse_attn_256
|
||||
out_s = attention(q, k, v, mask, variable_block_sizes)[0]
|
||||
if block_elements == 256:
|
||||
out_s = block_sparse_attn_256(q, k, v, mask, variable_block_sizes)[0]
|
||||
else:
|
||||
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
|
||||
|
||||
@@ -146,14 +142,14 @@ def video_sparse_attn_bshd(
|
||||
"""VSA entrypoint for [B, S, H, D] tensors.
|
||||
|
||||
Avoids the BHSD<->BSHD round-trip that ``video_sparse_attn`` performs on
|
||||
the CuTe 128/256-block paths; the 64-block path still expects BHSD and is not
|
||||
the CuTe 256-block path; the 64-block path still expects BHSD and is not
|
||||
supported here.
|
||||
"""
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
if block_elements not in (128, 256):
|
||||
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=128 or 256 "
|
||||
if block_elements != 256:
|
||||
raise ValueError("video_sparse_attn_bshd is only defined for block_elements=256 "
|
||||
f"(got {block_elements}); use video_sparse_attn for the 64-block path.")
|
||||
|
||||
batch, q_seq_len, heads, dim = q.shape
|
||||
@@ -175,15 +171,19 @@ def video_sparse_attn_bshd(
|
||||
raise ValueError(f"q_variable_block_sizes must have length q_num_blocks={q_num_blocks}, "
|
||||
f"got {q_variable_block_sizes.numel()}")
|
||||
|
||||
# Compression branch (BSHD-native: match fused_block_mean's semantics).
|
||||
# Padding values are expected to be zero; gradients are broadcast across
|
||||
# the full padded block, just like the BHSD fused common path.
|
||||
# Compression branch (BSHD-native: mean over the 256-token axis).
|
||||
token_idx = torch.arange(block_elements, device=q.device, dtype=torch.int32)
|
||||
q_token_valid = (token_idx.view(1, -1) < q_variable_block_sizes.view(-1,
|
||||
1)).view(1, q_num_blocks, block_elements, 1, 1)
|
||||
kv_token_valid = (token_idx.view(1, -1) < variable_block_sizes.view(-1,
|
||||
1)).view(1, kv_num_blocks, block_elements, 1, 1)
|
||||
|
||||
q_c = q.view(batch, q_num_blocks, block_elements, heads, dim)
|
||||
k_c = k.view(batch, kv_num_blocks, block_elements, heads, dim)
|
||||
v_c = v.view(batch, kv_num_blocks, block_elements, heads, dim)
|
||||
q_c = (q_c.float().sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
|
||||
k_c = (k_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
|
||||
v_c = (v_c.float().sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
|
||||
q_c = ((q_c.float() * q_token_valid).sum(dim=2) / q_variable_block_sizes.view(1, -1, 1, 1)).to(q.dtype)
|
||||
k_c = ((k_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(k.dtype)
|
||||
v_c = ((v_c.float() * kv_token_valid).sum(dim=2) / variable_block_sizes.view(1, -1, 1, 1)).to(v.dtype)
|
||||
q_ch = q_c.permute(0, 2, 1, 3).contiguous()
|
||||
k_ch = k_c.permute(0, 2, 1, 3).contiguous()
|
||||
v_ch = v_c.permute(0, 2, 1, 3).contiguous()
|
||||
@@ -195,15 +195,13 @@ def video_sparse_attn_bshd(
|
||||
|
||||
# Sparse branch (fused Triton topk mask + CuTe BSHD).
|
||||
mask = fused_topk_mask(scores, topk)
|
||||
attention = block_sparse_attn_128_bshd if block_elements == 128 else block_sparse_attn_256_bshd
|
||||
out_s, _ = attention(q, k, v, mask, variable_block_sizes)
|
||||
out_s, _ = block_sparse_attn_256_bshd(q, k, v, mask, variable_block_sizes)
|
||||
|
||||
# Out-of-place: ``out_s`` is the tensor FA4's autograd node saved for its
|
||||
# backward, so mutating it in place invalidates the graph.
|
||||
out_view = out_s.view(batch, q_num_blocks, block_elements, heads, dim)
|
||||
out = out_s
|
||||
out_view = out.view(batch, q_num_blocks, block_elements, heads, dim)
|
||||
if compress_attn_weight is not None:
|
||||
gate_view = compress_attn_weight.view(batch, q_num_blocks, block_elements, heads, dim)
|
||||
out = out_view + out_c_blk.unsqueeze(2) * gate_view
|
||||
out_view.add_(out_c_blk.unsqueeze(2) * gate_view)
|
||||
else:
|
||||
out = out_view + out_c_blk.unsqueeze(2)
|
||||
return out.view(batch, q_seq_len, heads, dim)
|
||||
out_view.add_(out_c_blk.unsqueeze(2))
|
||||
return out
|
||||
|
||||
+8
-19
@@ -237,12 +237,7 @@ def _attn_bwd_dkdv(
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
# Recompute logits exactly as the forward does: raw bf16 operands into
|
||||
# the dot, fp32 scale after accumulation. A bf16 pre-scaled K perturbs
|
||||
# the recomputed logits relative to the saved M by an error
|
||||
# proportional to |logit|, which exp2 amplifies into arbitrarily wrong
|
||||
# probabilities at large activations.
|
||||
qkT = tl.dot(k, qT) * (sm_scale * 1.4426950408889634)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
@@ -273,7 +268,6 @@ def _attn_bwd_dq(
|
||||
do,
|
||||
m,
|
||||
D,
|
||||
sm_scale,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
@@ -321,7 +315,7 @@ def _attn_bwd_dq(
|
||||
block_sparse_offset = (kv_idx * 2 + half) * step_n * stride_tok
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT) * (sm_scale * 1.4426950408889634)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
offs_in_block = half * step_n + tl.arange(0, BLOCK_N2)
|
||||
mask = offs_in_block < block_size
|
||||
@@ -330,7 +324,8 @@ def _attn_bwd_dq(
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ (kT is raw; the caller applies sm_scale once at the end).
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
@@ -458,7 +453,6 @@ def _attn_bwd(
|
||||
do,
|
||||
m,
|
||||
D, #
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
@@ -476,7 +470,7 @@ def _attn_bwd(
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= sm_scale
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
@@ -597,7 +591,6 @@ def _attn_bwd_dq_kernel(
|
||||
Q,
|
||||
K,
|
||||
V,
|
||||
sm_scale,
|
||||
DO, #
|
||||
DQ,
|
||||
M,
|
||||
@@ -670,7 +663,6 @@ def _attn_bwd_dq_kernel(
|
||||
do,
|
||||
m,
|
||||
D,
|
||||
sm_scale,
|
||||
q2k_index,
|
||||
q2k_num,
|
||||
max_kv_blks,
|
||||
@@ -688,7 +680,7 @@ def _attn_bwd_dq_kernel(
|
||||
)
|
||||
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq_acc *= sm_scale
|
||||
dq_acc *= LN2
|
||||
tl.store(dq_ptrs, dq_acc)
|
||||
|
||||
|
||||
@@ -756,11 +748,9 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD = q.shape[:2]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
# K stays raw: the backward kernels apply sm_scale in fp32 after the dot,
|
||||
# matching the forward's rounding exactly. (A bf16 pre-scaled K perturbs
|
||||
# the recomputed logits vs the saved M; exp2 turns that into unboundedly
|
||||
# wrong probabilities at large activations.)
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert Tq % PRE_BLOCK == 0
|
||||
pre_grid = (Tq // PRE_BLOCK, BATCH * N_HEAD)
|
||||
@@ -823,7 +813,6 @@ def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q
|
||||
q,
|
||||
arg_k,
|
||||
v,
|
||||
sm_scale,
|
||||
do,
|
||||
dq,
|
||||
M,
|
||||
|
||||
@@ -14,9 +14,7 @@ import math
|
||||
import torch
|
||||
|
||||
VSA_TILE_SIZE = (4, 4, 4)
|
||||
# 128 is served by the sm_100a CUDA backend (block_sparse_attn_sm100a); 64 and 256 by
|
||||
# Triton and the CuTe-DSL path. A volume here only needs a backend that accepts it.
|
||||
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 128, 256)
|
||||
_SUPPORTED_VSA_BLOCK_VOLUMES = (64, 256)
|
||||
|
||||
|
||||
def _canonicalize_device(device: torch.device | str) -> torch.device:
|
||||
|
||||
@@ -1,146 +0,0 @@
|
||||
"""VSA-128 CuTe/Triton forward and backward parity on Blackwell."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
|
||||
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_128
|
||||
|
||||
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
|
||||
|
||||
_BLOCK = 128
|
||||
_BLOCK_SIZE_3D = (2, 8, 8)
|
||||
|
||||
|
||||
def _select_backend(monkeypatch, backend: str) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required")
|
||||
if backend == "cute":
|
||||
pytest.importorskip(
|
||||
"flash_attn.cute.block_sparsity",
|
||||
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
|
||||
)
|
||||
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
|
||||
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
|
||||
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("FASTVIDEO_VSA_TRITON", "1")
|
||||
monkeypatch.delenv("FASTVIDEO_VSA_CUTEDSL", raising=False)
|
||||
|
||||
|
||||
def _dense_sparse_reference(q, k, v, block_map, variable_block_sizes):
|
||||
token_mask = block_map.repeat_interleave(_BLOCK, dim=2).repeat_interleave(_BLOCK, dim=3)
|
||||
kv_valid = torch.arange(_BLOCK, device=k.device) < variable_block_sizes[:, None]
|
||||
token_mask = token_mask & kv_valid.reshape(1, 1, 1, -1)
|
||||
logits = torch.matmul(q.float(), k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
|
||||
probabilities = torch.softmax(logits.masked_fill(~token_mask, float("-inf")), dim=-1)
|
||||
return torch.matmul(probabilities, v.float()).to(q.dtype)
|
||||
|
||||
|
||||
def _check(tag: str, expected: torch.Tensor, actual: torch.Tensor, avg_tol: float, rel_tol: float) -> None:
|
||||
assert torch.isfinite(actual).all().item(), f"{tag}: non-finite values"
|
||||
avg_abs, max_rel = _metrics(expected, actual)
|
||||
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
|
||||
assert avg_abs < avg_tol
|
||||
assert max_rel < rel_tol
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("backend", ["cute", "triton"])
|
||||
def test_vsa128_explicit_routes_forward_backward(backend: str, monkeypatch) -> None:
|
||||
"""Adjacent Q128 blocks must keep independent routes instead of merging."""
|
||||
_select_backend(monkeypatch, backend)
|
||||
torch.manual_seed(53)
|
||||
shape = (1, 1, 3 * _BLOCK, 128)
|
||||
base = [torch.randn(shape, device="cuda", dtype=torch.bfloat16) for _ in range(3)]
|
||||
grad_output = torch.randn(shape, device="cuda", dtype=torch.bfloat16)
|
||||
variable_block_sizes = torch.tensor([128, 91, 37], device="cuda", dtype=torch.int32)
|
||||
block_map = torch.eye(3, device="cuda", dtype=torch.bool).view(1, 1, 3, 3)
|
||||
|
||||
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
|
||||
actual, _ = block_sparse_attn_128(*actual_inputs, block_map, variable_block_sizes)
|
||||
(actual * grad_output).sum().backward()
|
||||
|
||||
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in base]
|
||||
expected = _dense_sparse_reference(*reference_inputs, block_map, variable_block_sizes)
|
||||
(expected * grad_output).sum().backward()
|
||||
|
||||
print(f"[vsa128-explicit-{backend}]")
|
||||
_check("out", expected, actual, 1e-3, 0.2)
|
||||
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_inputs, strict=True):
|
||||
_check(name, reference.grad, candidate.grad, 2e-2, 0.5)
|
||||
|
||||
|
||||
def _zero_kv_tail(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
|
||||
valid = torch.arange(_BLOCK, device=x.device) < variable_block_sizes[:, None]
|
||||
valid = valid.view(1, 1, -1, _BLOCK, 1).expand_as(x.view(1, x.shape[1], -1, _BLOCK, x.shape[-1]))
|
||||
return x * valid.reshape_as(x).to(x.dtype)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("backend", ["cute", "triton"])
|
||||
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
|
||||
def test_vsa128_wrapper_forward_backward(backend: str, layout: str, monkeypatch) -> None:
|
||||
_select_backend(monkeypatch, backend)
|
||||
torch.manual_seed(59)
|
||||
batch, heads, dim = 1, 2, 128
|
||||
q_blocks, kv_blocks, topk = 3, 4, 2
|
||||
q_shape = (batch, heads, q_blocks * _BLOCK, dim)
|
||||
kv_shape = (batch, heads, kv_blocks * _BLOCK, dim)
|
||||
q_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
|
||||
kv_sizes = torch.tensor([128, 91, 37, 128], device="cuda", dtype=torch.int32)
|
||||
q_sizes = torch.full((q_blocks, ), _BLOCK, device="cuda", dtype=torch.int32)
|
||||
k_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
|
||||
v_base = _zero_kv_tail(torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16), kv_sizes)
|
||||
gate_base = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16) * 0.1
|
||||
grad_output = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
if layout == "bhsd":
|
||||
actual_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
|
||||
actual_gate = gate_base.detach().clone().requires_grad_()
|
||||
actual = video_sparse_attn(
|
||||
*actual_inputs,
|
||||
kv_sizes,
|
||||
q_sizes,
|
||||
topk,
|
||||
block_size=_BLOCK_SIZE_3D,
|
||||
compress_attn_weight=actual_gate,
|
||||
)
|
||||
actual_grads = actual_inputs
|
||||
else:
|
||||
bshd_inputs = [tensor.transpose(1, 2).contiguous().detach().requires_grad_()
|
||||
for tensor in (q_base, k_base, v_base)]
|
||||
bshd_gate = gate_base.transpose(1, 2).contiguous().detach().requires_grad_()
|
||||
actual = video_sparse_attn_bshd(
|
||||
*bshd_inputs,
|
||||
kv_sizes,
|
||||
q_sizes,
|
||||
topk,
|
||||
block_size=_BLOCK_SIZE_3D,
|
||||
compress_attn_weight=bshd_gate,
|
||||
).transpose(1, 2)
|
||||
actual_grads = bshd_inputs
|
||||
(actual * grad_output).sum().backward()
|
||||
|
||||
reference_inputs = [tensor.detach().clone().requires_grad_() for tensor in (q_base, k_base, v_base)]
|
||||
reference_gate = gate_base.detach().clone().requires_grad_()
|
||||
expected = _torch_vsa256_reference(
|
||||
*reference_inputs,
|
||||
q_sizes,
|
||||
kv_sizes,
|
||||
topk,
|
||||
compress_attn_weight=reference_gate,
|
||||
)
|
||||
(expected * grad_output).sum().backward()
|
||||
|
||||
print(f"[vsa128-wrapper-{backend}-{layout}]")
|
||||
_check("out", expected, actual, 1e-3, 0.2)
|
||||
for name, reference, candidate in zip(("dq", "dk", "dv"), reference_inputs, actual_grads, strict=True):
|
||||
candidate_grad = candidate.grad if layout == "bhsd" else candidate.grad.transpose(1, 2)
|
||||
_check(name, reference.grad, candidate_grad, 2e-2, 0.5)
|
||||
actual_gate_grad = actual_gate.grad if layout == "bhsd" else bshd_gate.grad.transpose(1, 2)
|
||||
_check("dgate", reference_gate.grad, actual_gate_grad, 1e-3, 0.2)
|
||||
@@ -1,224 +0,0 @@
|
||||
"""VSA-256 FA4 CuTe forward/backward parity for BHSD and BSHD APIs.
|
||||
|
||||
Covers the shapes the CuTe backward actually sees in production: the gated
|
||||
compression branch (`compress_attn_weight`), partially filled Q tiles,
|
||||
and q_len != kv_len. Also pins the inference fast path, which must skip the
|
||||
KV-owned backward metadata without changing the forward result.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Tuple
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel import video_sparse_attn, video_sparse_attn_bshd
|
||||
|
||||
from .test_vsa256_triton import _metrics, _torch_vsa256_reference
|
||||
|
||||
_BLOCK = 256
|
||||
_BLOCK_SIZE_3D = (4, 8, 8) # prod == 256
|
||||
|
||||
# Measured on GB200 (sm_100) with bf16 inputs: grads land around 1e-4 avg_abs
|
||||
# and <=0.11 max_rel across every case below, so these leave ~10x headroom
|
||||
# without being loose enough to hide a real regression.
|
||||
_OUT_TOL = (1e-3, 0.2)
|
||||
_GRAD_TOL = (1e-3, 0.25)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _require_cute_backend(monkeypatch):
|
||||
pytest.importorskip(
|
||||
"flash_attn.cute.block_sparsity",
|
||||
reason="optional FA4 CuTe build (flash_attn.cute) not installed",
|
||||
)
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required")
|
||||
monkeypatch.setenv("FASTVIDEO_VSA_CUTEDSL", "1")
|
||||
monkeypatch.delenv("FASTVIDEO_VSA_TRITON", raising=False)
|
||||
monkeypatch.delenv("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", raising=False)
|
||||
|
||||
|
||||
def _zero_pad_tail(x: torch.Tensor, var: torch.Tensor) -> torch.Tensor:
|
||||
"""Zero the padded tail of every 256-token tile of a [B, H, S, D] tensor.
|
||||
|
||||
VSA callers scatter into a zeroed tile buffer, so padded slots are zero;
|
||||
both the kernel and the reference rely on that.
|
||||
"""
|
||||
bsz, heads, _, dim = x.shape
|
||||
blocks = var.numel()
|
||||
token_idx = torch.arange(_BLOCK, device=x.device, dtype=torch.int32)
|
||||
valid = (token_idx.view(1, -1) < var.view(-1, 1)).view(1, 1, blocks, _BLOCK, 1)
|
||||
valid = valid.expand(bsz, heads, blocks, _BLOCK, dim).reshape_as(x)
|
||||
return x * valid.to(x.dtype)
|
||||
|
||||
|
||||
def _make_inputs(
|
||||
q_blocks: int,
|
||||
kv_blocks: int,
|
||||
kv_var: torch.Tensor,
|
||||
q_var: torch.Tensor,
|
||||
heads: int = 2,
|
||||
dim: int = 128,
|
||||
seed: int = 42,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
torch.manual_seed(seed)
|
||||
device = torch.device("cuda")
|
||||
dtype = torch.bfloat16
|
||||
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
|
||||
q = torch.randn(1, heads, sq, dim, device=device, dtype=dtype)
|
||||
k = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
|
||||
v = torch.randn(1, heads, skv, dim, device=device, dtype=dtype)
|
||||
grad_out = torch.randn_like(q)
|
||||
return _zero_pad_tail(q, q_var), _zero_pad_tail(k, kv_var), _zero_pad_tail(v, kv_var), grad_out
|
||||
|
||||
|
||||
def _check(tag: str, ref: torch.Tensor, got: torch.Tensor, tol: Tuple[float, float]) -> None:
|
||||
assert torch.isfinite(got).all().item(), f"{tag}: non-finite values"
|
||||
avg_abs, max_rel = _metrics(ref, got)
|
||||
print(f" {tag}: avg_abs={avg_abs:.6e}, max_rel={max_rel:.6e}")
|
||||
assert avg_abs < tol[0], f"{tag}: avg_abs {avg_abs:.3e} >= {tol[0]:.3e}"
|
||||
assert max_rel < tol[1], f"{tag}: max_rel {max_rel:.3e} >= {tol[1]:.3e}"
|
||||
|
||||
|
||||
def _run_bhsd(q, k, v, kv_var, q_var, topk, gate=None):
|
||||
qg, kg, vg = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
|
||||
out = video_sparse_attn(qg, kg, vg, kv_var, q_var, topk, block_size=_BLOCK_SIZE_3D, compress_attn_weight=gate)
|
||||
return out, (qg, kg, vg)
|
||||
|
||||
|
||||
def _run_bshd(q, k, v, kv_var, q_var, topk, gate=None):
|
||||
qg, kg, vg = (t.transpose(1, 2).contiguous().requires_grad_(True) for t in (q, k, v))
|
||||
gate_bshd = None if gate is None else gate.transpose(1, 2).contiguous()
|
||||
out = video_sparse_attn_bshd(qg,
|
||||
kg,
|
||||
vg,
|
||||
kv_var,
|
||||
q_var,
|
||||
topk,
|
||||
block_size=_BLOCK_SIZE_3D,
|
||||
compress_attn_weight=gate_bshd)
|
||||
return out.transpose(1, 2), (qg, kg, vg)
|
||||
|
||||
|
||||
def _reference(q, k, v, q_var, kv_var, topk, gate=None):
|
||||
qr, kr, vr = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
|
||||
out = _torch_vsa256_reference(qr, kr, vr, q_var, kv_var, topk, compress_attn_weight=gate)
|
||||
return out, (qr, kr, vr)
|
||||
|
||||
|
||||
def _compare(tag, layout, q, k, v, kv_var, q_var, topk, grad_out, gate=None):
|
||||
runner = _run_bhsd if layout == "bhsd" else _run_bshd
|
||||
out, (qg, kg, vg) = runner(q, k, v, kv_var, q_var, topk, gate=gate)
|
||||
(out * grad_out).sum().backward()
|
||||
grads = [g.grad if g.grad.dim() == 4 and layout == "bhsd" else g.grad for g in (qg, kg, vg)]
|
||||
if layout == "bshd":
|
||||
grads = [g.transpose(1, 2) for g in grads]
|
||||
|
||||
out_ref, refs = _reference(q, k, v, q_var, kv_var, topk, gate=gate)
|
||||
(out_ref * grad_out).sum().backward()
|
||||
|
||||
print(f"[{tag}-{layout}]")
|
||||
_check("out", out_ref, out, _OUT_TOL)
|
||||
for name, ref, got in zip(("dq", "dk", "dv"), refs, grads):
|
||||
_check(name, ref.grad, got, _GRAD_TOL)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
|
||||
def test_vsa256_cute_forward_backward_vs_torch_ref(layout: str) -> None:
|
||||
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
|
||||
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
|
||||
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var)
|
||||
_compare("vsa256-cute", layout, q, k, v, kv_var, q_var, 2, grad_out)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
|
||||
def test_vsa256_cute_backward_with_compress_gate(layout: str) -> None:
|
||||
"""The gated compression branch is what Wan and MiniMax-H3 actually run.
|
||||
|
||||
It is also the branch that composes the sparse output with the compression
|
||||
output, so it is the one that breaks if that composition mutates FA4's
|
||||
saved output in place.
|
||||
"""
|
||||
kv_var = torch.tensor([256, 200, 256, 91], dtype=torch.int32, device="cuda")
|
||||
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
|
||||
q, k, v, grad_out = _make_inputs(3, 4, kv_var, q_var, seed=7)
|
||||
gate = torch.randn_like(q) * 0.1
|
||||
_compare("vsa256-cute-gated", layout, q, k, v, kv_var, q_var, 2, grad_out, gate=gate)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
|
||||
def test_vsa256_cute_backward_partial_q_blocks(layout: str) -> None:
|
||||
"""Q tiles that are not full: only the compression divisor depends on it,
|
||||
but it is the one axis the existing coverage held constant."""
|
||||
kv_var = torch.tensor([256, 128, 256], dtype=torch.int32, device="cuda")
|
||||
q_var = torch.tensor([256, 61, 199], dtype=torch.int32, device="cuda")
|
||||
q, k, v, grad_out = _make_inputs(3, 3, kv_var, q_var, seed=11)
|
||||
_compare("vsa256-cute-partial-q", layout, q, k, v, kv_var, q_var, 2, grad_out)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("layout", ["bhsd", "bshd"])
|
||||
def test_vsa256_cute_backward_cross_q_kv(layout: str) -> None:
|
||||
"""q_len != kv_len: forward has coverage, backward did not."""
|
||||
kv_var = torch.tensor([256, 143, 256, 256, 88], dtype=torch.int32, device="cuda")
|
||||
q_var = torch.full((2, ), _BLOCK, dtype=torch.int32, device="cuda")
|
||||
q, k, v, grad_out = _make_inputs(2, 5, kv_var, q_var, seed=13)
|
||||
_compare("vsa256-cute-cross", layout, q, k, v, kv_var, q_var, 3, grad_out)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
def test_vsa256_cute_inference_matches_training_forward() -> None:
|
||||
"""The KV-owned backward metadata is only built when something requires
|
||||
grad. Skipping it must not perturb the forward result."""
|
||||
kv_var = torch.tensor([256, 173, 79, 256], dtype=torch.int32, device="cuda")
|
||||
q_var = torch.full((3, ), _BLOCK, dtype=torch.int32, device="cuda")
|
||||
q, k, v, _ = _make_inputs(3, 4, kv_var, q_var, seed=5)
|
||||
|
||||
with torch.no_grad():
|
||||
out_infer = video_sparse_attn_bshd(
|
||||
q.transpose(1, 2).contiguous(),
|
||||
k.transpose(1, 2).contiguous(),
|
||||
v.transpose(1, 2).contiguous(),
|
||||
kv_var,
|
||||
q_var,
|
||||
2,
|
||||
block_size=_BLOCK_SIZE_3D,
|
||||
compress_attn_weight=None,
|
||||
)
|
||||
|
||||
out_train, _ = _run_bshd(q, k, v, kv_var, q_var, 2)
|
||||
torch.testing.assert_close(out_infer, out_train.transpose(1, 2).detach(), rtol=0, atol=0)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
def test_vsa256_cute_lse_is_bhs() -> None:
|
||||
"""The aux return is [B, H, S] on both entrypoints, matching the Triton
|
||||
path's contract."""
|
||||
from fastvideo_kernel.block_sparse_attn_256 import (block_sparse_attn_256, block_sparse_attn_256_bshd)
|
||||
|
||||
device = torch.device("cuda")
|
||||
heads, dim, q_blocks, kv_blocks = 2, 128, 3, 4
|
||||
sq, skv = q_blocks * _BLOCK, kv_blocks * _BLOCK
|
||||
q = torch.randn(1, heads, sq, dim, device=device, dtype=torch.bfloat16)
|
||||
k = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
|
||||
v = torch.randn(1, heads, skv, dim, device=device, dtype=torch.bfloat16)
|
||||
vbs = torch.full((kv_blocks, ), _BLOCK, dtype=torch.int32, device=device)
|
||||
mask = torch.zeros(1, heads, q_blocks, kv_blocks, dtype=torch.bool, device=device)
|
||||
mask[..., :2] = True
|
||||
|
||||
_, lse_bhsd = block_sparse_attn_256(q, k, v, mask, vbs)
|
||||
assert lse_bhsd.shape == (1, heads, sq), lse_bhsd.shape
|
||||
|
||||
_, lse_bshd = block_sparse_attn_256_bshd(
|
||||
q.transpose(1, 2).contiguous(),
|
||||
k.transpose(1, 2).contiguous(),
|
||||
v.transpose(1, 2).contiguous(),
|
||||
mask,
|
||||
vbs,
|
||||
)
|
||||
assert lse_bshd.shape == (1, heads, sq), lse_bshd.shape
|
||||
@@ -22,7 +22,6 @@ def _torch_vsa256_reference(
|
||||
q_var: torch.Tensor,
|
||||
kv_var: torch.Tensor,
|
||||
topk_logical: int,
|
||||
compress_attn_weight: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
bsz, heads, _sq, dim = q.shape
|
||||
q_blocks = q_var.numel()
|
||||
@@ -56,8 +55,6 @@ def _torch_vsa256_reference(
|
||||
logits = logits.masked_fill(~token_mask, float("-inf"))
|
||||
prob = torch.softmax(logits, dim=-1)
|
||||
out_s = torch.matmul(prob, vf).to(q.dtype)
|
||||
if compress_attn_weight is not None:
|
||||
return out_c * compress_attn_weight + out_s
|
||||
return out_c + out_s
|
||||
|
||||
|
||||
|
||||
@@ -1,104 +0,0 @@
|
||||
"""Regression: Triton block-sparse backward gradient parity at realistic activation scale.
|
||||
|
||||
The backward used to fold ``sm_scale / ln(2)`` into K in bf16 before the
|
||||
exp2-based logit recompute. The bf16 rounding error on the pre-scaled K grows
|
||||
proportionally to |logit| and exp2 amplifies it into exponentially wrong
|
||||
probabilities, so dQ/dK/dV were correct at unit scale (every pre-existing test)
|
||||
but off by orders of magnitude at real activation magnitudes.
|
||||
|
||||
This test sweeps the input scale and checks the Triton kernel's gradients
|
||||
against an fp32 masked-dense SDPA reference. The unit-scale case is the
|
||||
control (it passed even with the broken kernel); the large-scale cases are
|
||||
the regression.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn import _map_to_index, block_sparse_attn_triton
|
||||
|
||||
from .utils import generate_block_sparse_mask_for_function
|
||||
|
||||
BLOCK = 64
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _seed_rng():
|
||||
"""Pin the RNG so these cases do not depend on what ran before them.
|
||||
|
||||
Same convention as test_vsa_varlen.py: every tensor here comes from the
|
||||
global torch RNG and the checks use tight thresholds, so an unseeded run
|
||||
would shift inputs whenever an earlier test file draws a different number
|
||||
of randoms.
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed_all(0)
|
||||
|
||||
|
||||
def _dense_reference(q, k, v, block_mask):
|
||||
"""fp32 masked-dense SDPA over the token-expanded block mask.
|
||||
|
||||
q/k/v: [B, H, S, D]; block_mask: [B, H, S // BLOCK, S // BLOCK] bool.
|
||||
"""
|
||||
qf, kf, vf = q.float(), k.float(), v.float()
|
||||
token_mask = block_mask.repeat_interleave(BLOCK, dim=-2).repeat_interleave(BLOCK, dim=-1)
|
||||
logits = torch.matmul(qf, kf.transpose(-2, -1)) * (q.shape[-1]**-0.5)
|
||||
logits = logits.masked_fill(~token_mask, float("-inf"))
|
||||
return torch.matmul(logits.softmax(dim=-1), vf)
|
||||
|
||||
|
||||
@pytest.mark.cuda
|
||||
@pytest.mark.parametrize("scale", [1.0, 4.0, 16.0])
|
||||
def test_triton_backward_grad_parity_across_input_scales(scale: float) -> None:
|
||||
"""Kernel dQ/dK/dV must stay within a few percent of the fp32 reference
|
||||
regardless of input magnitude.
|
||||
|
||||
With the bf16 K pre-scaling bug, scale<=4.0 passes at this geometry while
|
||||
scale=16.0 fails (measured on GB200: dq relative L2 error 5.9e-1 vs 6.9e-3
|
||||
fixed); at larger geometries and real activation magnitudes the broken
|
||||
kernel is off by orders of magnitude. The passing unit-scale case is
|
||||
exactly how the bug survived the original test suite.
|
||||
"""
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required")
|
||||
|
||||
device = torch.device("cuda")
|
||||
dtype = torch.bfloat16
|
||||
batch, heads, dim = 1, 4, 128
|
||||
num_blocks = 8
|
||||
seq = num_blocks * BLOCK
|
||||
|
||||
q = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
|
||||
k = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype) * scale
|
||||
v = torch.randn(batch, heads, seq, dim, device=device, dtype=dtype)
|
||||
grad_out = torch.randn_like(q)
|
||||
|
||||
block_mask = generate_block_sparse_mask_for_function(heads, num_blocks, num_blocks, k=3,
|
||||
device=device).unsqueeze(0)
|
||||
q2k_idx, q2k_num = _map_to_index(block_mask)
|
||||
variable_block_sizes = torch.full((num_blocks, ), BLOCK, dtype=torch.int32, device=device)
|
||||
|
||||
q_ker, k_ker, v_ker = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
|
||||
out_ker, _ = block_sparse_attn_triton(q_ker, k_ker, v_ker, q2k_idx, q2k_num, variable_block_sizes)
|
||||
(out_ker.float() * grad_out.float()).sum().backward()
|
||||
|
||||
q_ref, k_ref, v_ref = (t.detach().clone().requires_grad_(True) for t in (q, k, v))
|
||||
out_ref = _dense_reference(q_ref, k_ref, v_ref, block_mask)
|
||||
(out_ref * grad_out.float()).sum().backward()
|
||||
|
||||
# Forward is exact at any scale; this pins the harness itself.
|
||||
fwd_rel = ((out_ker.float() - out_ref).norm() / out_ref.norm()).item()
|
||||
assert fwd_rel < 2e-2, f"scale={scale}: forward rel err {fwd_rel:.3e}"
|
||||
|
||||
for name, g_ker, g_ref in (
|
||||
("dq", q_ker.grad, q_ref.grad),
|
||||
("dk", k_ker.grad, k_ref.grad),
|
||||
("dv", v_ker.grad, v_ref.grad),
|
||||
):
|
||||
assert torch.isfinite(g_ker).all().item(), f"scale={scale}: non-finite {name}"
|
||||
ref_norm = g_ref.float().norm()
|
||||
rel = ((g_ker.float() - g_ref.float()).norm() / ref_norm.clamp_min(1e-12)).item()
|
||||
ratio = (g_ker.float().norm() / ref_norm.clamp_min(1e-12)).item()
|
||||
print(f"scale={scale} {name}: rel_l2={rel:.4e} norm_ratio={ratio:.4f}")
|
||||
assert rel < 5e-2, f"scale={scale}: {name} rel l2 err {rel:.3e} >= 5e-2"
|
||||
assert 0.98 < ratio < 1.02, f"scale={scale}: {name} grad-norm ratio {ratio:.4f}"
|
||||
@@ -23,20 +23,6 @@ from fastvideo_kernel.block_sparse_attn import (
|
||||
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _seed_rng():
|
||||
"""Pin the RNG so these cases do not depend on what ran before them.
|
||||
|
||||
Every tensor and every variable block size here comes from the global
|
||||
torch RNG, and the gradient checks use a tight max_rel threshold. Without
|
||||
a seed the inputs shift whenever an earlier test file draws a different
|
||||
number of randoms, which surfaces as an unrelated-looking failure in
|
||||
whichever case happens to land on unlucky data.
|
||||
"""
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed_all(0)
|
||||
|
||||
|
||||
def _reference_per_sequence(
|
||||
q_list,
|
||||
k_list,
|
||||
|
||||
@@ -100,7 +100,7 @@ class ComponentConfig:
|
||||
|
||||
@dataclass
|
||||
class PipelineSelection:
|
||||
workload_type: Literal["t2v", "i2v", "t2i", "i2i", "v2a", "t2a"] | None = None
|
||||
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
|
||||
preset: str | None = None
|
||||
preset_version: int | None = None
|
||||
components: ComponentConfig = field(default_factory=ComponentConfig)
|
||||
|
||||
@@ -19,12 +19,7 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
# Every worker records the loaded FlashAttention implementation so a
|
||||
# distributed profiling log contains one backend receipt per rank.
|
||||
logger.info("Worker %s Using FlashAttention-%s backend",
|
||||
os.environ.get("RANK", "0"),
|
||||
fa_version,
|
||||
local_main_process_only=False)
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
|
||||
@@ -173,12 +168,8 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor, key_len: int) -> torch.Tensor:
|
||||
# Normalize attn_mask to [B, key_len] where True means valid token.
|
||||
if attn_mask.dim() == 4:
|
||||
if attn_mask.shape[1] != 1 or attn_mask.shape[-2] != 1:
|
||||
raise ValueError("FLASH_ATTN only supports 4D key-padding masks with shape [B, 1, 1, K]")
|
||||
attn_mask = attn_mask[:, 0, 0, :]
|
||||
elif attn_mask.dim() == 3:
|
||||
if attn_mask.shape[-2] != 1:
|
||||
raise ValueError("FLASH_ATTN only supports 3D key-padding masks with shape [B, 1, K]")
|
||||
attn_mask = attn_mask[:, 0, :]
|
||||
elif attn_mask.dim() != 2:
|
||||
raise ValueError(f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}")
|
||||
@@ -283,14 +274,6 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
if getattr(attn_metadata, "is_causal", False):
|
||||
return flash_attn_func_compilable(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=True,
|
||||
)
|
||||
|
||||
# flash_attn_no_pad packs q/k/v as one tensor and assumes equal q/k
|
||||
# sequence lengths. Cross-attention can violate this.
|
||||
@@ -319,7 +302,7 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
|
||||
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
|
||||
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=self.causal, dropout_p=0, softmax_scale=None)
|
||||
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
|
||||
elif self.nvfp4_fa4:
|
||||
output = self._forward_nvfp4(query, key, value)
|
||||
|
||||
|
||||
@@ -41,8 +41,6 @@ class SDPABackend(AttentionBackend):
|
||||
class SDPAMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
attn_mask: torch.Tensor | None = None
|
||||
# The mask is exactly native causal attention, with no additional padding.
|
||||
is_causal: bool = False
|
||||
|
||||
|
||||
class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
@@ -5,17 +5,13 @@ H3 runs one joint bidirectional attention over
|
||||
``[text | condition keyframes | audio | generated video]``, so this
|
||||
backend differs from the Wan-tuned ``video_sparse_attn``:
|
||||
|
||||
- Tiles are ``[segment-pure prefix chunks] + [3D video tiles]``; prefix
|
||||
tiles never straddle segment boundaries. The tile size is selectable at
|
||||
metadata build time: 256 tokens ``(4,8,8)`` (default) or 64 tokens
|
||||
``(4,4,4)`` (see ``VSA_H3_TILE_SHAPES``).
|
||||
- Tiles are ``[segment-pure prefix chunks] + [3D (4,8,8) video tiles]``;
|
||||
prefix tiles never straddle segment boundaries.
|
||||
- Selection is pure Python on pooled tile scores; the block-sparse kernel
|
||||
consumes an explicit bool mask, so no kernel changes are needed.
|
||||
- The compression branch is gated by ``to_gate_compress``, which the base
|
||||
H3 checkpoint does not carry: the loader zero-initializes it, so
|
||||
untrained inference is exactly pure sparse and finetuning can learn the
|
||||
gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship
|
||||
trained gates, which load and activate the branch.
|
||||
- The compression branch is gated by ``to_gate_compress``, which the H3
|
||||
checkpoint does not carry: the loader zero-initializes it, so untrained
|
||||
inference is exactly pure sparse and finetuning can learn the gate.
|
||||
- Non-video *queries* are always dense. Non-video *keys* are either
|
||||
always-selected for every query ("exempt", default) or compete in
|
||||
top-k under a FLOP-matched budget ("compete") — the ablation axis,
|
||||
@@ -24,46 +20,22 @@ backend differs from the Wan-tuned ``video_sparse_attn``:
|
||||
(``vsa_dense_first_n_steps``, ``vsa_dense_layers``) let mixed schedules
|
||||
run the diffuse steps/layers dense while pushing the rest harder.
|
||||
|
||||
At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path
|
||||
Targets sm10.x through the FA4 CuTe 256-tile path
|
||||
(``FASTVIDEO_VSA_CUTEDSL=1``); the Triton 256→64 expansion is the
|
||||
fallback and keeps identical mask semantics. At tile 64 the block map is
|
||||
already at the kernels' native 64-token granularity, so both forward and
|
||||
backward run the Triton block-sparse kernels directly (no expansion,
|
||||
``FASTVIDEO_VSA_CUTEDSL`` does not apply). A third, opt-in route exists
|
||||
for the tile-64 FORWARD only: ``FASTVIDEO_VSA_SM100A=1`` sends no-grad
|
||||
forwards through the sm_100a CUDA block-sparse kernel
|
||||
(``fastvideo_kernel.block_sparse_attn_sm100a``, upstream PR #1719 plus
|
||||
our per-q-tile ``q2k_num`` fix) when the extension is built, the device
|
||||
is sm_100, and the geometry qualifies; grad-tracking forwards and every
|
||||
backward stay on Triton unchanged. If the env is set but a precondition
|
||||
fails, the route logs one warning and falls back.
|
||||
fallback and keeps identical mask semantics.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
from fastvideo_kernel.block_sparse_attn import block_sparse_attn as block_sparse_attn_64_bhsd
|
||||
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index
|
||||
except ImportError:
|
||||
block_sparse_attn_64_bhsd = None
|
||||
block_sparse_attn_256_bshd = None
|
||||
map_to_index = None
|
||||
|
||||
try:
|
||||
# Optional: only present in fastvideo_kernel builds that carry the sm_100a
|
||||
# CUDA block-sparse forward (upstream PR #1719). The module itself imports
|
||||
# fine without the compiled symbols (`_HAS_VSA_SM100A` is then False and
|
||||
# `is_supported` says no), so this only guards *module* availability.
|
||||
from fastvideo_kernel import block_sparse_attn_sm100a as _sm100a
|
||||
except ImportError:
|
||||
_sm100a = None
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder, layer_idx_from_prefix)
|
||||
@@ -71,115 +43,51 @@ from fastvideo.attention.backends.video_sparse_attn import (compute_topk, constr
|
||||
get_non_pad_index, get_tile_partition_indices,
|
||||
scatter_into_tile_buf)
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3_probe import probe_enabled, record_probe
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Opt-in switch for the sm_100a CUDA forward on the tile-64 no-grad path.
|
||||
VSA_SM100A_ENV = "FASTVIDEO_VSA_SM100A"
|
||||
|
||||
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x (default)
|
||||
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
|
||||
_TILE_ELEMS = math.prod(VSA_H3_TILE_SIZE)
|
||||
# Selectable tile geometries, keyed by element count (= the build-time
|
||||
# ``tile_size``). 64 runs the native 64-token Triton block-sparse kernels for
|
||||
# forward AND backward — the block map is already at kernel granularity, so no
|
||||
# 256->64 mask expansion is involved and FASTVIDEO_VSA_CUTEDSL does not apply.
|
||||
VSA_H3_TILE_SHAPES: dict[int, tuple[int, int, int]] = {
|
||||
_TILE_ELEMS: VSA_H3_TILE_SIZE,
|
||||
64: (4, 4, 4),
|
||||
}
|
||||
|
||||
|
||||
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
|
||||
tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
def token_tile_and_valid(variable_block_sizes: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Per padded-token tile id and pad-validity mask.
|
||||
|
||||
The single encoding of the padding contract, shared by the probe and the
|
||||
test oracle so they cannot drift from the backend's tile geometry.
|
||||
``tile_elems`` must match the metadata the sizes came from
|
||||
(``MiniMaxH3VSAMetadata.tile_elems``).
|
||||
"""
|
||||
device = variable_block_sizes.device
|
||||
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
|
||||
token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
|
||||
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(_TILE_ELEMS)
|
||||
token_valid = (torch.arange(_TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
|
||||
return token_tile, token_valid
|
||||
|
||||
|
||||
def _validate_h3_tile_geometry(
|
||||
prefix_segments: tuple[int, ...],
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
variable_block_sizes: torch.Tensor,
|
||||
untile_combined_index: torch.Tensor,
|
||||
tile_elems: int = _TILE_ELEMS,
|
||||
) -> None:
|
||||
"""Fail synchronously on out-of-bounds tile geometry.
|
||||
|
||||
Invariants the block-sparse kernel trusts without checking:
|
||||
every tile's valid size is in (0, tile_elems]; the sizes sum to the
|
||||
packed sequence length; and ``untile_combined_index`` maps each packed
|
||||
row to exactly one non-pad slot of the padded tile buffer. A violation
|
||||
would surface only as an async device fault at some later kernel or
|
||||
collective (e.g. an FSDP all-gather), which is unattributable — so raise
|
||||
here, once per cached geometry, with the numbers in hand.
|
||||
"""
|
||||
total = sum(prefix_segments) + math.prod(dit_seq_shape)
|
||||
n_pad = variable_block_sizes.numel() * tile_elems
|
||||
sizes_min = int(variable_block_sizes.min())
|
||||
sizes_max = int(variable_block_sizes.max())
|
||||
sizes_sum = int(variable_block_sizes.sum())
|
||||
if sizes_min < 1 or sizes_max > tile_elems or sizes_sum != total:
|
||||
raise ValueError(f"VSA-H3 tile sizes out of bounds for prefix={prefix_segments}, video={dit_seq_shape}, "
|
||||
f"tile_elems={tile_elems}: min={sizes_min}, max={sizes_max}, sum={sizes_sum}, "
|
||||
f"expected sum={total}.")
|
||||
if untile_combined_index.numel() != total:
|
||||
raise ValueError(f"VSA-H3 untile index has {untile_combined_index.numel()} entries for a packed "
|
||||
f"sequence of {total} rows (prefix={prefix_segments}, video={dit_seq_shape}).")
|
||||
idx_min = int(untile_combined_index.min())
|
||||
idx_max = int(untile_combined_index.max())
|
||||
if idx_min < 0 or idx_max >= n_pad:
|
||||
# Range first: the pad-slot gather below would itself index out of
|
||||
# bounds (the very async fault this guard exists to preempt).
|
||||
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: range "
|
||||
f"[{idx_min}, {idx_max}] vs padded length {n_pad} "
|
||||
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
|
||||
in_tile_offset = untile_combined_index % tile_elems
|
||||
maps_into_pad = bool((in_tile_offset >= variable_block_sizes[untile_combined_index // tile_elems]).any())
|
||||
if maps_into_pad or int(torch.unique(untile_combined_index).numel()) != total:
|
||||
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: "
|
||||
f"pad-slot hit={maps_into_pad} "
|
||||
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def _h3_tile_geometry(
|
||||
prefix_segments: tuple[int, ...],
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
tile_shape: tuple[int, int, int] = VSA_H3_TILE_SIZE,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
|
||||
"""Tile the packed sequence: segment-pure prefix chunks, then video tiles.
|
||||
|
||||
Returns (tile_partition_indices, variable_block_sizes,
|
||||
untile_combined_index, num_prefix_tiles, num_video_tiles).
|
||||
"""
|
||||
tile_elems = math.prod(tile_shape)
|
||||
prefix_len = sum(prefix_segments)
|
||||
|
||||
prefix_sizes: list[int] = []
|
||||
for segment in prefix_segments:
|
||||
full, rem = divmod(segment, tile_elems)
|
||||
prefix_sizes.extend([tile_elems] * full)
|
||||
full, rem = divmod(segment, _TILE_ELEMS)
|
||||
prefix_sizes.extend([_TILE_ELEMS] * full)
|
||||
if rem:
|
||||
prefix_sizes.append(rem)
|
||||
num_prefix_tiles = len(prefix_sizes)
|
||||
|
||||
ts_t, ts_h, ts_w = tile_shape
|
||||
ts_t, ts_h, ts_w = VSA_H3_TILE_SIZE
|
||||
t, h, w = dit_seq_shape
|
||||
num_tiles = (math.ceil(t / ts_t), math.ceil(h / ts_h), math.ceil(w / ts_w))
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, VSA_H3_TILE_SIZE)
|
||||
num_video_tiles = int(video_sizes.numel())
|
||||
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
|
||||
tile_partition_indices = torch.cat([
|
||||
torch.arange(prefix_len, device=device, dtype=torch.long),
|
||||
video_indices,
|
||||
@@ -192,11 +100,9 @@ def _h3_tile_geometry(
|
||||
|
||||
# get_non_pad_index is lru-cached on tensor identity; variable_block_sizes
|
||||
# is itself cached by this function, so the identity stays stable.
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, tile_elems)
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, _TILE_ELEMS)
|
||||
|
||||
untile_combined_index = non_pad_index[torch.argsort(tile_partition_indices)]
|
||||
# One-time (lru-cached) synchronous bounds check; see _validate_h3_tile_geometry.
|
||||
_validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
|
||||
return (tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles)
|
||||
|
||||
|
||||
@@ -233,9 +139,6 @@ class MiniMaxH3VSAMetadata(AttentionMetadata):
|
||||
exempt: bool
|
||||
variable_block_sizes: torch.Tensor
|
||||
untile_combined_index: torch.Tensor
|
||||
# tokens per tile (256 or 64); selects the tile geometry AND the kernel
|
||||
# route in forward() (256 -> VSA-256 CuTe/Triton, 64 -> native Triton)
|
||||
tile_elems: int = _TILE_ELEMS
|
||||
# layers forced dense regardless of sparsity (probe-guided opt-outs)
|
||||
dense_layers: tuple[int, ...] = ()
|
||||
# Single-slot holder for the padded tile buffer, owned by the BUILDER so
|
||||
@@ -255,28 +158,24 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
VSA_sparsity: float,
|
||||
prefix_segments: tuple[int, ...],
|
||||
device: torch.device,
|
||||
exempt: bool = True,
|
||||
dense_layers: tuple[int, ...] = (),
|
||||
tile_size: int = _TILE_ELEMS,
|
||||
**kwargs: dict[str, Any],
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
VSA_sparsity: float,
|
||||
prefix_segments: tuple[int, ...],
|
||||
device: torch.device,
|
||||
exempt: bool = True,
|
||||
dense_layers: tuple[int, ...] = (),
|
||||
**kwargs: dict[str, Any],
|
||||
) -> MiniMaxH3VSAMetadata:
|
||||
tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
|
||||
if tile_shape is None:
|
||||
raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
|
||||
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2])
|
||||
prefix_segments = tuple(int(s) for s in prefix_segments if s > 0)
|
||||
total_seq_length = sum(prefix_segments) + math.prod(dit_seq_shape)
|
||||
|
||||
(_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles,
|
||||
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
|
||||
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
|
||||
|
||||
return MiniMaxH3VSAMetadata(
|
||||
current_timestep=current_timestep,
|
||||
@@ -287,14 +186,13 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
|
||||
exempt=exempt,
|
||||
variable_block_sizes=variable_block_sizes,
|
||||
untile_combined_index=untile_combined_index,
|
||||
tile_elems=int(tile_size),
|
||||
dense_layers=tuple(int(layer) for layer in dense_layers),
|
||||
tile_buf_holder=self._tile_buf_holder,
|
||||
)
|
||||
|
||||
|
||||
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems: int = _TILE_ELEMS) -> torch.Tensor:
|
||||
"""fp32 mean over each tile_elems-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
|
||||
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
|
||||
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
|
||||
|
||||
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
|
||||
written), so a plain sum with fp32 accumulation needs no validity mask
|
||||
@@ -302,8 +200,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems:
|
||||
the masked mean exactly.
|
||||
"""
|
||||
batch, seq_len, heads, dim = x.shape
|
||||
n_tiles = seq_len // tile_elems
|
||||
pooled = x.view(batch, n_tiles, tile_elems, heads, dim).sum(dim=2, dtype=torch.float32)
|
||||
n_tiles = seq_len // _TILE_ELEMS
|
||||
pooled = x.view(batch, n_tiles, _TILE_ELEMS, heads, dim).sum(dim=2, dtype=torch.float32)
|
||||
pooled = pooled / variable_block_sizes.view(1, -1, 1, 1)
|
||||
return pooled.permute(0, 2, 1, 3)
|
||||
|
||||
@@ -334,24 +232,6 @@ def _build_block_mask(
|
||||
return mask
|
||||
|
||||
|
||||
def _sm100a_unavailable_reason(sm100a_mod: Any, query_bhsd: torch.Tensor, variable_block_sizes: torch.Tensor,
|
||||
grad_mode: bool) -> str | None:
|
||||
"""Why the opt-in sm_100a forward route cannot run here, or None if it can.
|
||||
|
||||
Pure decision logic, split out so the routing is unit-testable without a
|
||||
GPU or the compiled extension (tests substitute ``sm100a_mod``). Order
|
||||
matters only for the message: the cheapest, most actionable reason first.
|
||||
"""
|
||||
if sm100a_mod is None:
|
||||
return "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
|
||||
if grad_mode:
|
||||
return "inputs require grad and the sm_100a kernel is forward-only; grad paths keep Triton"
|
||||
if not sm100a_mod.is_supported(query_bhsd, variable_block_sizes):
|
||||
return ("block_sparse_attn_sm100a.is_supported returned False (needs an sm_100 device, a built "
|
||||
"extension, bf16, head_dim 128, an even tile count, and integer tile sizes)")
|
||||
return None
|
||||
|
||||
|
||||
class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -379,7 +259,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
|
||||
"routed to the VSA-H3 backend; exclude it from the supported backends.")
|
||||
n_tiles = attn_metadata.variable_block_sizes.numel()
|
||||
target_shape = (x.shape[0], n_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])
|
||||
target_shape = (x.shape[0], n_tiles * _TILE_ELEMS, x.shape[-2], x.shape[-1])
|
||||
|
||||
# single scatter: untile_combined_index maps original row i to its
|
||||
# padded slot, so this is exactly the inverse of postprocess_output
|
||||
@@ -401,11 +281,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
gate_compress: torch.Tensor | None,
|
||||
attn_metadata: MiniMaxH3VSAMetadata,
|
||||
) -> torch.Tensor:
|
||||
tile_elems = attn_metadata.tile_elems
|
||||
if tile_elems == 64:
|
||||
if block_sparse_attn_64_bhsd is None:
|
||||
raise NotImplementedError("fastvideo_kernel.block_sparse_attn is not installed")
|
||||
elif block_sparse_attn_256_bshd is None:
|
||||
if block_sparse_attn_256_bshd is None:
|
||||
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
|
||||
|
||||
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
|
||||
@@ -415,8 +291,8 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
|
||||
scores = None
|
||||
if layer_sparsity > 0.0 or gate_compress is not None or probe_dir is not None:
|
||||
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
|
||||
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
|
||||
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (query.shape[-1]**0.5)
|
||||
if probe_dir is not None:
|
||||
record_probe(probe_dir, self.layer_idx, query, key, scores, attn_metadata)
|
||||
@@ -433,75 +309,18 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
attn_metadata.exempt,
|
||||
)
|
||||
|
||||
if tile_elems == 64:
|
||||
# Native 64-token path: the block map is already at the kernels'
|
||||
# granularity. Both 64-token entries take BHSD ([B, H, S_pad, D]);
|
||||
# mirror block_sparse_attn_256_bshd's Triton branch and transpose
|
||||
# around the call.
|
||||
q_bhsd = query.transpose(1, 2).contiguous()
|
||||
k_bhsd = key.transpose(1, 2).contiguous()
|
||||
v_bhsd = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Opt-in sm_100a CUDA forward (upstream PR #1719 + per-q-tile
|
||||
# q2k_num fix). Forward-only: grad-tracking calls stay on Triton
|
||||
# so autograd keeps the Triton fwd+bwd pairing untouched. The
|
||||
# kernel does return an LSE in Triton's M format, so a future
|
||||
# fwd/bwd pairing is possible, but it is not built here.
|
||||
use_sm100a = False
|
||||
if os.environ.get(VSA_SM100A_ENV, "0") == "1":
|
||||
grad_mode = torch.is_grad_enabled() and (query.requires_grad or key.requires_grad
|
||||
or value.requires_grad)
|
||||
reason = _sm100a_unavailable_reason(_sm100a, q_bhsd, attn_metadata.variable_block_sizes, grad_mode)
|
||||
if reason is None and map_to_index is None:
|
||||
reason = "fastvideo_kernel.triton_kernels.index (map_to_index) is not importable"
|
||||
if reason is None:
|
||||
use_sm100a = True
|
||||
elif not torch.compiler.is_compiling():
|
||||
logger.warning_once(f"{VSA_SM100A_ENV}=1 but falling back to the Triton-64 kernels: {reason}")
|
||||
|
||||
if use_sm100a:
|
||||
# The sm_100a entry is index-native; compact the bool map the
|
||||
# same way the Triton bool entry does internally. Per-row
|
||||
# counts are NON-uniform here (prefix query tiles are dense,
|
||||
# video tiles run prefix+top-k) -- legal for the fixed kernel,
|
||||
# silently wrong on the pre-fix upstream one.
|
||||
q2k_idx, q2k_num = map_to_index(mask)
|
||||
out_bhsd, _ = _sm100a.block_sparse_attn_sm100a(
|
||||
q_bhsd,
|
||||
k_bhsd,
|
||||
v_bhsd,
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
attn_metadata.variable_block_sizes.to(torch.int32),
|
||||
need_lse=False,
|
||||
)
|
||||
else:
|
||||
out_bhsd, _ = block_sparse_attn_64_bhsd(
|
||||
q_bhsd,
|
||||
k_bhsd,
|
||||
v_bhsd,
|
||||
mask,
|
||||
attn_metadata.variable_block_sizes,
|
||||
)
|
||||
out = out_bhsd.transpose(1, 2).contiguous()
|
||||
else:
|
||||
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
|
||||
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
|
||||
|
||||
if gate_compress is not None:
|
||||
# Wan-style compression branch: dense attention over pooled tiles,
|
||||
# broadcast to each tile's rows, scaled by the learned gate
|
||||
# (zero-initialized for H3 => branch contributes nothing until
|
||||
# finetuned; the model layer skips it entirely for all-zero gates).
|
||||
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
|
||||
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
|
||||
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
|
||||
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
|
||||
batch, seq_len, heads, dim = out.shape
|
||||
batch, _, heads, dim = out.shape
|
||||
n_tiles = attn_metadata.variable_block_sizes.numel()
|
||||
# Out-of-place: on the CuTe backend ``out`` is the tensor FA4's
|
||||
# autograd node saved for its backward, so an in-place add here
|
||||
# bumps its version counter and backward dies with "one of the
|
||||
# variables needed for gradient computation has been modified".
|
||||
out_tiled = out.view(batch, n_tiles, tile_elems, heads, dim)
|
||||
gate_tiled = gate_compress.view(batch, n_tiles, tile_elems, heads, dim)
|
||||
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
|
||||
out.view(batch, n_tiles, _TILE_ELEMS, heads,
|
||||
dim).addcmul_(out_c.unsqueeze(2), gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim))
|
||||
return out
|
||||
|
||||
@@ -60,7 +60,7 @@ def record_probe(
|
||||
gen = torch.Generator(device="cpu").manual_seed(step * 1000 + layer)
|
||||
# sample among video rows in the PADDED/tiled domain that are non-pad
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import token_tile_and_valid
|
||||
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
|
||||
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes)
|
||||
video_rows = torch.nonzero((token_tile >= P) & token_valid, as_tuple=False).flatten()
|
||||
idx = video_rows[torch.randint(0, video_rows.numel(), (_TRUE_ROWS, ), generator=gen).to(query.device)]
|
||||
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -1,375 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Benchmark generate-fewer-frames plus MLX RIFE interpolation.
|
||||
|
||||
This script keeps the video diffusion path delegated to
|
||||
``fastvideo.benchmarks.mlx_fastwan_bench._generate_cell``. It only orchestrates
|
||||
two frame-count cells and the postprocess interpolation step.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
import time
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import imageio.v2 as imageio
|
||||
import numpy as np
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import encode_prompt, make_rotary_embeddings
|
||||
from fastvideo.benchmarks.mlx_fastwan_bench import _generate_cell, _ms_ssim
|
||||
from fastvideo.mlx_runtime.rife_interp import RIFEBackendError, interpolate, load_model
|
||||
from fastvideo.mlx_runtime.memory import add_memory_limit_args, apply_memory_limits
|
||||
|
||||
FOX_PROMPT = "A fox runs through a misty pine forest, leaves kicking up behind it."
|
||||
DEFAULT_MODEL_ROOT = Path("/Users/aryank/models/qad_int8_v2")
|
||||
|
||||
|
||||
def _parse_timesteps(raw: str) -> list[int]:
|
||||
timesteps = [int(part.strip()) for part in raw.split(",") if part.strip()]
|
||||
if not timesteps:
|
||||
raise SystemExit("No DMD timesteps parsed from --dmd-denoising-steps")
|
||||
return timesteps
|
||||
|
||||
|
||||
def _read_video(path: Path) -> list[np.ndarray]:
|
||||
if not path.is_file():
|
||||
raise FileNotFoundError(f"Video does not exist: {path}")
|
||||
frames = [np.asarray(frame[:, :, :3], dtype=np.uint8) for frame in imageio.mimread(path)]
|
||||
if not frames:
|
||||
raise RuntimeError(f"No frames decoded from {path}")
|
||||
return frames
|
||||
|
||||
|
||||
def _write_video(path: Path, frames: list[np.ndarray], fps: int) -> None:
|
||||
if not frames:
|
||||
raise ValueError("Cannot write an empty frame list")
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with imageio.get_writer(str(path), fps=fps, macro_block_size=1, codec="libx264", quality=8) as writer:
|
||||
for frame in frames:
|
||||
writer.append_data(np.asarray(frame, dtype=np.uint8))
|
||||
|
||||
|
||||
def _copy_video(src: Path, dst: Path) -> Path:
|
||||
dst.parent.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copy2(src, dst)
|
||||
return dst
|
||||
|
||||
|
||||
def _make_generation_inputs(args, num_frames: int):
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
config_path = args.model_root / "transformer" / "config.json"
|
||||
checkpoint_path = args.model_root / "transformer" / "diffusion_pytorch_model.safetensors"
|
||||
if not config_path.is_file():
|
||||
raise SystemExit(f"Missing DiT config: {config_path}")
|
||||
if not checkpoint_path.is_file():
|
||||
raise SystemExit(f"Missing DiT checkpoint: {checkpoint_path}")
|
||||
|
||||
config = json.loads(config_path.read_text())
|
||||
latent_frames = (num_frames - 1) // 4 + 1
|
||||
latent_height = args.height // 8
|
||||
latent_width = args.width // 8
|
||||
freqs_cis = make_rotary_embeddings(
|
||||
config,
|
||||
latent_frames=latent_frames,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
|
||||
generator = torch.Generator(device="cpu").manual_seed(args.seed)
|
||||
latents_seed = torch.randn(
|
||||
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
).numpy()
|
||||
timesteps = _parse_timesteps(args.dmd_denoising_steps)
|
||||
renoise_by_step = [
|
||||
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
|
||||
for _ in range(max(0,
|
||||
len(timesteps) - 1))
|
||||
]
|
||||
prompt_embeds = encode_prompt(
|
||||
model_root=args.model_root,
|
||||
prompt=args.prompt,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
)
|
||||
return {
|
||||
"checkpoint_path": checkpoint_path,
|
||||
"config_path": config_path,
|
||||
"encoder_hidden_states": mx.array(prompt_embeds.numpy()),
|
||||
"freqs_cis": freqs_cis,
|
||||
"timesteps": timesteps,
|
||||
"latents_seed": latents_seed,
|
||||
"renoise_by_step": renoise_by_step,
|
||||
"latent_frames": latent_frames,
|
||||
}
|
||||
|
||||
|
||||
def _generate(args, num_frames: int):
|
||||
cell_args = SimpleNamespace(
|
||||
model_root=args.model_root,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=num_frames,
|
||||
fps=args.fps,
|
||||
flow_shift=args.flow_shift,
|
||||
torch_device=args.torch_device,
|
||||
torch_dtype=args.torch_dtype,
|
||||
taehv_source_path=args.taehv_source_path,
|
||||
taehv_checkpoint_path=args.taehv_checkpoint_path,
|
||||
taehv_parallel=args.taehv_parallel,
|
||||
mlx_checkpoint_cache=args.mlx_checkpoint_cache,
|
||||
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
|
||||
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
|
||||
mlx_disable_cache=args.mlx_disable_cache,
|
||||
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
|
||||
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
|
||||
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
|
||||
benchmark_preset="metalfx-rife",
|
||||
current_prompt_id=f"fox-{num_frames}f",
|
||||
current_prompt=args.prompt,
|
||||
output_dir=args.output_dir,
|
||||
)
|
||||
inputs = _make_generation_inputs(args, num_frames)
|
||||
print(f"=== generate: {num_frames} frames ({inputs['latent_frames']} latent frames) ===", flush=True)
|
||||
return _generate_cell(
|
||||
args=cell_args,
|
||||
mode=args.mode,
|
||||
decoder=args.decoder,
|
||||
checkpoint_path=inputs["checkpoint_path"],
|
||||
config_path=inputs["config_path"],
|
||||
encoder_hidden_states=inputs["encoder_hidden_states"],
|
||||
freqs_cis=inputs["freqs_cis"],
|
||||
timesteps=inputs["timesteps"],
|
||||
latents_seed=inputs["latents_seed"],
|
||||
renoise_by_step=inputs["renoise_by_step"],
|
||||
)
|
||||
|
||||
|
||||
def _relative_to_output(path: Path, output_dir: Path) -> str:
|
||||
try:
|
||||
return str(path.relative_to(output_dir))
|
||||
except ValueError:
|
||||
return str(path)
|
||||
|
||||
|
||||
def _write_report(args, result: dict) -> Path:
|
||||
report_path = args.output_dir / "metalfx_rife_report.md"
|
||||
rows = result["speed_rows"]
|
||||
table_lines = [
|
||||
"| path | denoise_s | decode_s | gen_total_s | rife_s | net_s | speedup_vs_81 |",
|
||||
"| --- | ---: | ---: | ---: | ---: | ---: | ---: |",
|
||||
]
|
||||
for row in rows:
|
||||
table_lines.append(
|
||||
"| {path} | {denoise_s:.3f} | {decode_s:.3f} | {gen_total_s:.3f} | {rife_s:.3f} | {net_s:.3f} | {speedup_vs_81:.3f}x |"
|
||||
.format(**row))
|
||||
text = f"""# MetalFX/RIFE Generate-Fewer-Frames Benchmark Run
|
||||
|
||||
Prompt: {args.prompt}
|
||||
|
||||
Resolution: {args.height}x{args.width}, fps={args.fps}, mode={args.mode}, decoder={args.decoder}
|
||||
|
||||
| metric | value |
|
||||
| --- | ---: |
|
||||
| reconstruction_ms_ssim | {result['reconstruction_ms_ssim']:.6f} |
|
||||
| reference_frames | {result['reference_frames']} |
|
||||
| reduced_frames | {result['reduced_frames']} |
|
||||
| interpolated_frames | {result['interpolated_frames']} |
|
||||
|
||||
{chr(10).join(table_lines)}
|
||||
|
||||
Videos:
|
||||
|
||||
- reference: `{result['videos']['reference']}`
|
||||
- reference drop-41 RIFE reconstruction: `{result['videos']['drop41_rife81']}`
|
||||
- generated 41 RIFE to 81: `{result['videos']['generated41_rife81']}`
|
||||
- generated 41 direct: `{result['videos']['generated41']}`
|
||||
|
||||
Raw metrics are in `{_relative_to_output(args.output_dir / 'metrics.json', args.output_dir)}`.
|
||||
"""
|
||||
report_path.write_text(text)
|
||||
return report_path
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Evaluate 41-frame generation plus MLX RIFE interpolation vs 81-frame generation.")
|
||||
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
|
||||
parser.add_argument("--prompt", default=FOX_PROMPT)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=832)
|
||||
parser.add_argument("--reference-frames", type=int, default=81)
|
||||
parser.add_argument("--reduced-frames", type=int, default=41)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--mode", default="int8", choices=("fp16", "bf16", "int8", "int4", "mxfp8", "mxfp4", "nvfp4"))
|
||||
parser.add_argument("--decoder", default="taehv", choices=("taehv", "wan-vae"))
|
||||
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
|
||||
parser.add_argument("--flow-shift", type=float, default=8.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=512)
|
||||
parser.add_argument("--torch-device", default="auto")
|
||||
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument("--output-dir", type=Path, default=Path("bench/metalfx_rife"))
|
||||
parser.add_argument("--mlx-checkpoint-cache", type=Path, default=None)
|
||||
parser.add_argument("--compile", action="store_true", help="Enable FASTVIDEO_MLX_COMPILE=1 for DiT denoise.")
|
||||
parser.add_argument("--rife-scale", type=float, default=1.0)
|
||||
parser.add_argument("--taehv-source-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-parallel", action="store_true")
|
||||
add_memory_limit_args(parser)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.reference_frames != (args.reduced_frames - 1) * 2 + 1:
|
||||
raise SystemExit(
|
||||
"--reference-frames must equal (--reduced-frames - 1) * 2 + 1 for the default every-other-frame test")
|
||||
if args.compile:
|
||||
import os
|
||||
|
||||
os.environ["FASTVIDEO_MLX_COMPILE"] = "1"
|
||||
|
||||
args.model_root = args.model_root.expanduser().resolve()
|
||||
args.output_dir = args.output_dir.expanduser().resolve()
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
if args.mlx_checkpoint_cache is None:
|
||||
args.mlx_checkpoint_cache = args.output_dir / "mlx_checkpoint_cache"
|
||||
else:
|
||||
args.mlx_checkpoint_cache = args.mlx_checkpoint_cache.expanduser().resolve()
|
||||
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
runtime_limits = apply_memory_limits(
|
||||
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
|
||||
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
|
||||
mlx_disable_cache=args.mlx_disable_cache,
|
||||
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
|
||||
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
|
||||
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
|
||||
mx_module=mx,
|
||||
).as_metrics()
|
||||
mx.random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
|
||||
reference_cell = _generate(args, args.reference_frames)
|
||||
reduced_cell = _generate(args, args.reduced_frames)
|
||||
|
||||
reference_video = _copy_video(reference_cell.video_path, args.output_dir / "fox_reference_81.mp4")
|
||||
generated41_video = _copy_video(reduced_cell.video_path, args.output_dir / "fox_generated_41.mp4")
|
||||
|
||||
reference_frames = _read_video(reference_video)
|
||||
dropped_reference_frames = reference_frames[::2]
|
||||
if len(dropped_reference_frames) != args.reduced_frames:
|
||||
raise RuntimeError(f"Expected {args.reduced_frames} dropped frames, got {len(dropped_reference_frames)}")
|
||||
|
||||
try:
|
||||
rife_model = load_model("4.25")
|
||||
except RIFEBackendError:
|
||||
raise
|
||||
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
|
||||
raise RIFEBackendError(f"Unexpected RIFE load failure: {exc}") from exc
|
||||
|
||||
print("=== RIFE reconstruction: drop every other reference frame, 41 -> 81 ===", flush=True)
|
||||
start = time.perf_counter()
|
||||
reconstructed_frames = interpolate(dropped_reference_frames, factor=2, model=rife_model, scale=args.rife_scale)
|
||||
recon_rife_s = time.perf_counter() - start
|
||||
reconstructed_video = args.output_dir / "fox_reference_drop41_rife81.mp4"
|
||||
_write_video(reconstructed_video, reconstructed_frames, args.fps)
|
||||
reconstruction_ms_ssim = _ms_ssim(reference_video, reconstructed_video, required=True)
|
||||
if reconstruction_ms_ssim is None:
|
||||
raise RuntimeError("MS-SSIM returned None for reconstruction comparison")
|
||||
|
||||
print("=== RIFE actual reduced path: generated 41 -> 81 ===", flush=True)
|
||||
generated41_frames = _read_video(generated41_video)
|
||||
start = time.perf_counter()
|
||||
generated41_rife_frames = interpolate(generated41_frames, factor=2, model=rife_model, scale=args.rife_scale)
|
||||
generated41_rife_s = time.perf_counter() - start
|
||||
generated41_rife_video = args.output_dir / "fox_generated41_rife81.mp4"
|
||||
_write_video(generated41_rife_video, generated41_rife_frames, args.fps)
|
||||
|
||||
direct_denoise_s = float(reference_cell.metrics["denoise_s"])
|
||||
reduced_denoise_s = float(reduced_cell.metrics["denoise_s"])
|
||||
direct_gen_total_s = float(reference_cell.metrics["denoise_s"]) + float(reference_cell.metrics["decode_s"])
|
||||
reduced_gen_total_s = float(reduced_cell.metrics["denoise_s"]) + float(reduced_cell.metrics["decode_s"])
|
||||
net_denoise_rife_s = reduced_denoise_s + generated41_rife_s
|
||||
net_gen_rife_s = reduced_gen_total_s + generated41_rife_s
|
||||
|
||||
result = {
|
||||
"prompt":
|
||||
args.prompt,
|
||||
"model_root":
|
||||
str(args.model_root),
|
||||
"rife_impl":
|
||||
"rife-mlx vendored at fastvideo/third_party/rife_mlx, weights mlx-community/RIFE-4.25",
|
||||
"runtime_limits":
|
||||
runtime_limits,
|
||||
"reference_frames":
|
||||
len(reference_frames),
|
||||
"reduced_frames":
|
||||
len(generated41_frames),
|
||||
"interpolated_frames":
|
||||
len(generated41_rife_frames),
|
||||
"reconstruction_ms_ssim":
|
||||
reconstruction_ms_ssim,
|
||||
"reconstruction_rife_s":
|
||||
recon_rife_s,
|
||||
"generated41_rife_s":
|
||||
generated41_rife_s,
|
||||
"reference_metrics":
|
||||
reference_cell.metrics,
|
||||
"reduced_metrics":
|
||||
reduced_cell.metrics,
|
||||
"speed_rows": [
|
||||
{
|
||||
"path": "generate_81",
|
||||
"denoise_s": direct_denoise_s,
|
||||
"decode_s": float(reference_cell.metrics["decode_s"]),
|
||||
"gen_total_s": direct_gen_total_s,
|
||||
"rife_s": 0.0,
|
||||
"net_s": direct_gen_total_s,
|
||||
"speedup_vs_81": 1.0,
|
||||
},
|
||||
{
|
||||
"path": "generate_41_plus_rife81_denoise_only",
|
||||
"denoise_s": reduced_denoise_s,
|
||||
"decode_s": 0.0,
|
||||
"gen_total_s": reduced_denoise_s,
|
||||
"rife_s": generated41_rife_s,
|
||||
"net_s": net_denoise_rife_s,
|
||||
"speedup_vs_81": direct_denoise_s / net_denoise_rife_s,
|
||||
},
|
||||
{
|
||||
"path": "generate_41_plus_rife81_decode_included",
|
||||
"denoise_s": reduced_denoise_s,
|
||||
"decode_s": float(reduced_cell.metrics["decode_s"]),
|
||||
"gen_total_s": reduced_gen_total_s,
|
||||
"rife_s": generated41_rife_s,
|
||||
"net_s": net_gen_rife_s,
|
||||
"speedup_vs_81": direct_gen_total_s / net_gen_rife_s,
|
||||
},
|
||||
],
|
||||
"videos": {
|
||||
"reference": _relative_to_output(reference_video, args.output_dir),
|
||||
"drop41_rife81": _relative_to_output(reconstructed_video, args.output_dir),
|
||||
"generated41": _relative_to_output(generated41_video, args.output_dir),
|
||||
"generated41_rife81": _relative_to_output(generated41_rife_video, args.output_dir),
|
||||
},
|
||||
}
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(result, indent=2))
|
||||
report_path = _write_report(args, result)
|
||||
|
||||
print(json.dumps(result["speed_rows"], indent=2))
|
||||
print(f"reconstruction_ms_ssim={reconstruction_ms_ssim:.6f}")
|
||||
print(f"wrote {metrics_path}")
|
||||
print(f"wrote {report_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,826 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Prove-out benchmark for the MLX FastWan runtime (Apple Silicon).
|
||||
|
||||
Sweeps ``{dtype/quant} x {decoder}``, generates a clip per cell, and records the
|
||||
latency breakdown, peak unified memory, and MS-SSIM (optionally LPIPS) against a
|
||||
reference video. It emits a JSON blob and a markdown table -- the artifact that
|
||||
turns "int8 + TAEHV looks good" into defensible numbers, and (via
|
||||
``--assert-min-ssim``) a regression gate for the ``mx.compile`` work.
|
||||
|
||||
Design notes:
|
||||
- Generation reuses the hybrid POC helpers in
|
||||
``examples/inference/basic/mlx_wan_prompt_to_video.py`` (torch-MPS UMT5 encode
|
||||
and Wan-VAE/TAEHV decode) plus the on-device MLX DMD sampler
|
||||
(``fastvideo/mlx_runtime/sampling.py``); the denoise loop never leaves the
|
||||
device.
|
||||
- Quality reuses the tested MS-SSIM primitive
|
||||
``fastvideo/tests/utils.py::compute_video_ssim_torchvision``.
|
||||
- Reference: by default each cell is scored against the highest-fidelity cell
|
||||
in the sweep (``fp16`` + ``wan-vae``), which needs no CUDA box and answers
|
||||
"how much does int8/int4/TAEHV degrade vs the best local config". Pass
|
||||
``--reference PATH`` to score against an external clip instead (e.g. the
|
||||
torch-MPS or CUDA FastVideo output of the same model) for a "vs. the original
|
||||
model" column.
|
||||
|
||||
Run on an Apple Silicon Mac (needs ``mlx`` + a torch build with MPS):
|
||||
|
||||
python fastvideo/benchmarks/mlx_fastwan_bench.py \
|
||||
--modes fp16,bf16,int8,int4 --decoders taehv,wan-vae
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import html
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from examples.inference.basic.mlx_wan_prompt_to_video import (
|
||||
DEFAULT_MODEL_ROOT,
|
||||
decode_latents_to_video,
|
||||
encode_prompt,
|
||||
make_rotary_embeddings,
|
||||
)
|
||||
from fastvideo.mlx_runtime.memory import add_memory_limit_args, apply_memory_limits, cleanup_mlx
|
||||
|
||||
# The highest-fidelity cell; used as the default SSIM reference when no external
|
||||
# reference video is supplied.
|
||||
REFERENCE_MODE = "fp16"
|
||||
REFERENCE_DECODER = "wan-vae"
|
||||
|
||||
ALLOWED_MODES = ("fp16", "bf16", "int8", "int4", "mxfp8", "mxfp4", "nvfp4")
|
||||
ALLOWED_DECODERS = ("taehv", "wan-vae")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PromptCase:
|
||||
id: str
|
||||
prompt: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class BenchmarkPreset:
|
||||
height: int
|
||||
width: int
|
||||
num_frames: int
|
||||
modes: str
|
||||
decoders: str
|
||||
mlx_memory_limit_gib: float | None = None
|
||||
mlx_disable_cache: bool = False
|
||||
torch_mps_high_watermark_ratio: float | None = None
|
||||
torch_mps_low_watermark_ratio: float | None = None
|
||||
|
||||
|
||||
PROMPT_SETS = {
|
||||
"motion7": (
|
||||
PromptCase("beach-sunset", "A slow cinematic sunset over ocean waves at a quiet beach."),
|
||||
PromptCase("fox-forest", "A fox runs through a misty pine forest, leaves kicking up behind it."),
|
||||
PromptCase("raccoon-sunflowers", "A raccoon walks through a sunflower field as petals move in the wind."),
|
||||
PromptCase("surfing-cat", "A cat wearing sunglasses surfs across a bright blue ocean wave."),
|
||||
PromptCase("burning-clock", "A vintage table clock burns on a wooden desk, flames flickering realistically."),
|
||||
PromptCase("forest-walk", "Video game style footage of a man walking through a dense forest path."),
|
||||
PromptCase("sea-dock-yachts", "Several yachts are parked at a sea dock while water ripples around them."),
|
||||
),
|
||||
}
|
||||
|
||||
BENCHMARK_PRESETS = {
|
||||
"default":
|
||||
BenchmarkPreset(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=81,
|
||||
modes="fp16,bf16,int8,int4",
|
||||
decoders="taehv,wan-vae",
|
||||
),
|
||||
"mac-16gb":
|
||||
BenchmarkPreset(
|
||||
height=448,
|
||||
width=832,
|
||||
num_frames=61,
|
||||
modes="int8",
|
||||
decoders="taehv",
|
||||
mlx_memory_limit_gib=16.0,
|
||||
mlx_disable_cache=True,
|
||||
torch_mps_high_watermark_ratio=0.57,
|
||||
torch_mps_low_watermark_ratio=0.0,
|
||||
),
|
||||
"mac-32gb":
|
||||
BenchmarkPreset(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=81,
|
||||
modes="int8,fp16",
|
||||
decoders="taehv",
|
||||
),
|
||||
"mac-64gb":
|
||||
BenchmarkPreset(
|
||||
height=480,
|
||||
width=832,
|
||||
num_frames=81,
|
||||
modes="int8,fp16",
|
||||
decoders="taehv,wan-vae",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cell:
|
||||
prompt_id: str
|
||||
prompt: str
|
||||
mode: str
|
||||
decoder: str
|
||||
video_path: Path
|
||||
latents: np.ndarray
|
||||
metrics: dict[str, float | int | str | bool | None] = field(default_factory=dict)
|
||||
|
||||
|
||||
def _mode_to_dtype_quant(mode: str) -> tuple[str, str | None]:
|
||||
"""Map a sweep mode to (MLX compute dtype, quantization spec).
|
||||
|
||||
Quantized modes keep fp16 activations and quantize only the DiT linear
|
||||
weights (matching ``mlx_dit_from_diffusers_safetensors``).
|
||||
"""
|
||||
if mode == "bf16":
|
||||
return "bf16", None
|
||||
if mode == "fp16":
|
||||
return "fp16", None
|
||||
# int8/int4/mxfp*/nvfp4 -> fp16 activations + quantized weights.
|
||||
return "fp16", mode
|
||||
|
||||
|
||||
def _mx_dtype(mx, base: str):
|
||||
return {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[base]
|
||||
|
||||
|
||||
def _parse_list(raw: str, allowed: tuple[str, ...], label: str) -> list[str]:
|
||||
items = [x.strip() for x in raw.split(",") if x.strip()]
|
||||
unknown = sorted(set(items) - set(allowed))
|
||||
if unknown:
|
||||
raise ValueError(f"Unsupported {label}: {unknown} (allowed: {list(allowed)})")
|
||||
return items
|
||||
|
||||
|
||||
def _safe_slug(value: str, *, fallback: str) -> str:
|
||||
slug = "".join(ch.lower() if ch.isalnum() else "-" for ch in value.strip())
|
||||
slug = "-".join(part for part in slug.split("-") if part)
|
||||
return slug[:64] or fallback
|
||||
|
||||
|
||||
def _load_prompt_cases(prompt: str, prompt_file: Path | None, prompt_set: str = "single") -> list[PromptCase]:
|
||||
"""Load one prompt, a built-in prompt set, or a text/jsonl prompt file."""
|
||||
if prompt_file is None:
|
||||
if prompt_set == "single":
|
||||
return [PromptCase(id="prompt-001", prompt=prompt)]
|
||||
if prompt_set not in PROMPT_SETS:
|
||||
raise ValueError(f"Unsupported prompt set: {prompt_set} (allowed: {sorted(PROMPT_SETS) + ['single']})")
|
||||
return list(PROMPT_SETS[prompt_set])
|
||||
cases: list[PromptCase] = []
|
||||
for line_index, raw_line in enumerate(prompt_file.read_text().splitlines(), start=1):
|
||||
line = raw_line.strip()
|
||||
if not line or line.startswith("#"):
|
||||
continue
|
||||
prompt_id = f"prompt-{len(cases) + 1:03d}"
|
||||
prompt_text = line
|
||||
if prompt_file.suffix.lower() == ".jsonl":
|
||||
item = json.loads(line)
|
||||
prompt_text = str(item.get("prompt") or item.get("text") or item.get("caption") or "").strip()
|
||||
if not prompt_text:
|
||||
raise ValueError(f"{prompt_file}:{line_index} has no prompt/text/caption field")
|
||||
prompt_id = str(item.get("id") or item.get("name") or prompt_id)
|
||||
cases.append(PromptCase(id=_safe_slug(prompt_id, fallback=f"prompt-{len(cases) + 1:03d}"), prompt=prompt_text))
|
||||
if not cases:
|
||||
raise ValueError(f"No prompts found in {prompt_file}")
|
||||
return cases
|
||||
|
||||
|
||||
def denoise_dmd_on_device(
|
||||
*,
|
||||
mx,
|
||||
dit,
|
||||
latents,
|
||||
encoder_hidden_states,
|
||||
freqs_cis,
|
||||
timesteps: list[int],
|
||||
renoise_by_step: list[np.ndarray],
|
||||
schedule,
|
||||
dmd_step,
|
||||
mx_dtype,
|
||||
) -> tuple[np.ndarray, list[float]]:
|
||||
"""Run the FastWan DMD loop entirely on the MLX device.
|
||||
|
||||
Mirrors the loop in ``mlx_wan_prompt_to_video.py`` (fp32 affine math, MLX RNG
|
||||
re-noise) so the benchmark measures exactly the shipped path.
|
||||
|
||||
Returns the final latents plus per-step wall times. The first step carries
|
||||
one-time costs (mx.compile tracing, kernel warm-up), so first-vs-steady
|
||||
step timing is how the benchmark separates cold-start from steady-state
|
||||
denoise throughput.
|
||||
|
||||
All host-side tensors (timesteps, re-noise draws) are uploaded before the
|
||||
loop starts, so the per-step body performs no bulk host->device transfers
|
||||
and step timings measure device work rather than staging copies.
|
||||
"""
|
||||
timesteps_mx = [mx.array([float(timestep)]).astype(mx.float32) for timestep in timesteps]
|
||||
renoise_mx = [mx.array(renoise).astype(mx.float32) for renoise in renoise_by_step]
|
||||
if timesteps_mx or renoise_mx:
|
||||
mx.eval(*timesteps_mx, *renoise_mx)
|
||||
|
||||
step_times: list[float] = []
|
||||
for step_index, timestep in enumerate(timesteps):
|
||||
step_start = time.perf_counter()
|
||||
noise_input_latent = latents
|
||||
noise_pred = dit(latents.astype(mx_dtype), encoder_hidden_states, timesteps_mx[step_index], freqs_cis)
|
||||
|
||||
noise_input_f32 = noise_input_latent.astype(mx.float32)
|
||||
pred_noise_f32 = noise_pred.astype(mx.float32)
|
||||
if step_index < len(timesteps) - 1:
|
||||
next_ts: float | None = float(timesteps[step_index + 1])
|
||||
renoise = renoise_mx[step_index]
|
||||
else:
|
||||
next_ts, renoise = None, None
|
||||
latents = dmd_step(
|
||||
latents=noise_input_f32,
|
||||
noise_input_latent=noise_input_f32,
|
||||
pred_noise=pred_noise_f32,
|
||||
schedule=schedule,
|
||||
timestep=float(timestep),
|
||||
next_timestep=next_ts,
|
||||
noise=renoise,
|
||||
).astype(mx_dtype)
|
||||
mx.eval(latents)
|
||||
step_times.append(time.perf_counter() - step_start)
|
||||
return np.array(latents.astype(mx.float32)), step_times
|
||||
|
||||
|
||||
def _peak_memory_bytes(mx) -> int:
|
||||
try:
|
||||
return int(mx.get_peak_memory())
|
||||
except Exception: # noqa: BLE001 - best-effort telemetry only.
|
||||
return 0
|
||||
|
||||
|
||||
def _latent_delta_metrics(candidate: np.ndarray, baseline: np.ndarray) -> dict[str, float]:
|
||||
diff = candidate.astype(np.float32) - baseline.astype(np.float32)
|
||||
mse = float(np.mean(np.square(diff)))
|
||||
signal = float(np.mean(np.square(baseline.astype(np.float32))))
|
||||
return {
|
||||
"latent_mse_vs_ref_mode": mse,
|
||||
"latent_snr_db_vs_ref_mode": float(10.0 * np.log10(signal / mse)) if mse > 0 else float("inf"),
|
||||
}
|
||||
|
||||
|
||||
def _ms_ssim(reference_video: Path, candidate_video: Path, *, required: bool = False) -> float | None:
|
||||
"""Mean MS-SSIM between two mp4s, via the repo's tested helper."""
|
||||
if not reference_video.exists() or not candidate_video.exists():
|
||||
return None
|
||||
try:
|
||||
from fastvideo.tests.utils import compute_video_ssim_torchvision
|
||||
except ImportError as exc:
|
||||
message = ("MS-SSIM is unavailable because `pytorch-msssim` is not installed. "
|
||||
"Install FastVideo with the test extra, e.g. `uv pip install -e '.[mlx,test]'`, "
|
||||
"or run without an SSIM assertion.")
|
||||
if required:
|
||||
raise RuntimeError(message) from exc
|
||||
print(f"{message} Skipping MS-SSIM.")
|
||||
return None
|
||||
|
||||
try:
|
||||
ssim_values = compute_video_ssim_torchvision(str(reference_video), str(candidate_video), use_ms_ssim=True)
|
||||
except ImportError as exc:
|
||||
message = ("MS-SSIM is unavailable because `pytorch-msssim` is not installed. "
|
||||
"Install FastVideo with the test extra, e.g. `uv pip install -e '.[mlx,test]'`, "
|
||||
"or run without an SSIM assertion.")
|
||||
if required:
|
||||
raise RuntimeError(message) from exc
|
||||
print(f"{message} Skipping MS-SSIM.")
|
||||
return None
|
||||
return float(ssim_values[0])
|
||||
|
||||
|
||||
def _markdown_table(rows: list[dict]) -> str:
|
||||
columns = [
|
||||
("prompt_id", "prompt"),
|
||||
("mode", "mode"),
|
||||
("decoder", "decoder"),
|
||||
("status", "status"),
|
||||
("denoise_s", "denoise s"),
|
||||
("decode_s", "decode s"),
|
||||
("total_s", "total s"),
|
||||
("peak_gib", "peak GiB"),
|
||||
("ms_ssim_vs_ref", "MS-SSIM"),
|
||||
("lpips_vs_ref", "LPIPS"),
|
||||
]
|
||||
header = "| " + " | ".join(label for _, label in columns) + " |"
|
||||
sep = "| " + " | ".join("---" for _ in columns) + " |"
|
||||
lines = [header, sep]
|
||||
for row in rows:
|
||||
cells = []
|
||||
for key, _ in columns:
|
||||
value = row.get(key)
|
||||
if isinstance(value, float):
|
||||
cells.append(f"{value:.3f}")
|
||||
elif value is None:
|
||||
cells.append("-")
|
||||
else:
|
||||
cells.append(str(value))
|
||||
lines.append("| " + " | ".join(cells) + " |")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _format_metric(value) -> str:
|
||||
if isinstance(value, float):
|
||||
return f"{value:.3f}"
|
||||
if value is None:
|
||||
return "-"
|
||||
return str(value)
|
||||
|
||||
|
||||
def _html_grid(rows: list[dict]) -> str:
|
||||
groups: dict[str, list[dict]] = {}
|
||||
for row in rows:
|
||||
groups.setdefault(str(row.get("prompt_id", "prompt")), []).append(row)
|
||||
|
||||
sections = []
|
||||
for prompt_id, group_rows in groups.items():
|
||||
prompt = next((str(row.get("prompt", "")) for row in group_rows if row.get("prompt")), "")
|
||||
cards = []
|
||||
for row in group_rows:
|
||||
title = f"{row.get('mode', '-')} / {row.get('decoder', '-')}"
|
||||
status = row.get("status", "-")
|
||||
video_path = row.get("video_path")
|
||||
if video_path:
|
||||
media = f'<video src="{html.escape(str(video_path))}" muted loop controls playsinline></video>'
|
||||
else:
|
||||
media = f'<div class="missing">No video<br>{html.escape(str(row.get("error", "")))}</div>'
|
||||
metrics = (f"status={status} · total={_format_metric(row.get('total_s'))}s · "
|
||||
f"denoise={_format_metric(row.get('denoise_s'))}s · "
|
||||
f"decode={_format_metric(row.get('decode_s'))}s · "
|
||||
f"peak={_format_metric(row.get('peak_gib'))}GiB")
|
||||
cards.append("<article>"
|
||||
f"<h3>{html.escape(title)}</h3>"
|
||||
f"{media}"
|
||||
f"<p>{html.escape(metrics)}</p>"
|
||||
"</article>")
|
||||
sections.append("<section>"
|
||||
f"<h2>{html.escape(prompt_id)}</h2>"
|
||||
f"<p class=\"prompt\">{html.escape(prompt)}</p>"
|
||||
f"<div class=\"grid\">{''.join(cards)}</div>"
|
||||
"</section>")
|
||||
|
||||
return """<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1">
|
||||
<title>FastVideo MLX benchmark grid</title>
|
||||
<style>
|
||||
body { font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif; margin: 24px; background: #111; color: #eee; }
|
||||
button { margin-right: 8px; padding: 8px 12px; border-radius: 8px; border: 1px solid #555; background: #222; color: #eee; }
|
||||
section { margin-top: 28px; }
|
||||
.prompt { color: #bbb; max-width: 900px; }
|
||||
.grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(280px, 1fr)); gap: 16px; }
|
||||
article { background: #1b1b1b; border: 1px solid #333; border-radius: 12px; padding: 12px; }
|
||||
h1, h2, h3 { margin: 0 0 10px; }
|
||||
video { width: 100%; border-radius: 8px; background: #000; }
|
||||
article p { color: #bbb; font-size: 13px; line-height: 1.4; }
|
||||
.missing { min-height: 160px; display: grid; place-items: center; text-align: center; color: #f5b5b5; background: #2a1515; border-radius: 8px; padding: 12px; }
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<h1>FastVideo MLX benchmark grid</h1>
|
||||
<p>Use the controls below to start/stop every clip together for side-by-side inspection.</p>
|
||||
<button onclick="for (const v of document.querySelectorAll('video')) { v.currentTime = 0; v.play(); }">Restart + play all</button>
|
||||
<button onclick="for (const v of document.querySelectorAll('video')) v.pause();">Pause all</button>
|
||||
""" + "\n".join(sections) + """
|
||||
</body>
|
||||
</html>
|
||||
"""
|
||||
|
||||
|
||||
def _write_html_grid(rows: list[dict], output_dir: Path) -> Path:
|
||||
html_path = output_dir / "index.html"
|
||||
html_path.write_text(_html_grid(rows))
|
||||
return html_path
|
||||
|
||||
|
||||
def _generate_cell(
|
||||
*,
|
||||
args,
|
||||
mode: str,
|
||||
decoder: str,
|
||||
checkpoint_path: Path,
|
||||
config_path: Path,
|
||||
encoder_hidden_states,
|
||||
freqs_cis,
|
||||
timesteps: list[int],
|
||||
latents_seed: np.ndarray,
|
||||
renoise_by_step: list[np.ndarray],
|
||||
) -> Cell:
|
||||
import mlx.core as mx
|
||||
|
||||
from fastvideo.mlx_runtime.fastwan import mlx_dit_from_diffusers_safetensors
|
||||
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
base_dtype, quantization = _mode_to_dtype_quant(mode)
|
||||
mx_dtype = _mx_dtype(mx, base_dtype)
|
||||
|
||||
mx.clear_cache()
|
||||
mx.reset_peak_memory()
|
||||
load_start = time.perf_counter()
|
||||
load_source = "diffusers"
|
||||
if args.mlx_checkpoint_cache is not None:
|
||||
from fastvideo.mlx_runtime.checkpoint import (
|
||||
load_mlx_dit_checkpoint,
|
||||
save_mlx_dit_checkpoint,
|
||||
)
|
||||
|
||||
mode_ckpt_dir = args.mlx_checkpoint_cache / mode
|
||||
if (mode_ckpt_dir / "mlx_dit.json").exists():
|
||||
dit = load_mlx_dit_checkpoint(mode_ckpt_dir)
|
||||
load_source = "mlx_checkpoint"
|
||||
else:
|
||||
dit = mlx_dit_from_diffusers_safetensors(
|
||||
checkpoint_path,
|
||||
config_path,
|
||||
dtype=base_dtype,
|
||||
quantization=quantization,
|
||||
)
|
||||
save_mlx_dit_checkpoint(dit, mode_ckpt_dir)
|
||||
load_source = "diffusers_then_saved"
|
||||
else:
|
||||
dit = mlx_dit_from_diffusers_safetensors(
|
||||
checkpoint_path,
|
||||
config_path,
|
||||
dtype=base_dtype,
|
||||
quantization=quantization,
|
||||
)
|
||||
load_s = time.perf_counter() - load_start
|
||||
load_peak = _peak_memory_bytes(mx)
|
||||
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=args.flow_shift)
|
||||
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
|
||||
|
||||
latents = mx.array(latents_seed).astype(mx_dtype)
|
||||
mx.reset_peak_memory()
|
||||
denoise_start = time.perf_counter()
|
||||
latents_np, step_times = denoise_dmd_on_device(
|
||||
mx=mx,
|
||||
dit=dit,
|
||||
latents=latents,
|
||||
encoder_hidden_states=encoder_hidden_states.astype(mx_dtype),
|
||||
freqs_cis=freqs_cis,
|
||||
timesteps=timesteps,
|
||||
renoise_by_step=renoise_by_step,
|
||||
schedule=schedule,
|
||||
dmd_step=dmd_step,
|
||||
mx_dtype=mx_dtype,
|
||||
)
|
||||
denoise_s = time.perf_counter() - denoise_start
|
||||
denoise_peak = _peak_memory_bytes(mx)
|
||||
del dit, latents
|
||||
cleanup_mlx(mx)
|
||||
|
||||
video_path = (args.output_dir / f"{args.current_prompt_id}" /
|
||||
f"video_{mode}_{decoder}_{args.height}x{args.width}x{args.num_frames}.mp4")
|
||||
decode_start = time.perf_counter()
|
||||
decode_latents_to_video(
|
||||
model_root=args.model_root,
|
||||
latents_np=latents_np,
|
||||
output_path=video_path,
|
||||
fps=args.fps,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
backend=decoder,
|
||||
taehv_source_path=args.taehv_source_path,
|
||||
taehv_checkpoint_path=args.taehv_checkpoint_path,
|
||||
taehv_parallel=args.taehv_parallel,
|
||||
)
|
||||
decode_s = time.perf_counter() - decode_start
|
||||
|
||||
metrics: dict[str, float | int | str | bool | None] = {
|
||||
"prompt_id": args.current_prompt_id,
|
||||
"prompt": args.current_prompt,
|
||||
"benchmark_preset": args.benchmark_preset,
|
||||
"mode": mode,
|
||||
"decoder": decoder,
|
||||
"status": "ok",
|
||||
"video_path": str(video_path.relative_to(args.output_dir)),
|
||||
"load_s": load_s,
|
||||
"load_source": load_source,
|
||||
"denoise_s": denoise_s,
|
||||
# The first step carries one-time costs (mx.compile tracing, kernel
|
||||
# warm-up); steady-state throughput is the median of the rest.
|
||||
"denoise_first_step_s": step_times[0] if step_times else None,
|
||||
"denoise_steady_step_s": (float(np.median(step_times[1:])) if len(step_times) > 1 else None),
|
||||
"decode_s": decode_s,
|
||||
"total_s": load_s + denoise_s + decode_s,
|
||||
"load_peak_gib": load_peak / (1024**3),
|
||||
"peak_gib": max(load_peak, denoise_peak) / (1024**3),
|
||||
"quantization": quantization or "none",
|
||||
"compute_dtype": base_dtype,
|
||||
"compile": os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1",
|
||||
"fast_norm": os.environ.get("FASTVIDEO_MLX_FAST_NORM", "0") == "1",
|
||||
"mlx_memory_limit_gib": args.mlx_memory_limit_gib,
|
||||
"mlx_cache_limit_gib": args.mlx_cache_limit_gib,
|
||||
"mlx_disable_cache": args.mlx_disable_cache,
|
||||
"mlx_wired_limit_gib": args.mlx_wired_limit_gib,
|
||||
"torch_mps_high_watermark_ratio": args.torch_mps_high_watermark_ratio,
|
||||
"torch_mps_low_watermark_ratio": args.torch_mps_low_watermark_ratio,
|
||||
}
|
||||
return Cell(
|
||||
prompt_id=args.current_prompt_id,
|
||||
prompt=args.current_prompt,
|
||||
mode=mode,
|
||||
decoder=decoder,
|
||||
video_path=video_path,
|
||||
latents=latents_np,
|
||||
metrics=metrics,
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
preset_parser = argparse.ArgumentParser(add_help=False)
|
||||
preset_parser.add_argument("--benchmark-preset", choices=tuple(BENCHMARK_PRESETS), default="default")
|
||||
preset_args, _ = preset_parser.parse_known_args()
|
||||
preset = BENCHMARK_PRESETS[preset_args.benchmark_preset]
|
||||
|
||||
parser = argparse.ArgumentParser(description="MLX FastWan prove-out benchmark (latency + quality).")
|
||||
parser.add_argument("--benchmark-preset",
|
||||
choices=tuple(BENCHMARK_PRESETS),
|
||||
default=preset_args.benchmark_preset,
|
||||
help="Memory-tier benchmark defaults. Explicit CLI flags override preset values.")
|
||||
parser.add_argument("--model-root", type=Path, default=DEFAULT_MODEL_ROOT)
|
||||
parser.add_argument("--prompt", default="A paper boat sails through a shallow stream in a mossy forest.")
|
||||
parser.add_argument(
|
||||
"--prompt-file",
|
||||
type=Path,
|
||||
default=None,
|
||||
help=
|
||||
"Optional prompt set. Plain text uses one prompt per non-empty line; .jsonl accepts prompt/text/caption plus optional id/name.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-set",
|
||||
choices=("single", *PROMPT_SETS.keys()),
|
||||
default="single",
|
||||
help="Built-in standard prompt set. Ignored when --prompt-file is supplied.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=preset.height)
|
||||
parser.add_argument("--width", type=int, default=preset.width)
|
||||
parser.add_argument("--num-frames", type=int, default=preset.num_frames)
|
||||
parser.add_argument("--dmd-denoising-steps", default="1000,757,522")
|
||||
parser.add_argument("--flow-shift", type=float, default=8.0)
|
||||
parser.add_argument("--max-sequence-length", type=int, default=512)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--modes", default=preset.modes)
|
||||
parser.add_argument("--decoders", default=preset.decoders)
|
||||
parser.add_argument("--output-dir", type=Path, default=Path("video_samples/mlx_fastwan_bench"))
|
||||
parser.add_argument("--torch-device", default="auto")
|
||||
parser.add_argument("--torch-dtype", choices=("fp16", "fp32"), default="fp16")
|
||||
parser.add_argument(
|
||||
"--reference",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="External reference mp4 to score every cell against. Defaults to the fp16+wan-vae cell.",
|
||||
)
|
||||
parser.add_argument("--assert-min-ssim",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Fail if any cell's MS-SSIM vs the reference falls below this value.")
|
||||
parser.add_argument("--compile",
|
||||
action="store_true",
|
||||
help="Enable mx.compile on the DiT forward (sets FASTVIDEO_MLX_COMPILE=1).")
|
||||
parser.add_argument("--lpips", action="store_true", help="Also compute LPIPS (needs the `lpips` package).")
|
||||
parser.add_argument("--taehv-source-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-checkpoint-path", type=Path, default=None)
|
||||
parser.add_argument("--taehv-parallel", action="store_true")
|
||||
parser.add_argument(
|
||||
"--mlx-checkpoint-cache",
|
||||
type=Path,
|
||||
default=None,
|
||||
help="Directory of per-mode pre-quantized MLX checkpoints. The first cell of a mode "
|
||||
"converts from Diffusers weights and saves here (load_source=diffusers_then_saved); "
|
||||
"later cells and later runs reload without requantizing (load_source=mlx_checkpoint), "
|
||||
"which is also how the checkpoint load-time win is measured.",
|
||||
)
|
||||
add_memory_limit_args(
|
||||
parser,
|
||||
mlx_memory_limit_gib=preset.mlx_memory_limit_gib,
|
||||
mlx_disable_cache=preset.mlx_disable_cache,
|
||||
torch_mps_high_watermark_ratio=preset.torch_mps_high_watermark_ratio,
|
||||
torch_mps_low_watermark_ratio=preset.torch_mps_low_watermark_ratio,
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.compile:
|
||||
os.environ["FASTVIDEO_MLX_COMPILE"] = "1"
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
runtime_limits = apply_memory_limits(
|
||||
mlx_memory_limit_gib=args.mlx_memory_limit_gib,
|
||||
mlx_cache_limit_gib=args.mlx_cache_limit_gib,
|
||||
mlx_disable_cache=args.mlx_disable_cache,
|
||||
mlx_wired_limit_gib=args.mlx_wired_limit_gib,
|
||||
torch_mps_high_watermark_ratio=args.torch_mps_high_watermark_ratio,
|
||||
torch_mps_low_watermark_ratio=args.torch_mps_low_watermark_ratio,
|
||||
mx_module=mx,
|
||||
).as_metrics()
|
||||
import torch
|
||||
|
||||
mx.random.seed(args.seed)
|
||||
torch.manual_seed(args.seed)
|
||||
args.output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
modes = _parse_list(args.modes, ALLOWED_MODES, "modes")
|
||||
decoders = _parse_list(args.decoders, ALLOWED_DECODERS, "decoders")
|
||||
|
||||
config_path = args.model_root / "transformer/config.json"
|
||||
checkpoint_path = args.model_root / "transformer/diffusion_pytorch_model.safetensors"
|
||||
config = json.loads(config_path.read_text())
|
||||
latent_frames = (args.num_frames - 1) // 4 + 1
|
||||
latent_height = args.height // 8
|
||||
latent_width = args.width // 8
|
||||
freqs_cis = make_rotary_embeddings(
|
||||
config,
|
||||
latent_frames=latent_frames,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
)
|
||||
generator = torch.Generator(device="cpu").manual_seed(args.seed)
|
||||
latents_seed = torch.randn(
|
||||
(1, int(config["in_channels"]), latent_frames, latent_height, latent_width),
|
||||
generator=generator,
|
||||
dtype=torch.float32,
|
||||
).numpy()
|
||||
|
||||
timesteps = [int(step.strip()) for step in args.dmd_denoising_steps.split(",") if step.strip()]
|
||||
# Keep DMD stochasticity identical across benchmark cells. Without this,
|
||||
# FP16/INT8/decoder comparisons can accidentally measure different re-noise
|
||||
# samples instead of only quantization or decoder differences.
|
||||
renoise_by_step = [
|
||||
torch.randn(latents_seed.shape, generator=generator, dtype=torch.float32).numpy()
|
||||
for _ in range(max(0,
|
||||
len(timesteps) - 1))
|
||||
]
|
||||
|
||||
from fastvideo.mlx_runtime.fastwan import UnsupportedMLXQuantizationError
|
||||
|
||||
prompt_cases = _load_prompt_cases(args.prompt, args.prompt_file, args.prompt_set)
|
||||
cells: list[Cell] = []
|
||||
unsupported_rows: list[dict] = []
|
||||
for prompt_case in prompt_cases:
|
||||
args.current_prompt_id = prompt_case.id
|
||||
args.current_prompt = prompt_case.prompt
|
||||
prompt_embeds = encode_prompt(
|
||||
model_root=args.model_root,
|
||||
prompt=prompt_case.prompt,
|
||||
max_sequence_length=args.max_sequence_length,
|
||||
device_arg=args.torch_device,
|
||||
dtype_arg=args.torch_dtype,
|
||||
)
|
||||
encoder_hidden_states = mx.array(prompt_embeds.numpy())
|
||||
|
||||
for mode in modes:
|
||||
for decoder in decoders:
|
||||
print(f"=== cell: prompt={prompt_case.id} mode={mode} decoder={decoder} ===")
|
||||
try:
|
||||
cells.append(
|
||||
_generate_cell(
|
||||
args=args,
|
||||
mode=mode,
|
||||
decoder=decoder,
|
||||
checkpoint_path=checkpoint_path,
|
||||
config_path=config_path,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
freqs_cis=freqs_cis,
|
||||
timesteps=timesteps,
|
||||
latents_seed=latents_seed,
|
||||
renoise_by_step=renoise_by_step,
|
||||
))
|
||||
except UnsupportedMLXQuantizationError as exc:
|
||||
# Record the cell as unsupported and keep sweeping: a partial
|
||||
# report on this MLX build beats crashing the whole run.
|
||||
print(f"skipping cell (unsupported by this MLX build): {exc}")
|
||||
unsupported_rows.append({
|
||||
"prompt_id": prompt_case.id,
|
||||
"prompt": prompt_case.prompt,
|
||||
"mode": mode,
|
||||
"decoder": decoder,
|
||||
"status": "unsupported_by_mlx",
|
||||
"error": str(exc),
|
||||
})
|
||||
|
||||
if not cells:
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(unsupported_rows, indent=2))
|
||||
raise SystemExit(f"No benchmark cell could run: every requested mode is unsupported by this MLX build. "
|
||||
f"Wrote {metrics_path}.")
|
||||
|
||||
# Resolve one internal reference per prompt. A single external reference, if
|
||||
# supplied, is used for every prompt and only video metrics are computed.
|
||||
reference_by_prompt: dict[str, tuple[Path, np.ndarray | None]] = {}
|
||||
if args.reference is not None:
|
||||
for prompt_case in prompt_cases:
|
||||
reference_by_prompt[prompt_case.id] = (args.reference, None)
|
||||
else:
|
||||
for prompt_case in prompt_cases:
|
||||
prompt_cells = [c for c in cells if c.prompt_id == prompt_case.id]
|
||||
if not prompt_cells:
|
||||
continue
|
||||
ref_cell = next(
|
||||
(c for c in prompt_cells if c.mode == REFERENCE_MODE and c.decoder == REFERENCE_DECODER),
|
||||
prompt_cells[0],
|
||||
)
|
||||
reference_by_prompt[prompt_case.id] = (ref_cell.video_path, ref_cell.latents)
|
||||
print(f"Using internal reference cell for {prompt_case.id}: "
|
||||
f"mode={ref_cell.mode} decoder={ref_cell.decoder}")
|
||||
|
||||
lpips_fn = _load_lpips() if args.lpips else None
|
||||
|
||||
rows: list[dict] = []
|
||||
failures: list[str] = []
|
||||
for cell in cells:
|
||||
reference_video, reference_latents = reference_by_prompt[cell.prompt_id]
|
||||
ms_ssim = _ms_ssim(Path(reference_video), cell.video_path, required=args.assert_min_ssim is not None)
|
||||
cell.metrics["ms_ssim_vs_ref"] = ms_ssim
|
||||
cell.metrics.update(runtime_limits)
|
||||
if reference_latents is not None:
|
||||
cell.metrics.update(_latent_delta_metrics(cell.latents, reference_latents))
|
||||
cell.metrics["lpips_vs_ref"] = (_lpips_between(lpips_fn, Path(reference_video), cell.video_path)
|
||||
if lpips_fn else None)
|
||||
if args.assert_min_ssim is not None and ms_ssim is not None and ms_ssim < args.assert_min_ssim:
|
||||
failures.append(
|
||||
f"{cell.prompt_id}/{cell.mode}/{cell.decoder}: MS-SSIM {ms_ssim:.4f} < {args.assert_min_ssim}")
|
||||
rows.append(dict(cell.metrics))
|
||||
print(json.dumps(cell.metrics, indent=2))
|
||||
rows.extend(unsupported_rows)
|
||||
|
||||
metrics_path = args.output_dir / "metrics.json"
|
||||
metrics_path.write_text(json.dumps(rows, indent=2))
|
||||
table_path = args.output_dir / "metrics.md"
|
||||
table = _markdown_table(rows)
|
||||
table_path.write_text(table + "\n")
|
||||
html_path = _write_html_grid(rows, args.output_dir)
|
||||
print("\n" + table)
|
||||
print(f"\nWrote {metrics_path}, {table_path}, and {html_path}")
|
||||
|
||||
if failures:
|
||||
raise SystemExit("SSIM regression gate failed:\n " + "\n ".join(failures))
|
||||
|
||||
|
||||
def _load_lpips() -> object | None:
|
||||
"""Return an LPIPS model, or ``None`` if the optional dep is unavailable."""
|
||||
try:
|
||||
import lpips # noqa: PLC0415 - optional dependency.
|
||||
except ImportError:
|
||||
print("LPIPS requested but the `lpips` package is not installed; skipping (install `.[eval]`).")
|
||||
return None
|
||||
return lpips.LPIPS(net="alex")
|
||||
|
||||
|
||||
def _lpips_between(lpips_fn, reference_video: Path, candidate_video: Path) -> float | None:
|
||||
if lpips_fn is None or not reference_video.exists() or not candidate_video.exists():
|
||||
return None
|
||||
import torch
|
||||
|
||||
ref = _read_video_frames(reference_video)
|
||||
cand = _read_video_frames(candidate_video)
|
||||
if ref is None or cand is None or ref.shape != cand.shape:
|
||||
return None
|
||||
# LPIPS expects NCHW in [-1, 1].
|
||||
ref_t = torch.from_numpy(ref).permute(0, 3, 1, 2).float() / 127.5 - 1.0
|
||||
cand_t = torch.from_numpy(cand).permute(0, 3, 1, 2).float() / 127.5 - 1.0
|
||||
with torch.no_grad():
|
||||
scores = lpips_fn(ref_t, cand_t)
|
||||
return float(scores.mean().item())
|
||||
|
||||
|
||||
def _read_video_frames(path: Path) -> np.ndarray | None:
|
||||
try:
|
||||
import cv2
|
||||
except ImportError:
|
||||
return None
|
||||
cap = cv2.VideoCapture(str(path))
|
||||
frames = []
|
||||
try:
|
||||
while True:
|
||||
ok, frame_bgr = cap.read()
|
||||
if not ok:
|
||||
break
|
||||
frames.append(cv2.cvtColor(frame_bgr, cv2.COLOR_BGR2RGB))
|
||||
finally:
|
||||
cap.release()
|
||||
if not frames:
|
||||
return None
|
||||
return np.stack(frames, axis=0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -2,13 +2,7 @@ from fastvideo.configs.models.base import ModelConfig
|
||||
from fastvideo.configs.models.dits.base import DiTConfig
|
||||
from fastvideo.configs.models.encoders.base import EncoderConfig
|
||||
from fastvideo.configs.models.vaes.base import VAEConfig
|
||||
from fastvideo.configs.models.audio import (
|
||||
BigVGANV2Config,
|
||||
LTX2AudioDecoderConfig,
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
MMAudioVAEConfig,
|
||||
)
|
||||
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig, LTX2AudioEncoderConfig, LTX2VocoderConfig)
|
||||
from fastvideo.configs.models.upsamplers.base import UpsamplerConfig
|
||||
|
||||
__all__ = [
|
||||
@@ -19,7 +13,5 @@ __all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
"MMAudioVAEConfig",
|
||||
"BigVGANV2Config",
|
||||
"UpsamplerConfig",
|
||||
]
|
||||
|
||||
@@ -5,21 +5,9 @@ from fastvideo.configs.models.audio.ltx2_audio_vae import (
|
||||
LTX2AudioEncoderConfig,
|
||||
LTX2VocoderConfig,
|
||||
)
|
||||
from fastvideo.configs.models.audio.mmaudio_vae import (
|
||||
MMAudioVAEArchConfig,
|
||||
MMAudioVAEConfig,
|
||||
)
|
||||
from fastvideo.configs.models.audio.bigvgan import (
|
||||
BigVGANV2ArchConfig,
|
||||
BigVGANV2Config,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTX2AudioEncoderConfig",
|
||||
"LTX2AudioDecoderConfig",
|
||||
"LTX2VocoderConfig",
|
||||
"MMAudioVAEArchConfig",
|
||||
"MMAudioVAEConfig",
|
||||
"BigVGANV2ArchConfig",
|
||||
"BigVGANV2Config",
|
||||
]
|
||||
|
||||
@@ -1,18 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Reusable BigVGAN-v2 vocoder configuration."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class BigVGANV2ArchConfig(ArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["BigVGANV2"])
|
||||
sample_rate: int = 44100
|
||||
num_mels: int = 128
|
||||
|
||||
|
||||
@dataclass
|
||||
class BigVGANV2Config(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=BigVGANV2ArchConfig)
|
||||
@@ -1,21 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MMAudio audio VAE configuration."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig, ModelConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioVAEArchConfig(ArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["MMAudioVAE"])
|
||||
mode: str = "44k"
|
||||
data_dim: int = 128
|
||||
embed_dim: int = 40
|
||||
hidden_dim: int = 512
|
||||
need_encoder: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioVAEConfig(ModelConfig):
|
||||
arch_config: ArchConfig = field(default_factory=MMAudioVAEArchConfig)
|
||||
@@ -11,7 +11,6 @@ from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.magi_human import MagiHumanVideoConfig
|
||||
from fastvideo.configs.models.dits.minimax_h3 import MiniMaxH3Config
|
||||
from fastvideo.configs.models.dits.mmaudio import MMAudioArchConfig, MMAudioTransformerConfig
|
||||
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
|
||||
@@ -25,5 +24,5 @@ __all__ = [
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
|
||||
"MiniMaxH3Config", "ZImageDiTConfig", "MMAudioArchConfig", "MMAudioTransformerConfig"
|
||||
"MiniMaxH3Config", "ZImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -1,49 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_mmaudio_transformer_block(name: str, module) -> bool:
|
||||
del module
|
||||
parts = name.split(".")
|
||||
return len(parts) >= 2 and parts[-1].isdigit() and parts[-2] in {"joint_blocks", "fused_blocks"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_mmaudio_transformer_block])
|
||||
param_names_mapping: dict = field(default_factory=lambda: {r"^(.*)$": r"\1"})
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.TORCH_SDPA, )
|
||||
|
||||
latent_dim: int = 40
|
||||
clip_dim: int = 1024
|
||||
sync_dim: int = 768
|
||||
text_dim: int = 1024
|
||||
hidden_dim: int = 896
|
||||
depth: int = 21
|
||||
fused_depth: int = 14
|
||||
num_heads: int = 14
|
||||
mlp_ratio: float = 4.0
|
||||
latent_seq_len: int = 345
|
||||
clip_seq_len: int = 64
|
||||
sync_seq_len: int = 192
|
||||
text_seq_len: int = 77
|
||||
v2: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.hidden_size = self.hidden_dim
|
||||
self.num_attention_heads = self.num_heads
|
||||
self.num_channels_latents = self.latent_dim
|
||||
self.in_channels = self.latent_dim
|
||||
self.out_channels = self.latent_dim
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioTransformerConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=MMAudioArchConfig)
|
||||
prefix: str = "MMAudio"
|
||||
@@ -1,10 +1,6 @@
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
BaseEncoderOutput,
|
||||
EncoderConfig,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
from fastvideo.configs.models.encoders.clip import CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig
|
||||
from fastvideo.configs.models.encoders.base import (BaseEncoderOutput, EncoderConfig, ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
|
||||
from fastvideo.configs.models.encoders.llama import LlamaConfig
|
||||
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5ArchConfig, LingBotWorld2UMT5Config
|
||||
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
|
||||
@@ -16,49 +12,15 @@ from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
|
||||
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import (MiniMaxH3Qwen3VLArchConfig, MiniMaxH3Qwen3VLConfig)
|
||||
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
|
||||
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (
|
||||
StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig,
|
||||
)
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig)
|
||||
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
|
||||
from fastvideo.configs.models.encoders.mmaudio_synchformer import MMAudioSynchformerArchConfig, MMAudioSynchformerConfig
|
||||
from fastvideo.configs.models.encoders.mmaudio_clip import (
|
||||
MMAudioDFNCLIPTextArchConfig,
|
||||
MMAudioDFNCLIPTextConfig,
|
||||
MMAudioDFNCLIPVisionArchConfig,
|
||||
MMAudioDFNCLIPVisionConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig",
|
||||
"TextEncoderConfig",
|
||||
"ImageEncoderConfig",
|
||||
"BaseEncoderOutput",
|
||||
"CLIPTextConfig",
|
||||
"CLIPVisionConfig",
|
||||
"WAN2_1ControlCLIPVisionConfig",
|
||||
"LlamaConfig",
|
||||
"T5Config",
|
||||
"T5LargeConfig",
|
||||
"Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig",
|
||||
"Reason1Config",
|
||||
"LTX2GemmaConfig",
|
||||
"SiglipVisionConfig",
|
||||
"StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig",
|
||||
"T5GemmaEncoderConfig",
|
||||
"Qwen3TextConfig",
|
||||
"Mistral3TextConfig",
|
||||
"LingBotWorld2UMT5ArchConfig",
|
||||
"LingBotWorld2UMT5Config",
|
||||
"LingBotVideoQwen3VLTextConfig",
|
||||
"MiniMaxH3Qwen3VLArchConfig",
|
||||
"MiniMaxH3Qwen3VLConfig",
|
||||
"MMAudioSynchformerArchConfig",
|
||||
"MMAudioSynchformerConfig",
|
||||
"MMAudioDFNCLIPTextArchConfig",
|
||||
"MMAudioDFNCLIPTextConfig",
|
||||
"MMAudioDFNCLIPVisionArchConfig",
|
||||
"MMAudioDFNCLIPVisionConfig",
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
|
||||
"LingBotWorld2UMT5ArchConfig", "LingBotWorld2UMT5Config", "LingBotVideoQwen3VLTextConfig",
|
||||
"MiniMaxH3Qwen3VLArchConfig", "MiniMaxH3Qwen3VLConfig"
|
||||
]
|
||||
|
||||
@@ -62,8 +62,6 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
hidden_size: int = 5120
|
||||
intermediate_size: int = 25600
|
||||
num_hidden_layers: int = 64
|
||||
output_hidden_state_index: int = 50
|
||||
num_hidden_layers_override: int | None = 50
|
||||
num_attention_heads: int = 64
|
||||
num_key_value_heads: int = 8
|
||||
head_dim: int = 128
|
||||
@@ -109,7 +107,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
vision_initializer_range: float = 0.02
|
||||
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
|
||||
|
||||
output_hidden_states: bool = False
|
||||
output_hidden_states: bool = True
|
||||
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
_is_language_transformer_layer,
|
||||
@@ -120,17 +118,6 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
|
||||
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
|
||||
if self.num_hidden_layers_override is not None:
|
||||
if self.num_hidden_layers_override <= 0:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
|
||||
if self.num_hidden_layers_override < self.output_hidden_state_index:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
|
||||
f"hidden_states[{self.output_hidden_state_index}], got "
|
||||
f"{self.num_hidden_layers_override}.")
|
||||
|
||||
rope_scaling = dict(self.rope_scaling or {})
|
||||
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
|
||||
if not self.mrope_interleaved:
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Native DFN5B CLIP conditioner configurations for MMAudio."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.clip import (
|
||||
CLIPTextArchConfig,
|
||||
CLIPTextConfig,
|
||||
CLIPVisionArchConfig,
|
||||
CLIPVisionConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioDFNCLIPTextArchConfig(CLIPTextArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["MMAudioDFNCLIPTextEncoder"])
|
||||
vocab_size: int = 49408
|
||||
hidden_size: int = 1024
|
||||
intermediate_size: int = 4096
|
||||
projection_dim: int = 1024
|
||||
num_hidden_layers: int = 24
|
||||
num_attention_heads: int = 16
|
||||
max_position_embeddings: int = 77
|
||||
text_len: int = 77
|
||||
hidden_act: str = "quick_gelu"
|
||||
layer_norm_eps: float = 1e-5
|
||||
pad_token_id: int = 0
|
||||
bos_token_id: int = 49406
|
||||
eos_token_id: int = 49407
|
||||
# MMAudio moves this encoder as one unit between CPU and GPU. Keeping it a
|
||||
# plain module avoids nesting FSDP CPU-offload semantics inside the custom
|
||||
# OpenCLIP causal-mask forward used by this pipeline.
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioDFNCLIPTextConfig(CLIPTextConfig):
|
||||
arch_config: CLIPTextArchConfig = field(default_factory=MMAudioDFNCLIPTextArchConfig)
|
||||
# OpenCLIP supplies an explicit additive triangular mask to
|
||||
# nn.MultiheadAttention. The MMAudio adapter reproduces that path instead
|
||||
# of using SDPA's is_causal shortcut, which rounds differently in bf16.
|
||||
is_causal: bool = False
|
||||
prefix: str = "mmaudio_dfn_clip_text"
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioDFNCLIPVisionArchConfig(CLIPVisionArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["MMAudioDFNCLIPVisionEncoder"])
|
||||
hidden_size: int = 1280
|
||||
intermediate_size: int = 5120
|
||||
projection_dim: int = 1024
|
||||
num_hidden_layers: int = 32
|
||||
num_attention_heads: int = 16
|
||||
image_size: int = 378
|
||||
patch_size: int = 14
|
||||
hidden_act: str = "quick_gelu"
|
||||
layer_norm_eps: float = 1e-5
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioDFNCLIPVisionConfig(CLIPVisionConfig):
|
||||
arch_config: CLIPVisionArchConfig = field(default_factory=MMAudioDFNCLIPVisionArchConfig)
|
||||
num_hidden_layers_override: int | None = None
|
||||
require_post_norm: bool | None = True
|
||||
enable_scale: bool = True
|
||||
is_causal: bool = False
|
||||
prefix: str = "mmaudio_dfn_clip_vision"
|
||||
@@ -1,26 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Configuration for MMAudio's Synchformer visual conditioner."""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioSynchformerArchConfig(ImageEncoderArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["MMAudioSynchformerVisualEncoder"])
|
||||
image_size: int = 224
|
||||
num_channels: int = 3
|
||||
segment_size: int = 16
|
||||
segment_stride: int = 8
|
||||
hidden_size: int = 768
|
||||
tokens_per_segment: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioSynchformerConfig(ImageEncoderConfig):
|
||||
arch_config: ImageEncoderArchConfig = field(default_factory=MMAudioSynchformerArchConfig)
|
||||
prefix: str = "synchformer"
|
||||
@@ -11,7 +11,6 @@ from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V
|
||||
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.configs.pipelines.mmaudio import MMAudioV2AConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.wan import (LucyEditDevConfig, SelfForcingWanT2V480PConfig, WanI2V480PConfig,
|
||||
@@ -23,5 +22,5 @@ __all__ = [
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
|
||||
"Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig",
|
||||
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "MMAudioV2AConfig", "get_pipeline_config_cls_from_name"
|
||||
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -63,11 +63,6 @@ class PipelineConfig:
|
||||
# Image encoder configuration
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
# Optional multi-encoder contract. Existing pipelines continue to use the
|
||||
# singular fields above; V2A and other multimodal pipelines can opt into
|
||||
# indexed ``image_encoder``, ``image_encoder_2``, ... components.
|
||||
image_encoder_configs: tuple[EncoderConfig, ...] | None = None
|
||||
image_encoder_precisions: tuple[str, ...] | None = None
|
||||
|
||||
# Text encoder configuration
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp32", )
|
||||
|
||||
@@ -1,66 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pipeline configuration for the native MMAudio video-to-audio port."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, ModelConfig
|
||||
from fastvideo.configs.models.audio import BigVGANV2Config, MMAudioVAEConfig
|
||||
from fastvideo.configs.models.dits import MMAudioTransformerConfig
|
||||
from fastvideo.configs.models.encoders import (
|
||||
MMAudioDFNCLIPTextConfig,
|
||||
MMAudioDFNCLIPVisionConfig,
|
||||
MMAudioSynchformerConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class MMAudioV2AConfig(PipelineConfig):
|
||||
"""MMAudio large-44k-v2 inference defaults.
|
||||
|
||||
The published demo moves every module to bfloat16. Keeping the same
|
||||
per-component precision here is important: condition features seed the
|
||||
complete flow trajectory, so silently encoding them in fp32 changes the
|
||||
generated waveform even when the transformer weights are identical.
|
||||
"""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=MMAudioTransformerConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (MMAudioDFNCLIPTextConfig(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
image_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (
|
||||
MMAudioDFNCLIPVisionConfig(),
|
||||
MMAudioSynchformerConfig(),
|
||||
))
|
||||
image_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
|
||||
|
||||
audio_decoder_config: ModelConfig = field(default_factory=MMAudioVAEConfig)
|
||||
audio_decoder_precision: str = "bf16"
|
||||
vocoder_config: ModelConfig = field(default_factory=BigVGANV2Config)
|
||||
vocoder_precision: str = "bf16"
|
||||
|
||||
# Published large_44k_v2 default sequence contract. The official demo
|
||||
# supports other durations, although quality can drop far away from the
|
||||
# eight-second training duration.
|
||||
duration_s: float = 8.0
|
||||
max_audio_duration_s: float | None = None
|
||||
sampling_rate: int = 44100
|
||||
spectrogram_frame_rate: int = 512
|
||||
latent_downsample_rate: int = 2
|
||||
clip_frame_rate: int = 8
|
||||
sync_frame_rate: int = 25
|
||||
sync_segment_size: int = 16
|
||||
sync_segment_stride: int = 8
|
||||
sync_downsample_rate: int = 2
|
||||
clip_image_size: int = 384
|
||||
sync_image_size: int = 224
|
||||
clip_batch_size_multiplier: int = 40
|
||||
sync_batch_size_multiplier: int = 40
|
||||
|
||||
num_inference_steps: int = 25
|
||||
guidance_scale: float = 4.5
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
@@ -50,7 +50,7 @@ from fastvideo.api.schema import (
|
||||
SamplingConfig,
|
||||
)
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines import ForwardBatch
|
||||
from fastvideo.utils import align_to, shallow_asdict
|
||||
@@ -605,13 +605,7 @@ class VideoGenerator:
|
||||
|
||||
# Single prompt generation (original behavior)
|
||||
if prompt is None:
|
||||
if fastvideo_args.workload_type is WorkloadType.V2A:
|
||||
# Video semantics are sufficient conditioning for V2A models;
|
||||
# model-specific text stages interpret the empty string using
|
||||
# their native tokenizer/empty-prompt contract.
|
||||
prompt = ""
|
||||
else:
|
||||
raise ValueError("Either prompt or prompt_txt must be provided")
|
||||
raise ValueError("Either prompt or prompt_txt must be provided")
|
||||
output_path = self._prepare_output_path(sampling_param.output_path, prompt)
|
||||
kwargs["output_path"] = output_path
|
||||
if prompt_embeds is not None:
|
||||
@@ -630,13 +624,6 @@ class VideoGenerator:
|
||||
return False
|
||||
return args.workload_type.value.endswith("2i")
|
||||
|
||||
def _is_audio_workload(self) -> bool:
|
||||
"""Return True when the workload produces standalone audio."""
|
||||
args = getattr(self, "fastvideo_args", None)
|
||||
if args is None:
|
||||
return False
|
||||
return args.workload_type.value.endswith("2a")
|
||||
|
||||
def _prepare_output_path(
|
||||
self,
|
||||
output_path: str,
|
||||
@@ -656,12 +643,7 @@ class VideoGenerator:
|
||||
warning is logged.
|
||||
- If the target path already exists, a numeric suffix is appended.
|
||||
"""
|
||||
if self._is_image_workload():
|
||||
target_ext = ".png"
|
||||
elif self._is_audio_workload():
|
||||
target_ext = ".wav"
|
||||
else:
|
||||
target_ext = ".mp4"
|
||||
target_ext = ".png" if self._is_image_workload() else ".mp4"
|
||||
|
||||
def _sanitize_filename_component(name: str) -> str:
|
||||
# Remove characters invalid on common filesystems, strip spaces/dots
|
||||
@@ -808,19 +790,14 @@ class VideoGenerator:
|
||||
latent_batch_size = _infer_latent_batch_size(batch)
|
||||
is_latent_output = fastvideo_args.output_type == "latent"
|
||||
needs_frame_output = batch.return_frames or (batch.save_video and not is_latent_output)
|
||||
# A populated ``samples`` has exactly one consumer — the result
|
||||
# dict (``"samples": samples if batch.return_frames else None``).
|
||||
# Post-decode frame building reads ``output_batch.output``
|
||||
# directly (the GPU ``vid_u8`` path), not ``samples``. So when
|
||||
# ``return_frames=False`` the pinned fp32 alloc + D->H copy are
|
||||
# dead weight — the CLI generate flow (``save_video=True``,
|
||||
# ``return_frames=False``) hits this on every call.
|
||||
# ``output_type == "latent"`` keeps its existing branch (shape
|
||||
# mismatch falls through to ``.cpu()`` below) for callers that
|
||||
# *do* ask for the latent samples via ``return_frames=True``.
|
||||
needs_samples_buffer = batch.return_frames or needs_frame_output
|
||||
# When ``output_type == "latent"`` the forward output has latent
|
||||
# shape (e.g. ``[B, C_latent, T_latent, H_latent, W_latent]``)
|
||||
# rather than the pre-allocation's pixel shape. Skip the pinned
|
||||
# ~50 MB buffer entirely. Also skip it for metadata-only calls;
|
||||
# neither the result nor save path will consume the decoded tensor.
|
||||
# ``skip_pixel_prealloc`` also gates the slow-path warning.
|
||||
needs_samples_out = batch.return_frames
|
||||
skip_pixel_prealloc = is_latent_output or not needs_samples_out
|
||||
skip_pixel_prealloc = is_latent_output or not needs_samples_buffer
|
||||
if skip_pixel_prealloc:
|
||||
samples = torch.empty(0, device='cpu')
|
||||
else:
|
||||
@@ -840,11 +817,9 @@ class VideoGenerator:
|
||||
"This usually means the executor/pipeline failed earlier.")
|
||||
|
||||
audio_only = bool(output_batch.extra.get("audio_only"))
|
||||
if not needs_samples_out:
|
||||
# Nothing downstream reads ``samples`` (the result dict
|
||||
# returns None when ``return_frames=False``); keep the empty
|
||||
# placeholder allocated above and skip the fp32 D->H copy
|
||||
# entirely.
|
||||
if not needs_samples_buffer or (audio_only and not batch.return_frames):
|
||||
# Metadata-only/audio-only request: keep the empty placeholder and
|
||||
# avoid the decoded tensor D->H copy.
|
||||
pass
|
||||
elif audio_only:
|
||||
# Audio-only return-frames requests expose the small placeholder
|
||||
@@ -876,13 +851,8 @@ class VideoGenerator:
|
||||
# `GenerationResult.size` describes the produced media, not only the
|
||||
# base-stage request. Refiner pipelines can change the final pixel
|
||||
# dimensions, so derive this result metadata from the decoded output.
|
||||
# Read the geometry from `output_batch.output` (a shape-only access,
|
||||
# no D->H copy): when `return_frames=False` the `samples` mirror
|
||||
# stays an empty placeholder and no longer carries the decoded
|
||||
# shape. Metadata-only calls keep the request fallback and never
|
||||
# inspect the (possibly dropped) worker output.
|
||||
output_size = _resolve_output_size(
|
||||
output_batch.output if needs_frame_output else samples,
|
||||
samples,
|
||||
(target_height, target_width, batch.num_frames),
|
||||
pixel_output=not is_latent_output and not audio_only,
|
||||
)
|
||||
@@ -894,26 +864,13 @@ class VideoGenerator:
|
||||
elif not needs_frame_output:
|
||||
frames = None
|
||||
else:
|
||||
# Quantize on the source device (typically CUDA) BEFORE the
|
||||
# device->host copy. `samples` above is just the pinned-CPU
|
||||
# mirror of `output_batch.output` (`samples.copy_(output)` or
|
||||
# `output.cpu()`) with no intervening preprocessing, so reading
|
||||
# `output_batch.output` here is the same data. The old path
|
||||
# paid a full fp32 video D->H copy (which scales with
|
||||
# resolution x frames x batch) and then a single-threaded
|
||||
# per-frame CPU *255/cast loop. Casting to uint8 on-device
|
||||
# makes the transfer 4x smaller, ships it in a single copy,
|
||||
# and moves the elementwise work onto the GPU. clamp_() also
|
||||
# fixes a latent overflow bug: VAE output slightly outside
|
||||
# [0, 1] wrapped mod 256 in the old unclamped cast.
|
||||
# (Equivalence is SSIM-gated, not bit-exact: float->uint8
|
||||
# differs <=1 LSB CPU vs GPU.)
|
||||
src = output_batch.output
|
||||
vid_u8 = (src * 255).clamp_(0, 255).to(torch.uint8)
|
||||
vid_u8 = rearrange(vid_u8, "b c t h w -> t b c h w").cpu()
|
||||
frames = [
|
||||
torchvision.utils.make_grid(x, nrow=6).permute(1, 2, 0).squeeze(-1).contiguous().numpy() for x in vid_u8
|
||||
]
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.permute(1, 2, 0).squeeze(-1)
|
||||
x = (x * 255).to(torch.uint8)
|
||||
frames.append(x.contiguous().cpu().numpy())
|
||||
postprocess_time = time.perf_counter() - postprocess_start
|
||||
logger.info("PostDecodeFrameProcessStage completed in %.3f s", postprocess_time)
|
||||
if logging_info is not None:
|
||||
|
||||
+14
-34
@@ -21,22 +21,20 @@ 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
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_WAIT_STEPS: int = 2
|
||||
FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS: int = 1
|
||||
FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS: int = 2
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
|
||||
FASTVIDEO_TRACE_ACTIVATIONS: bool = False
|
||||
FASTVIDEO_TRACE_LAYERS: str = ""
|
||||
@@ -222,34 +220,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# If set (=1), MiniMax-H3 VAE decode (and, with the ENCODE variant,
|
||||
# reference-video encode) round-robins its temporal chunks across the
|
||||
# sequence-parallel ranks instead of running serially on the output rank.
|
||||
# Folded into FastVideoArgs.vae_parallel_decode / vae_parallel_encode at
|
||||
# construction (parse-once). The STRATEGY variant picks the chunk
|
||||
# transport collective: "gather" (default) or "all_gather".
|
||||
"FASTVIDEO_VAE_PARALLEL_DECODE":
|
||||
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE", "0") != "0",
|
||||
"FASTVIDEO_VAE_PARALLEL_ENCODE":
|
||||
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_ENCODE", "0") != "0",
|
||||
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY":
|
||||
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY", None),
|
||||
|
||||
# Opt-in MiniMax-H3 inference-only Triton fusions adapted from the
|
||||
# NVlabs/Sana Sol-Engine implementation. Accepts `all`, `1`, or a
|
||||
# comma-separated subset of `modulate,qknorm_rope,swiglu`. An empty value
|
||||
# (the default), `0`, or `none` keeps the eager implementation.
|
||||
"FASTVIDEO_MINIMAX_H3_FUSIONS":
|
||||
lambda: os.getenv("FASTVIDEO_MINIMAX_H3_FUSIONS", ""),
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
# Emit lightweight NVTX ranges for external profilers such as Nsight Systems.
|
||||
"FASTVIDEO_NVTX_PROFILE":
|
||||
lambda: os.getenv("FASTVIDEO_NVTX_PROFILE", "0") != "0",
|
||||
|
||||
# Enables torch profiler if set. Path to the directory where torch profiler
|
||||
# traces are saved. Note that it must be an absolute path.
|
||||
"FASTVIDEO_TORCH_PROFILER_DIR":
|
||||
@@ -268,11 +242,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY":
|
||||
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY", "0") != "0"),
|
||||
|
||||
# Enable torch profiler stack capture with
|
||||
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. Off by default: stack capture
|
||||
# costs ~1.5x runtime overhead and ~1.4x trace size.
|
||||
# Enable torch profiler to profile stack if set
|
||||
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. If not set, torch profiler WILL
|
||||
# profile stack by default.
|
||||
"FASTVIDEO_TORCH_PROFILER_WITH_STACK":
|
||||
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "0") != "0"),
|
||||
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "1") != "0"),
|
||||
|
||||
# Enable torch profiler to profile flops if set
|
||||
# FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1. If not set, torch profiler will
|
||||
@@ -281,10 +255,16 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
|
||||
# Wait steps per profiling cycle (torch.profiler.schedule wait parameter)
|
||||
# Defaults to 2 if not set.
|
||||
"FASTVIDEO_TORCH_PROFILER_WAIT_STEPS":
|
||||
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WAIT_STEPS", "2")),
|
||||
# Warmup steps per profiling cycle (torch.profiler.schedule warmup parameter)
|
||||
# Defaults to 1 if not set.
|
||||
"FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS":
|
||||
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS", "1")),
|
||||
# Active steps per profiling cycle (torch.profiler.schedule active parameter)
|
||||
# Defaults to 2 if not set.
|
||||
"FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS":
|
||||
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS", "2")),
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS":
|
||||
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
|
||||
|
||||
|
||||
@@ -61,8 +61,6 @@ class WorkloadType(str, Enum):
|
||||
T2V = "t2v" # Text to Video
|
||||
T2I = "t2i" # Text to Image
|
||||
I2I = "i2i" # Image to Image
|
||||
V2A = "v2a" # Video to Audio
|
||||
T2A = "t2a" # Text to Audio
|
||||
|
||||
@classmethod
|
||||
def from_string(cls, value: str) -> "WorkloadType":
|
||||
@@ -146,19 +144,6 @@ class FastVideoArgs:
|
||||
vae_cpu_offload: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
# Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the
|
||||
# video VAE's temporal chunks (decode) and clips (reference encode) are
|
||||
# round-robined across the sequence-parallel ranks and reassembled
|
||||
# bit-exactly on the group's first rank instead of running serially on
|
||||
# one rank while the others idle. ``__post_init__`` folds the
|
||||
# FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars
|
||||
# into these fields (parse-once, like attention_backend), and
|
||||
# FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport
|
||||
# collective ("gather" or "all_gather").
|
||||
vae_parallel_decode: bool = False
|
||||
vae_parallel_encode: bool = False
|
||||
vae_parallel_decode_strategy: str | None = None
|
||||
|
||||
# Compilation
|
||||
# ``enable_torch_compile`` covers the DiT path (transformer,
|
||||
# transformer_2, and the LTX-2 stage-2 transformer_refine).
|
||||
@@ -182,7 +167,6 @@ class FastVideoArgs:
|
||||
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
VSA_tile_size: int = 256 # VSA-H3 tile size (256 or 64); 64 = native Triton path
|
||||
|
||||
# V-MoBA parameters
|
||||
moba_config_path: str | None = None
|
||||
@@ -300,27 +284,8 @@ class FastVideoArgs:
|
||||
env_backend = envs.FASTVIDEO_ATTENTION_BACKEND
|
||||
if env_backend is not None and backend_name_to_enum(env_backend) is not None:
|
||||
self.attention_backend = env_backend
|
||||
self._fold_vae_parallel_env()
|
||||
self.check_fastvideo_args()
|
||||
|
||||
def _fold_vae_parallel_env(self) -> None:
|
||||
"""Parse-once adapters for the sequence-parallel VAE env vars."""
|
||||
import fastvideo.envs as envs
|
||||
|
||||
# Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES /
|
||||
# DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args
|
||||
# never imports model modules; a unit test pins the two in sync).
|
||||
strategies = ("gather", "all_gather")
|
||||
if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE:
|
||||
self.vae_parallel_decode = True
|
||||
if not self.vae_parallel_encode and envs.FASTVIDEO_VAE_PARALLEL_ENCODE:
|
||||
self.vae_parallel_encode = True
|
||||
if self.vae_parallel_decode_strategy is None:
|
||||
self.vae_parallel_decode_strategy = envs.FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY or "gather"
|
||||
if self.vae_parallel_decode_strategy not in strategies:
|
||||
raise ValueError(f"vae_parallel_decode_strategy must be one of {strategies}, "
|
||||
f"got {self.vae_parallel_decode_strategy!r}.")
|
||||
|
||||
def _apply_transformer_quant(self) -> None:
|
||||
"""Pin the typed ``transformer_quant`` instance onto ``dit_config``.
|
||||
|
||||
@@ -664,18 +629,6 @@ class FastVideoArgs:
|
||||
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
|
||||
"Should be enabled in almost all cases",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-parallel-decode",
|
||||
action=StoreBoolean,
|
||||
help="With sequence parallelism, round-robin MiniMax-H3 VAE decode chunks across the SP ranks "
|
||||
"and reassemble bit-exactly on the output rank (default: serial decode on the output rank)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-parallel-encode",
|
||||
action=StoreBoolean,
|
||||
help="With sequence parallelism, round-robin MiniMax-H3 reference-video VAE encode clips across "
|
||||
"the SP ranks; every rank keeps the identical full encoding (default: serial encode on every rank)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action=StoreBoolean,
|
||||
@@ -689,12 +642,6 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.VSA_sparsity,
|
||||
help="Validation sparsity for VSA",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--VSA-tile-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.VSA_tile_size,
|
||||
help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path",
|
||||
)
|
||||
|
||||
# Master port for distributed training/inference
|
||||
parser.add_argument(
|
||||
|
||||
+1
-4
@@ -114,10 +114,7 @@ def _info(logger: Logger,
|
||||
is_local_main_process = local_rank == 0
|
||||
|
||||
if (main_process_only and is_main_process) or (local_main_process_only and is_local_main_process):
|
||||
# Honor an explicit stacklevel (info_once routes through here with
|
||||
# stacklevel already set) instead of passing the keyword twice.
|
||||
stacklevel = kwargs.pop("stacklevel", 2)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=stacklevel, **kwargs)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
|
||||
@@ -1,115 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Experimental Apple MLX runtime helpers.
|
||||
|
||||
This package is intentionally small for now. It exists to grow the Apple-native
|
||||
FastWan path in measurable steps: shape planning, primitive benchmarks, then
|
||||
Wan block parity, then full DiT/runtime support.
|
||||
"""
|
||||
|
||||
from fastvideo.mlx_runtime.fastwan import (
|
||||
FastWanShape,
|
||||
MLXQuantizationSpec,
|
||||
MLXWanDiT,
|
||||
MLXWanTransformerBlock,
|
||||
UnsupportedMLXQuantizationError,
|
||||
ensure_quantization_supported,
|
||||
fastwan_shape,
|
||||
fastwan_shape_from_config,
|
||||
mlx_dit_from_diffusers_safetensors,
|
||||
mlx_block_weights_from_torch,
|
||||
mlx_block_weights_from_diffusers_safetensors,
|
||||
quantization_support_error,
|
||||
)
|
||||
from fastvideo.mlx_runtime.checkpoint import (
|
||||
load_mlx_dit_checkpoint,
|
||||
save_mlx_dit_checkpoint,
|
||||
)
|
||||
from fastvideo.mlx_runtime.memory import (
|
||||
AppliedMemoryLimits,
|
||||
add_memory_limit_args,
|
||||
apply_memory_limits,
|
||||
gib_to_bytes,
|
||||
)
|
||||
from fastvideo.mlx_runtime.refine import (
|
||||
DEFAULT_REFINE_SIGMA,
|
||||
RefinePlan,
|
||||
TwoPassResult,
|
||||
default_refine_timesteps,
|
||||
plan_refine_resolutions,
|
||||
prepare_refine_latents,
|
||||
refine_sigma_from_schedule,
|
||||
run_dmd_loop,
|
||||
run_two_pass_dmd,
|
||||
upsample_latents_spatial,
|
||||
)
|
||||
from fastvideo.mlx_runtime.frame_upsample import (
|
||||
DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
PIXEL_UPSAMPLE_MODES,
|
||||
unsharp,
|
||||
upsample_frame,
|
||||
upsample_frames,
|
||||
)
|
||||
from fastvideo.mlx_runtime.fast_spatial import (
|
||||
DEFAULT_FAST_SPATIAL_SHARPEN,
|
||||
FastSpatialPlan,
|
||||
apply_fast_spatial_upsample,
|
||||
plan_fast_spatial,
|
||||
resolve_spatial_mode,
|
||||
)
|
||||
from fastvideo.mlx_runtime.prompt_enhance import (
|
||||
DEFAULT_ENHANCE_SYSTEM_PROMPT,
|
||||
DEFAULT_MLX_LM_MODEL,
|
||||
EnhanceResult,
|
||||
enhance_prompt,
|
||||
enhance_prompt_template,
|
||||
enhance_result_as_metrics,
|
||||
load_or_enhance_prompt,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AppliedMemoryLimits",
|
||||
"DEFAULT_ENHANCE_SYSTEM_PROMPT",
|
||||
"DEFAULT_MLX_LM_MODEL",
|
||||
"DEFAULT_FAST_SPATIAL_SHARPEN",
|
||||
"DEFAULT_PIXEL_UPSAMPLE_MODE",
|
||||
"DEFAULT_REFINE_SIGMA",
|
||||
"EnhanceResult",
|
||||
"FastSpatialPlan",
|
||||
"FastWanShape",
|
||||
"MLXQuantizationSpec",
|
||||
"MLXWanDiT",
|
||||
"MLXWanTransformerBlock",
|
||||
"RefinePlan",
|
||||
"TwoPassResult",
|
||||
"UnsupportedMLXQuantizationError",
|
||||
"add_memory_limit_args",
|
||||
"apply_fast_spatial_upsample",
|
||||
"apply_memory_limits",
|
||||
"enhance_prompt",
|
||||
"enhance_prompt_template",
|
||||
"enhance_result_as_metrics",
|
||||
"ensure_quantization_supported",
|
||||
"fastwan_shape",
|
||||
"fastwan_shape_from_config",
|
||||
"gib_to_bytes",
|
||||
"load_mlx_dit_checkpoint",
|
||||
"load_or_enhance_prompt",
|
||||
"mlx_dit_from_diffusers_safetensors",
|
||||
"mlx_block_weights_from_diffusers_safetensors",
|
||||
"mlx_block_weights_from_torch",
|
||||
"PIXEL_UPSAMPLE_MODES",
|
||||
"default_refine_timesteps",
|
||||
"plan_fast_spatial",
|
||||
"plan_refine_resolutions",
|
||||
"prepare_refine_latents",
|
||||
"quantization_support_error",
|
||||
"refine_sigma_from_schedule",
|
||||
"resolve_spatial_mode",
|
||||
"run_dmd_loop",
|
||||
"run_two_pass_dmd",
|
||||
"save_mlx_dit_checkpoint",
|
||||
"unsharp",
|
||||
"upsample_frame",
|
||||
"upsample_frames",
|
||||
"upsample_latents_spatial",
|
||||
]
|
||||
@@ -1,273 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pre-quantized MLX checkpoint save/load for the FastWan DiT.
|
||||
|
||||
Loading the Diffusers fp32/fp16 checkpoint and quantizing at startup costs
|
||||
both download size and load time on every run. This module persists an already
|
||||
cast (and optionally already quantized) ``MLXWanDiT`` so 16 GB users download
|
||||
and load roughly half the bytes and skip requantization entirely:
|
||||
|
||||
dit = mlx_dit_from_diffusers_safetensors(ckpt, cfg, quantization="int8")
|
||||
save_mlx_dit_checkpoint(dit, "FastWan2.1-T2V-1.3B-mlx-int8")
|
||||
...
|
||||
dit = load_mlx_dit_checkpoint("FastWan2.1-T2V-1.3B-mlx-int8")
|
||||
|
||||
Format (one directory):
|
||||
|
||||
- ``mlx_dit.safetensors`` — every array, saved with ``mx.save_safetensors``.
|
||||
Plain weights keep their key; a quantized weight ``K`` is stored as the
|
||||
packed ``K`` plus ``K.scales`` (and ``K.biases`` for affine modes).
|
||||
- ``mlx_dit.json`` — format version, the model config, the quantization spec,
|
||||
and which keys are quantized, so the loader can rebuild ``QuantizedMatrix``
|
||||
objects without guessing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.mlx_runtime.fastwan import (
|
||||
MLXQuantizationSpec,
|
||||
MLXWanDiT,
|
||||
MLXWanTransformerBlock,
|
||||
QuantizedMatrix,
|
||||
ensure_quantization_supported,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
FORMAT_VERSION = 1
|
||||
WEIGHTS_FILENAME = "mlx_dit.safetensors"
|
||||
MANIFEST_FILENAME = "mlx_dit.json"
|
||||
|
||||
_BLOCK_PREFIX = "blocks"
|
||||
|
||||
_DTYPE_TO_NAME = {"float16": "fp16", "bfloat16": "bf16", "float32": "fp32"}
|
||||
|
||||
|
||||
def _dtype_name(dtype) -> str:
|
||||
"""Return the manifest name for a supported MLX data type.
|
||||
|
||||
Parameters:
|
||||
dtype: The MLX data type to convert.
|
||||
|
||||
Returns:
|
||||
str: The manifest name corresponding to the data type.
|
||||
|
||||
Raises:
|
||||
ValueError: If the data type is not supported for checkpointing.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
for raw, name in _DTYPE_TO_NAME.items():
|
||||
if dtype == getattr(mx, raw):
|
||||
return name
|
||||
raise ValueError(f"Unsupported MLX dtype for checkpointing: {dtype}")
|
||||
|
||||
|
||||
def _name_to_dtype(name: str):
|
||||
"""Convert a manifest dtype name to its corresponding MLX dtype.
|
||||
|
||||
Parameters:
|
||||
name (str): Manifest name, such as ``"fp16"``, ``"bf16"``, or ``"fp32"``.
|
||||
|
||||
Returns:
|
||||
The corresponding MLX dtype.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
return {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[name]
|
||||
|
||||
|
||||
def _flatten_weights(dit: MLXWanDiT) -> dict[str, Any]:
|
||||
"""Combine model and transformer-block weights into a single flattened mapping.
|
||||
|
||||
Parameters:
|
||||
dit (MLXWanDiT): Model whose weights should be flattened.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: Mapping containing top-level weights and indexed transformer-block weights.
|
||||
"""
|
||||
flat: dict[str, Any] = dict(dit.weights)
|
||||
for index, block in enumerate(dit.blocks):
|
||||
for name, value in block.weights.items():
|
||||
flat[f"{_BLOCK_PREFIX}.{index}.{name}"] = value
|
||||
return flat
|
||||
|
||||
|
||||
def save_mlx_dit_checkpoint(dit: MLXWanDiT, checkpoint_dir: str | Path) -> Path:
|
||||
"""Save a plain or quantized MLX Wan DiT checkpoint to a directory.
|
||||
|
||||
Parameters:
|
||||
dit (MLXWanDiT): Model whose weights and configuration will be saved.
|
||||
checkpoint_dir (str | Path): Destination directory for the checkpoint.
|
||||
|
||||
Returns:
|
||||
Path: Path to the checkpoint directory.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
checkpoint_dir = Path(checkpoint_dir)
|
||||
arrays: dict[str, Any] = {}
|
||||
quantized: dict[str, dict[str, Any]] = {}
|
||||
spec: MLXQuantizationSpec | None = None
|
||||
for key, value in _flatten_weights(dit).items():
|
||||
if isinstance(value, QuantizedMatrix):
|
||||
if spec is not None and value.spec != spec:
|
||||
raise ValueError(f"Mixed quantization specs in one checkpoint ({spec} vs {value.spec} at '{key}') "
|
||||
"are not supported.")
|
||||
spec = value.spec
|
||||
arrays[key] = value.weight
|
||||
arrays[f"{key}.scales"] = value.scales
|
||||
if value.biases is not None:
|
||||
arrays[f"{key}.biases"] = value.biases
|
||||
quantized[key] = {
|
||||
"dequantized_dtype": _dtype_name(value.dequantized_dtype),
|
||||
"has_biases": value.biases is not None,
|
||||
}
|
||||
else:
|
||||
arrays[key] = value
|
||||
|
||||
manifest = {
|
||||
"format_version": FORMAT_VERSION,
|
||||
"config": dit.config,
|
||||
"num_blocks": len(dit.blocks),
|
||||
"quantization": None if spec is None else {
|
||||
"mode": spec.mode,
|
||||
"bits": spec.bits,
|
||||
"group_size": spec.group_size,
|
||||
},
|
||||
"quantized_keys": quantized,
|
||||
}
|
||||
|
||||
manifest_json = json.dumps(manifest, indent=2)
|
||||
checkpoint_dir.parent.mkdir(parents=True, exist_ok=True)
|
||||
staging_dir = Path(tempfile.mkdtemp(dir=checkpoint_dir.parent, prefix=f".{checkpoint_dir.name}.staging-"))
|
||||
backup_root: Path | None = None
|
||||
try:
|
||||
staged_weights = staging_dir / WEIGHTS_FILENAME
|
||||
staged_manifest = staging_dir / MANIFEST_FILENAME
|
||||
mx.save_safetensors(str(staged_weights), arrays)
|
||||
staged_manifest.write_text(manifest_json)
|
||||
if checkpoint_dir.exists():
|
||||
backup_root = Path(tempfile.mkdtemp(dir=checkpoint_dir.parent, prefix=f".{checkpoint_dir.name}.backup-"))
|
||||
try:
|
||||
checkpoint_dir.replace(backup_root / checkpoint_dir.name)
|
||||
except Exception:
|
||||
shutil.rmtree(backup_root, ignore_errors=True)
|
||||
raise
|
||||
try:
|
||||
staging_dir.replace(checkpoint_dir)
|
||||
except Exception:
|
||||
if backup_root is not None:
|
||||
(backup_root / checkpoint_dir.name).replace(checkpoint_dir)
|
||||
shutil.rmtree(backup_root, ignore_errors=True)
|
||||
raise
|
||||
if backup_root is not None:
|
||||
shutil.rmtree(backup_root, ignore_errors=True)
|
||||
finally:
|
||||
shutil.rmtree(staging_dir, ignore_errors=True)
|
||||
logger.info("Saved MLX DiT checkpoint (%d arrays, quantization=%s) to %s", len(arrays),
|
||||
spec.label if spec else "none", checkpoint_dir)
|
||||
return checkpoint_dir
|
||||
|
||||
|
||||
def load_mlx_dit_checkpoint(checkpoint_dir: str | Path, *, compile: bool = False) -> MLXWanDiT:
|
||||
"""
|
||||
Reconstruct an MLXWanDiT model from a versioned checkpoint.
|
||||
|
||||
Parameters:
|
||||
checkpoint_dir (str | Path): Directory containing the checkpoint manifest and weights.
|
||||
compile (bool): Whether to configure the reconstructed model for compilation.
|
||||
|
||||
Returns:
|
||||
MLXWanDiT: The reconstructed model.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the checkpoint manifest or weights file is missing.
|
||||
ValueError: If the checkpoint format is unsupported or block weights are incomplete.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
checkpoint_dir = Path(checkpoint_dir)
|
||||
manifest_path = checkpoint_dir / MANIFEST_FILENAME
|
||||
weights_path = checkpoint_dir / WEIGHTS_FILENAME
|
||||
if not manifest_path.exists() or not weights_path.exists():
|
||||
raise FileNotFoundError(f"Not an MLX DiT checkpoint directory: {checkpoint_dir} "
|
||||
f"(expected {MANIFEST_FILENAME} and {WEIGHTS_FILENAME}).")
|
||||
|
||||
manifest = json.loads(manifest_path.read_text())
|
||||
version = manifest.get("format_version")
|
||||
if version != FORMAT_VERSION:
|
||||
raise ValueError(f"MLX DiT checkpoint {checkpoint_dir} has format_version={version}; "
|
||||
f"this FastVideo build reads version {FORMAT_VERSION}. Re-export the checkpoint.")
|
||||
|
||||
spec = None
|
||||
if manifest["quantization"] is not None:
|
||||
spec = MLXQuantizationSpec(**manifest["quantization"])
|
||||
# The packed layout of mx.quantize output is mode-specific, so a build
|
||||
# that cannot run the mode cannot use these arrays at all.
|
||||
ensure_quantization_supported(spec)
|
||||
|
||||
arrays = mx.load(str(weights_path))
|
||||
quantized_keys: dict[str, dict[str, Any]] = manifest["quantized_keys"]
|
||||
|
||||
def rebuild(key: str):
|
||||
"""
|
||||
Reconstructs a weight array or quantized matrix from checkpoint data.
|
||||
|
||||
Parameters:
|
||||
key (str): The weight key to rebuild.
|
||||
|
||||
Returns:
|
||||
The stored array for an unquantized weight or a reconstructed quantized matrix.
|
||||
"""
|
||||
if key not in quantized_keys:
|
||||
return arrays[key]
|
||||
info = quantized_keys[key]
|
||||
assert spec is not None, f"Quantized key '{key}' in a checkpoint without a quantization spec"
|
||||
return QuantizedMatrix(
|
||||
weight=arrays[key],
|
||||
scales=arrays[f"{key}.scales"],
|
||||
biases=arrays[f"{key}.biases"] if info["has_biases"] else None,
|
||||
spec=spec,
|
||||
dequantized_dtype=_name_to_dtype(info["dequantized_dtype"]),
|
||||
)
|
||||
|
||||
config = manifest["config"]
|
||||
block_keys: dict[int, list[str]] = {}
|
||||
top_level_keys: list[str] = []
|
||||
for key in arrays:
|
||||
if key.endswith(".scales") or key.endswith(".biases"):
|
||||
continue
|
||||
if key.startswith(f"{_BLOCK_PREFIX}."):
|
||||
index_str, _, _ = key[len(_BLOCK_PREFIX) + 1:].partition(".")
|
||||
block_keys.setdefault(int(index_str), []).append(key)
|
||||
else:
|
||||
top_level_keys.append(key)
|
||||
|
||||
weights = {key: rebuild(key) for key in top_level_keys}
|
||||
|
||||
num_blocks = int(manifest["num_blocks"])
|
||||
if sorted(block_keys) != list(range(num_blocks)):
|
||||
raise ValueError(f"MLX DiT checkpoint {checkpoint_dir} is missing block weights: "
|
||||
f"manifest says {num_blocks} blocks, found indices {sorted(block_keys)}.")
|
||||
|
||||
inner_dim = int(config["num_attention_heads"]) * int(config["attention_head_dim"])
|
||||
blocks = []
|
||||
for index in range(num_blocks):
|
||||
prefix = f"{_BLOCK_PREFIX}.{index}."
|
||||
block_weights = {key[len(prefix):]: rebuild(key) for key in block_keys[index]}
|
||||
blocks.append(
|
||||
MLXWanTransformerBlock(
|
||||
block_weights,
|
||||
dim=inner_dim,
|
||||
ffn_dim=int(config["ffn_dim"]),
|
||||
num_heads=int(config["num_attention_heads"]),
|
||||
eps=float(config["eps"]),
|
||||
))
|
||||
return MLXWanDiT(weights, blocks, config, compile=compile)
|
||||
@@ -1,228 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Spatial fast mode for the MLX Wan runtime (RIFE's spatial twin).
|
||||
|
||||
RIFE ``--fast`` cuts *frames* (temporal). This module cuts *pixels*
|
||||
(spatial): denoise at ``target // scale``, decode at that size, then
|
||||
resample the decoded frames up to the target. No second denoise pass —
|
||||
that is ``--refine`` (quality). The two compose:
|
||||
|
||||
* ``--fast-spatial`` alone → speed (≈ scale² fewer tokens)
|
||||
* ``--refine`` alone → quality two-pass (H3 / LTX-2)
|
||||
* ``--fast`` + ``--refine`` → fewer frames at base res, full-res refine
|
||||
* ``--fast`` + ``--fast-spatial`` → fewer frames *and* fewer pixels
|
||||
|
||||
The upsample runs in **pixel** space, after the VAE decode. It used to run
|
||||
in latent space (bilinear over the latent H/W plane, sharing the refine
|
||||
hand-off primitive) and that is what made spatial fast mode incoherent: an
|
||||
interpolated Wan latent is off the decoder's manifold, so decode returned
|
||||
the right silhouette under a smeared veil. ``--refine`` can get away with
|
||||
the latent-space upsample because a second DMD pass re-denoises the result;
|
||||
spatial fast mode hands the latent straight to the decoder, so it cannot.
|
||||
See :mod:`fastvideo.mlx_runtime.frame_upsample` for the full rationale.
|
||||
|
||||
MetalFX is intentionally not used: it needs game-engine motion vectors
|
||||
and depth that diffusion output lacks.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.mlx_runtime.frame_upsample import (
|
||||
DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
PIXEL_UPSAMPLE_MODES,
|
||||
upsample_frames,
|
||||
)
|
||||
from fastvideo.mlx_runtime.refine import RefinePlan, plan_refine_resolutions
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Resampling from a smaller decode loses high-frequency detail the same way
|
||||
# RIFE's flow warp does, so spatial fast mode borrows ``--fast``'s remedy: a
|
||||
# light unsharp pass. 0.4 recovers perceived crispness on Wan2.1 output at 2x
|
||||
# without the halos that show up by ~0.8.
|
||||
DEFAULT_FAST_SPATIAL_SHARPEN = 0.4
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FastSpatialPlan:
|
||||
"""Resolved geometry for a spatial-fast (upsample-only) run."""
|
||||
|
||||
plan: RefinePlan
|
||||
upsample_mode: str
|
||||
sharpen: float = DEFAULT_FAST_SPATIAL_SHARPEN
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
"""
|
||||
Determine whether spatial scaling is enabled.
|
||||
|
||||
Returns:
|
||||
`true` if the spatial scale is greater than one, `false` otherwise.
|
||||
"""
|
||||
return self.plan.spatial_scale > 1
|
||||
|
||||
@property
|
||||
def scale(self) -> int:
|
||||
"""Provides the configured spatial scaling factor.
|
||||
|
||||
Returns:
|
||||
int: The spatial scaling factor.
|
||||
"""
|
||||
return self.plan.spatial_scale
|
||||
|
||||
@property
|
||||
def target_height(self) -> int:
|
||||
"""
|
||||
Return the target output height for the spatial plan.
|
||||
|
||||
Returns:
|
||||
int: Target output height in pixels.
|
||||
"""
|
||||
return self.plan.target_height
|
||||
|
||||
@property
|
||||
def target_width(self) -> int:
|
||||
"""Return the target image width in pixels.
|
||||
|
||||
Returns:
|
||||
int: The target image width.
|
||||
"""
|
||||
return self.plan.target_width
|
||||
|
||||
@property
|
||||
def stage1_height(self) -> int:
|
||||
"""
|
||||
Provide the stage-one latent height used for reduced-resolution processing.
|
||||
|
||||
Returns:
|
||||
int: The stage-one latent height.
|
||||
"""
|
||||
return self.plan.stage1_height
|
||||
|
||||
@property
|
||||
def stage1_width(self) -> int:
|
||||
"""Get the stage-one latent width.
|
||||
|
||||
Returns:
|
||||
int: The stage-one latent width.
|
||||
"""
|
||||
return self.plan.stage1_width
|
||||
|
||||
|
||||
def plan_fast_spatial(
|
||||
*,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
spatial_scale: int = 2,
|
||||
vae_spatial_compression: int = 8,
|
||||
vae_temporal_compression: int = 4,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
upsample_mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
sharpen: float = DEFAULT_FAST_SPATIAL_SHARPEN,
|
||||
enabled: bool = True,
|
||||
) -> FastSpatialPlan:
|
||||
"""
|
||||
Build a plan for reduced-resolution denoising followed by pixel-space upsampling.
|
||||
|
||||
Parameters:
|
||||
upsample_mode (str): Pixel interpolation kernel, one of
|
||||
:data:`~fastvideo.mlx_runtime.frame_upsample.PIXEL_UPSAMPLE_MODES`.
|
||||
sharpen (float): Unsharp strength applied after the resize.
|
||||
|
||||
Returns:
|
||||
FastSpatialPlan: The validated spatial-fast processing plan.
|
||||
|
||||
Raises:
|
||||
ValueError: If the upsample mode is unsupported or ``sharpen`` is negative.
|
||||
"""
|
||||
if upsample_mode not in PIXEL_UPSAMPLE_MODES:
|
||||
raise ValueError(f"Unsupported upsample mode: {upsample_mode!r} "
|
||||
f"(expected one of {', '.join(PIXEL_UPSAMPLE_MODES)})")
|
||||
if sharpen < 0.0:
|
||||
raise ValueError(f"sharpen must be >= 0, got {sharpen}")
|
||||
plan = plan_refine_resolutions(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
spatial_scale=spatial_scale,
|
||||
vae_spatial_compression=vae_spatial_compression,
|
||||
vae_temporal_compression=vae_temporal_compression,
|
||||
patch_size=patch_size,
|
||||
enabled=enabled,
|
||||
mode_label="fast-spatial",
|
||||
)
|
||||
if plan.spatial_scale > 1:
|
||||
logger.info(
|
||||
"[MLX fast-spatial] denoise+decode %dx%d → upsample %dx to %dx%d (%s, sharpen=%.2f)",
|
||||
plan.stage1_width,
|
||||
plan.stage1_height,
|
||||
plan.spatial_scale,
|
||||
plan.target_width,
|
||||
plan.target_height,
|
||||
upsample_mode,
|
||||
sharpen,
|
||||
)
|
||||
return FastSpatialPlan(plan=plan, upsample_mode=upsample_mode, sharpen=sharpen)
|
||||
|
||||
|
||||
def apply_fast_spatial_upsample(
|
||||
frames: Iterable[np.ndarray],
|
||||
spatial: FastSpatialPlan,
|
||||
) -> list[np.ndarray]:
|
||||
"""Resample decoded stage-1 frames up to the target resolution.
|
||||
|
||||
This runs on decoded RGB frames, *not* on latents: see the module
|
||||
docstring for why the latent-space version produced a blurred veil.
|
||||
|
||||
Parameters:
|
||||
frames (Iterable[np.ndarray]): Decoded HxWx3 uint8 RGB frames, produced
|
||||
by decoding at the stage-one resolution.
|
||||
spatial (FastSpatialPlan): Plan defining the target size, interpolation
|
||||
kernel, and unsharp strength.
|
||||
|
||||
Returns:
|
||||
list[np.ndarray]: Frames at the target resolution. When spatial scaling
|
||||
is disabled the frames are returned unchanged, as a list.
|
||||
"""
|
||||
if not spatial.enabled:
|
||||
return list(frames)
|
||||
return upsample_frames(
|
||||
frames,
|
||||
width=spatial.target_width,
|
||||
height=spatial.target_height,
|
||||
mode=spatial.upsample_mode,
|
||||
sharpen=spatial.sharpen,
|
||||
)
|
||||
|
||||
|
||||
def resolve_spatial_mode(
|
||||
*,
|
||||
refine: bool,
|
||||
fast_spatial: bool,
|
||||
) -> str:
|
||||
"""Select the active spatial processing mode, with refinement taking precedence.
|
||||
|
||||
Returns:
|
||||
str: ``"refine"`` when refinement is enabled, ``"fast_spatial"`` when
|
||||
spatial-fast processing is enabled, or ``"off"`` otherwise.
|
||||
"""
|
||||
if refine:
|
||||
return "refine"
|
||||
if fast_spatial:
|
||||
return "fast_spatial"
|
||||
return "off"
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_FAST_SPATIAL_SHARPEN",
|
||||
"FastSpatialPlan",
|
||||
"apply_fast_spatial_upsample",
|
||||
"plan_fast_spatial",
|
||||
"resolve_spatial_mode",
|
||||
]
|
||||
@@ -1,980 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# mypy: disable-error-code=no-untyped-call
|
||||
"""FastWan-oriented helpers for the experimental MLX runtime path."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import statistics
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from collections.abc import Callable
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FastWanShape:
|
||||
height: int
|
||||
width: int
|
||||
num_frames: int
|
||||
latent_frames: int
|
||||
latent_height: int
|
||||
latent_width: int
|
||||
patch_frames: int
|
||||
patch_height: int
|
||||
patch_width: int
|
||||
tokens: int
|
||||
hidden_size: int
|
||||
num_heads: int
|
||||
head_dim: int
|
||||
|
||||
|
||||
class UnsupportedMLXQuantizationError(ValueError):
|
||||
"""A quantization mode the installed MLX build cannot execute.
|
||||
|
||||
Raised by :func:`ensure_quantization_supported` before any model weights
|
||||
are loaded, so callers (CLI flags, benchmark sweeps) can fail fast with an
|
||||
actionable message -- or skip the mode -- instead of crashing deep inside
|
||||
``mx.quantize`` mid-load.
|
||||
"""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MLXQuantizationSpec:
|
||||
"""MLX quantized-matmul configuration for DiT linear weights."""
|
||||
|
||||
mode: str
|
||||
bits: int | None = None
|
||||
group_size: int | None = None
|
||||
|
||||
@classmethod
|
||||
def from_name(cls, name: str | None) -> MLXQuantizationSpec | None:
|
||||
if name is None or name in {"", "none", "fp16", "fp32"}:
|
||||
return None
|
||||
if name == "int8":
|
||||
return cls(mode="affine", bits=8, group_size=64)
|
||||
if name == "int4":
|
||||
return cls(mode="affine", bits=4, group_size=64)
|
||||
if name == "mxfp8":
|
||||
return cls(mode="mxfp8")
|
||||
if name == "mxfp4":
|
||||
return cls(mode="mxfp4")
|
||||
if name == "nvfp4":
|
||||
return cls(mode="nvfp4")
|
||||
raise ValueError(f"Unsupported MLX quantization mode: {name}")
|
||||
|
||||
@property
|
||||
def label(self) -> str:
|
||||
if self.mode == "affine":
|
||||
return f"int{self.bits}"
|
||||
return self.mode
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantizedMatrix:
|
||||
weight: mx.array
|
||||
scales: mx.array
|
||||
biases: mx.array | None
|
||||
spec: MLXQuantizationSpec
|
||||
dequantized_dtype: mx.Dtype
|
||||
|
||||
|
||||
def fastwan_shape(
|
||||
*,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
vae_temporal_compression: int = 4,
|
||||
vae_spatial_compression: int = 8,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
num_heads: int = 12,
|
||||
head_dim: int = 128,
|
||||
) -> FastWanShape:
|
||||
"""Return the approximate DiT token shape for Wan/FastWan T2V inference."""
|
||||
latent_frames = (num_frames - 1) // vae_temporal_compression + 1
|
||||
latent_height = height // vae_spatial_compression
|
||||
latent_width = width // vae_spatial_compression
|
||||
patch_frames = latent_frames // patch_size[0]
|
||||
patch_height = latent_height // patch_size[1]
|
||||
patch_width = latent_width // patch_size[2]
|
||||
tokens = patch_frames * patch_height * patch_width
|
||||
return FastWanShape(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
latent_frames=latent_frames,
|
||||
latent_height=latent_height,
|
||||
latent_width=latent_width,
|
||||
patch_frames=patch_frames,
|
||||
patch_height=patch_height,
|
||||
patch_width=patch_width,
|
||||
tokens=tokens,
|
||||
hidden_size=num_heads * head_dim,
|
||||
num_heads=num_heads,
|
||||
head_dim=head_dim,
|
||||
)
|
||||
|
||||
|
||||
def fastwan_shape_from_config(
|
||||
config_path: str | Path,
|
||||
*,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
) -> FastWanShape:
|
||||
config = json.loads(Path(config_path).read_text())
|
||||
return fastwan_shape(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
patch_size=tuple(config["patch_size"]),
|
||||
num_heads=int(config["num_attention_heads"]),
|
||||
head_dim=int(config["attention_head_dim"]),
|
||||
)
|
||||
|
||||
|
||||
def replace_tokens(shape: FastWanShape, tokens: int) -> FastWanShape:
|
||||
return FastWanShape(**{**shape.__dict__, "tokens": tokens})
|
||||
|
||||
|
||||
def median_ms(samples: list[float]) -> float:
|
||||
return statistics.median(samples) * 1000.0
|
||||
|
||||
|
||||
def benchmark_mlx_attention(shape: FastWanShape, warmup: int, iters: int) -> float:
|
||||
import mlx.core as mx
|
||||
|
||||
q = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
|
||||
k = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
|
||||
v = mx.random.normal((1, shape.num_heads, shape.tokens, shape.head_dim), dtype=mx.float16)
|
||||
scale = shape.head_dim**-0.5
|
||||
|
||||
for _ in range(warmup):
|
||||
y = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
|
||||
mx.eval(y)
|
||||
|
||||
samples = []
|
||||
for _ in range(iters):
|
||||
start = time.perf_counter()
|
||||
y = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale)
|
||||
mx.eval(y)
|
||||
samples.append(time.perf_counter() - start)
|
||||
return median_ms(samples)
|
||||
|
||||
|
||||
def benchmark_mlx_linear(shape: FastWanShape, warmup: int, iters: int) -> float:
|
||||
import mlx.core as mx
|
||||
|
||||
x = mx.random.normal((shape.tokens, shape.hidden_size), dtype=mx.float16)
|
||||
w = mx.random.normal((shape.hidden_size, shape.hidden_size), dtype=mx.float16)
|
||||
b = mx.zeros((shape.hidden_size, ), dtype=mx.float16)
|
||||
|
||||
for _ in range(warmup):
|
||||
y = x @ w + b
|
||||
mx.eval(y)
|
||||
|
||||
samples = []
|
||||
for _ in range(iters):
|
||||
start = time.perf_counter()
|
||||
y = x @ w + b
|
||||
mx.eval(y)
|
||||
samples.append(time.perf_counter() - start)
|
||||
return median_ms(samples)
|
||||
|
||||
|
||||
def benchmark_torch_mps_attention(shape: FastWanShape, warmup: int, iters: int) -> float | None:
|
||||
try:
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
except ImportError:
|
||||
return None
|
||||
|
||||
if not torch.backends.mps.is_available():
|
||||
return None
|
||||
|
||||
device = torch.device("mps")
|
||||
q = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
|
||||
k = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
|
||||
v = torch.randn((1, shape.num_heads, shape.tokens, shape.head_dim), device=device, dtype=torch.float16)
|
||||
|
||||
for _ in range(warmup):
|
||||
y = F.scaled_dot_product_attention(q, k, v)
|
||||
torch.mps.synchronize()
|
||||
_ = y
|
||||
|
||||
samples = []
|
||||
for _ in range(iters):
|
||||
start = time.perf_counter()
|
||||
y = F.scaled_dot_product_attention(q, k, v)
|
||||
torch.mps.synchronize()
|
||||
_ = y
|
||||
samples.append(time.perf_counter() - start)
|
||||
return median_ms(samples)
|
||||
|
||||
|
||||
def torch_to_mx(tensor) -> mx.array:
|
||||
import mlx.core as mx
|
||||
|
||||
return mx.array(tensor.detach().cpu().float().numpy())
|
||||
|
||||
|
||||
def weight_dtype(weight):
|
||||
if isinstance(weight, QuantizedMatrix):
|
||||
return weight.dequantized_dtype
|
||||
return weight.dtype
|
||||
|
||||
|
||||
_QUANT_SUPPORT_CACHE: dict[tuple[str, int | None, int | None], str | None] = {}
|
||||
|
||||
|
||||
def quantization_support_error(spec: MLXQuantizationSpec) -> str | None:
|
||||
"""Probe whether the installed MLX build supports ``spec``.
|
||||
|
||||
Runs a tiny ``mx.quantize`` + ``mx.quantized_matmul`` with exactly the
|
||||
arguments :func:`quantize_matrix` / :func:`linear` use, so the result
|
||||
reflects the real runtime path. The affine (int8/int4) modes are stable
|
||||
across MLX releases, but the ``mxfp8``/``mxfp4``/``nvfp4`` mode strings
|
||||
require newer MLX builds and raise otherwise. Returns ``None`` when the
|
||||
mode works, else the underlying error message. Cached per spec.
|
||||
"""
|
||||
key = (spec.mode, spec.bits, spec.group_size)
|
||||
if key not in _QUANT_SUPPORT_CACHE:
|
||||
import mlx.core as mx
|
||||
|
||||
try:
|
||||
probe_dim = max(spec.group_size or 0, 64)
|
||||
weight = mx.zeros((probe_dim, probe_dim), dtype=mx.float16)
|
||||
quantized = quantize_matrix(weight, spec)
|
||||
y = linear(mx.zeros((1, probe_dim), dtype=mx.float16), quantized)
|
||||
mx.eval(y)
|
||||
_QUANT_SUPPORT_CACHE[key] = None
|
||||
except Exception as exc: # noqa: BLE001 - MLX raises varied error types per backend/version.
|
||||
_QUANT_SUPPORT_CACHE[key] = f"{type(exc).__name__}: {exc}"
|
||||
return _QUANT_SUPPORT_CACHE[key]
|
||||
|
||||
|
||||
def ensure_quantization_supported(spec: MLXQuantizationSpec | None) -> None:
|
||||
"""Raise :class:`UnsupportedMLXQuantizationError` if ``spec`` cannot run here."""
|
||||
if spec is None:
|
||||
return
|
||||
error = quantization_support_error(spec)
|
||||
if error is None:
|
||||
return
|
||||
import mlx.core as mx
|
||||
|
||||
mlx_version = getattr(mx, "__version__", "unknown")
|
||||
raise UnsupportedMLXQuantizationError(f"MLX quantization mode '{spec.label}' is not supported by the installed mlx "
|
||||
f"({mlx_version}): {error}. Upgrade mlx or pick a supported mode "
|
||||
f"(int8 is currently the most reliable quality/memory target).")
|
||||
|
||||
|
||||
def quantize_matrix(weight, spec: MLXQuantizationSpec | None):
|
||||
if spec is None:
|
||||
return weight
|
||||
import mlx.core as mx
|
||||
|
||||
if len(weight.shape) < 2:
|
||||
return weight
|
||||
q = mx.quantize(weight, group_size=spec.group_size, bits=spec.bits, mode=spec.mode)
|
||||
biases = q[2] if len(q) == 3 else None
|
||||
eval_args = [q[0], q[1]]
|
||||
if biases is not None:
|
||||
eval_args.append(biases)
|
||||
mx.eval(*eval_args)
|
||||
return QuantizedMatrix(
|
||||
weight=q[0],
|
||||
scales=q[1],
|
||||
biases=biases,
|
||||
spec=spec,
|
||||
dequantized_dtype=weight.dtype,
|
||||
)
|
||||
|
||||
|
||||
def linear(x, weight, bias=None):
|
||||
import mlx.core as mx
|
||||
|
||||
if isinstance(weight, QuantizedMatrix):
|
||||
y = mx.quantized_matmul(
|
||||
x,
|
||||
weight.weight,
|
||||
weight.scales,
|
||||
weight.biases,
|
||||
transpose=True,
|
||||
group_size=weight.spec.group_size,
|
||||
bits=weight.spec.bits,
|
||||
mode=weight.spec.mode,
|
||||
).astype(x.dtype)
|
||||
else:
|
||||
y = x @ weight.T
|
||||
if bias is not None:
|
||||
y = y + bias
|
||||
return y
|
||||
|
||||
|
||||
def _use_fast_norm() -> bool:
|
||||
"""Opt-in to MLX's fused ``mx.fast`` normalization kernels.
|
||||
|
||||
Off by default so the numerically-explicit reference path stays the
|
||||
baseline. Set ``FASTVIDEO_MLX_FAST_NORM=1`` to route LayerNorm/RMSNorm
|
||||
through single fused Metal kernels (fewer intermediates, less memory
|
||||
traffic) and benchmark the speedup.
|
||||
"""
|
||||
import os
|
||||
|
||||
return os.environ.get("FASTVIDEO_MLX_FAST_NORM", "0") == "1"
|
||||
|
||||
|
||||
def layer_norm(x, weight=None, bias=None, eps: float = 1e-6):
|
||||
import mlx.core as mx
|
||||
|
||||
if _use_fast_norm():
|
||||
# Compute in fp32 (matching the reference below) so downstream dtype
|
||||
# and precision are identical across call sites.
|
||||
w = weight.astype(mx.float32) if weight is not None else None
|
||||
b = bias.astype(mx.float32) if bias is not None else None
|
||||
return mx.fast.layer_norm(x.astype(mx.float32), w, b, eps)
|
||||
|
||||
x_float = x.astype(mx.float32)
|
||||
mean = mx.mean(x_float, axis=-1, keepdims=True)
|
||||
var = mx.mean(mx.square(x_float - mean), axis=-1, keepdims=True)
|
||||
y = (x_float - mean) * mx.rsqrt(var + eps)
|
||||
if weight is not None:
|
||||
y = y * weight
|
||||
if bias is not None:
|
||||
y = y + bias
|
||||
return y
|
||||
|
||||
|
||||
def rms_norm(x, weight, eps: float = 1e-6):
|
||||
import mlx.core as mx
|
||||
|
||||
if _use_fast_norm():
|
||||
return mx.fast.rms_norm(x, weight, eps)
|
||||
|
||||
orig_dtype = x.dtype
|
||||
x_float = x.astype(mx.float32)
|
||||
variance = mx.mean(mx.square(x_float), axis=-1, keepdims=True)
|
||||
y = x_float * mx.rsqrt(variance + eps)
|
||||
return y.astype(orig_dtype) * weight
|
||||
|
||||
|
||||
def apply_rotary_emb(x, cos, sin, *, is_neox_style: bool = False):
|
||||
"""Apply FastVideo's rotary convention to MLX tensors.
|
||||
|
||||
Args:
|
||||
x: [batch, seq, heads, head_dim]
|
||||
cos/sin: [seq, head_dim] for Wan's full-dimension rotate-pair style,
|
||||
or [seq, head_dim // 2] for traditional RoPE.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
head_size = x.shape[-1]
|
||||
rope_dim = cos.shape[-1]
|
||||
cos = cos[None, :, None, :]
|
||||
sin = sin[None, :, None, :]
|
||||
x_float = x.astype(mx.float32)
|
||||
|
||||
if rope_dim == head_size:
|
||||
x_pairs = x_float.reshape(*x.shape[:-1], -1, 2)
|
||||
x_real = x_pairs[..., 0]
|
||||
x_imag = x_pairs[..., 1]
|
||||
x_rotated = mx.stack([-x_imag, x_real], axis=-1).reshape(*x.shape)
|
||||
return (x_float * cos + x_rotated * sin).astype(x.dtype)
|
||||
|
||||
if is_neox_style:
|
||||
x1, x2 = mx.split(x_float, 2, axis=-1)
|
||||
o1 = x1 * cos - x2 * sin
|
||||
o2 = x2 * cos + x1 * sin
|
||||
return mx.concatenate([o1, o2], axis=-1).astype(x.dtype)
|
||||
|
||||
x1 = x_float[..., ::2]
|
||||
x2 = x_float[..., 1::2]
|
||||
o1 = x1 * cos - x2 * sin
|
||||
o2 = x2 * cos + x1 * sin
|
||||
return mx.stack([o1, o2], axis=-1).reshape(*x.shape).astype(x.dtype)
|
||||
|
||||
|
||||
_WINDOWED_ATTENTION_WARNED = False
|
||||
|
||||
|
||||
def _warn_windowed_attention_once(window: int) -> None:
|
||||
"""Warn that FASTVIDEO_MLX_WINDOW degrades output on a dense-trained DiT.
|
||||
|
||||
Sliding-window self-attention is fast (6.6x at a +-3-frame window on 1.3B)
|
||||
but these checkpoints were trained with dense attention, and restricting it
|
||||
at inference produces heavy colour-block noise: structural agreement with
|
||||
the dense baseline drops to 0.25 at +-3 frames and 0.03 at +-5. Sparsity of
|
||||
this kind is a training-time method. Kept as a research knob, but it should
|
||||
never be on by accident.
|
||||
"""
|
||||
global _WINDOWED_ATTENTION_WARNED
|
||||
if _WINDOWED_ATTENTION_WARNED:
|
||||
return
|
||||
_WINDOWED_ATTENTION_WARNED = True
|
||||
logger.warning(
|
||||
"FASTVIDEO_MLX_WINDOW=%d enables sliding-window self-attention. These "
|
||||
"checkpoints are trained dense; expect severely degraded output. This is "
|
||||
"a research knob, not a speed setting — use --fast-spatial for real "
|
||||
"denoise savings.",
|
||||
window,
|
||||
)
|
||||
|
||||
|
||||
def gelu_tanh(x):
|
||||
"""tanh-approximate GELU, as used by Wan's FFN.
|
||||
|
||||
``mlx.nn.gelu_approx`` is the same tanh approximation behind a fused
|
||||
kernel. On the 1.3B FFN shape (32760x8960) it is bit-identical to the
|
||||
expanded expression below and 3.3x faster — 28.9ms -> 8.7ms per layer,
|
||||
which is 0.6s per denoise step across 30 layers.
|
||||
"""
|
||||
import mlx.nn as nn
|
||||
|
||||
return nn.gelu_approx(x)
|
||||
|
||||
|
||||
def silu(x):
|
||||
import mlx.core as mx
|
||||
|
||||
return x * mx.sigmoid(x)
|
||||
|
||||
|
||||
def timestep_embedding(t, dim: int, max_period: int = 10000):
|
||||
import mlx.core as mx
|
||||
|
||||
half = dim // 2
|
||||
freqs = mx.exp(-math.log(max_period) * mx.arange(0, half, dtype=mx.float32) / half)
|
||||
args = t[:, None].astype(mx.float32) * freqs[None]
|
||||
embedding = mx.concatenate([mx.cos(args), mx.sin(args)], axis=-1)
|
||||
if dim % 2:
|
||||
embedding = mx.concatenate([embedding, mx.zeros_like(embedding[:, :1])], axis=-1)
|
||||
return embedding
|
||||
|
||||
|
||||
def scale_residual(residual, x, gate):
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
def scale_residual_layer_norm_scale_shift(residual, x, gate, shift, scale, weight=None, bias=None, eps: float = 1e-6):
|
||||
if isinstance(gate, int):
|
||||
assert gate == 1
|
||||
residual_output = residual + x
|
||||
else:
|
||||
residual_output = residual + x * gate
|
||||
normalized = layer_norm(residual_output, weight=weight, bias=bias, eps=eps)
|
||||
modulated = normalized * (1.0 + scale) + shift
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
class MLXWanT2VCrossAttention:
|
||||
|
||||
def __init__(self, weights: dict[str, mx.array], *, dim: int, num_heads: int, eps: float = 1e-6) -> None:
|
||||
self.weights = weights
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.eps = eps
|
||||
|
||||
def __call__(self, x, context):
|
||||
import mlx.core as mx
|
||||
|
||||
batch = x.shape[0]
|
||||
q = linear(x, self.weights["attn2.to_q.weight"], self.weights.get("attn2.to_q.bias"))
|
||||
q = rms_norm(q, self.weights["attn2.norm_q.weight"], eps=self.eps).reshape(batch, -1, self.num_heads,
|
||||
self.head_dim)
|
||||
|
||||
if context.shape[1] == 0:
|
||||
attended = mx.zeros_like(q)
|
||||
else:
|
||||
k = linear(context, self.weights["attn2.to_k.weight"], self.weights.get("attn2.to_k.bias"))
|
||||
k = rms_norm(k, self.weights["attn2.norm_k.weight"],
|
||||
eps=self.eps).reshape(batch, -1, self.num_heads, self.head_dim)
|
||||
v = linear(context, self.weights["attn2.to_v.weight"],
|
||||
self.weights.get("attn2.to_v.bias")).reshape(batch, -1, self.num_heads, self.head_dim)
|
||||
attended = mx.fast.scaled_dot_product_attention(
|
||||
q.transpose(0, 2, 1, 3),
|
||||
k.transpose(0, 2, 1, 3),
|
||||
v.transpose(0, 2, 1, 3),
|
||||
scale=self.head_dim**-0.5,
|
||||
).transpose(0, 2, 1, 3)
|
||||
|
||||
attended = attended.reshape(batch, -1, self.dim)
|
||||
return linear(attended, self.weights["attn2.to_out.weight"], self.weights.get("attn2.to_out.bias"))
|
||||
|
||||
|
||||
class MLXWanTransformerBlock:
|
||||
"""Dense T2V Wan transformer block for the experimental MLX runtime.
|
||||
|
||||
This mirrors the non-VSA PyTorch block for single-process dense attention.
|
||||
Rotary embeddings and sequence-parallel paths are intentionally left out of
|
||||
this first parity target.
|
||||
"""
|
||||
|
||||
def __init__(self, weights: dict[str, mx.array], *, dim: int, ffn_dim: int, num_heads: int, eps: float = 1e-6):
|
||||
self.weights = weights
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.eps = eps
|
||||
self.attn2 = MLXWanT2VCrossAttention(weights, dim=dim, num_heads=num_heads, eps=eps)
|
||||
|
||||
def __call__(self, hidden_states, encoder_hidden_states, temb, freqs_cis=None):
|
||||
import mlx.core as mx
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
e = self.weights["scale_shift_table"] + temb.astype(mx.float32)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = mx.split(e, 6, axis=1)
|
||||
|
||||
norm_hidden_states = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
|
||||
norm_hidden_states = (norm_hidden_states * (1.0 + scale_msa) + shift_msa).astype(orig_dtype)
|
||||
|
||||
query = linear(norm_hidden_states, self.weights["to_q.weight"], self.weights.get("to_q.bias"))
|
||||
key = linear(norm_hidden_states, self.weights["to_k.weight"], self.weights.get("to_k.bias"))
|
||||
value = linear(norm_hidden_states, self.weights["to_v.weight"], self.weights.get("to_v.bias"))
|
||||
|
||||
query = rms_norm(query, self.weights["norm_q.weight"],
|
||||
eps=self.eps).reshape(hidden_states.shape[0], -1, self.num_heads, self.head_dim)
|
||||
key = rms_norm(key, self.weights["norm_k.weight"], eps=self.eps).reshape(hidden_states.shape[0], -1,
|
||||
self.num_heads, self.head_dim)
|
||||
value = value.reshape(hidden_states.shape[0], -1, self.num_heads, self.head_dim)
|
||||
|
||||
if freqs_cis is not None:
|
||||
cos, sin = freqs_cis
|
||||
query = apply_rotary_emb(query, cos, sin, is_neox_style=False)
|
||||
key = apply_rotary_emb(key, cos, sin, is_neox_style=False)
|
||||
|
||||
# Self-attention only. FASTVIDEO_MLX_WINDOW=0/unset → full SDPA (byte-identical
|
||||
# to the historical path). When >0, use chunked symmetric sliding-window
|
||||
# attention (see windowed_attention.py). Cross-attn (attn2) stays full.
|
||||
# Optional FASTVIDEO_MLX_WINDOW_SINK (default 0) adds global sink tokens.
|
||||
q_bh = query.transpose(0, 2, 1, 3) # (B, H, S, D)
|
||||
k_bh = key.transpose(0, 2, 1, 3)
|
||||
v_bh = value.transpose(0, 2, 1, 3)
|
||||
scale = self.head_dim**-0.5
|
||||
window = int(os.environ.get("FASTVIDEO_MLX_WINDOW", "0") or "0")
|
||||
if window > 0:
|
||||
from fastvideo.mlx_runtime.windowed_attention import windowed_attention
|
||||
|
||||
_warn_windowed_attention_once(window)
|
||||
sink = int(os.environ.get("FASTVIDEO_MLX_WINDOW_SINK", "0") or "0")
|
||||
attn_output = windowed_attention(q_bh, k_bh, v_bh, window=window, sink=sink, scale=scale)
|
||||
else:
|
||||
attn_output = mx.fast.scaled_dot_product_attention(q_bh, k_bh, v_bh, scale=scale)
|
||||
attn_output = attn_output.transpose(0, 2, 1, 3)
|
||||
attn_output = attn_output.reshape(hidden_states.shape[0], -1, self.dim)
|
||||
attn_output = linear(attn_output, self.weights["to_out.weight"], self.weights.get("to_out.bias"))
|
||||
|
||||
norm_hidden_states, hidden_states = scale_residual_layer_norm_scale_shift(
|
||||
hidden_states,
|
||||
attn_output,
|
||||
gate_msa,
|
||||
0.0,
|
||||
0.0,
|
||||
weight=self.weights["self_attn_residual_norm.norm.weight"],
|
||||
bias=self.weights["self_attn_residual_norm.norm.bias"],
|
||||
eps=self.eps,
|
||||
)
|
||||
norm_hidden_states = norm_hidden_states.astype(orig_dtype)
|
||||
hidden_states = hidden_states.astype(orig_dtype)
|
||||
|
||||
attn_output = self.attn2(norm_hidden_states, encoder_hidden_states)
|
||||
norm_hidden_states, hidden_states = scale_residual_layer_norm_scale_shift(
|
||||
hidden_states,
|
||||
attn_output,
|
||||
1,
|
||||
c_shift_msa,
|
||||
c_scale_msa,
|
||||
eps=self.eps,
|
||||
)
|
||||
norm_hidden_states = norm_hidden_states.astype(orig_dtype)
|
||||
hidden_states = hidden_states.astype(orig_dtype)
|
||||
|
||||
ff_output = linear(norm_hidden_states, self.weights["ffn.fc_in.weight"], self.weights.get("ffn.fc_in.bias"))
|
||||
ff_output = gelu_tanh(ff_output)
|
||||
ff_output = linear(ff_output, self.weights["ffn.fc_out.weight"], self.weights.get("ffn.fc_out.bias"))
|
||||
hidden_states = scale_residual(hidden_states, ff_output, c_gate_msa)
|
||||
return hidden_states.astype(orig_dtype)
|
||||
|
||||
|
||||
def mlx_block_weights_from_torch(torch_block) -> dict[str, mx.array]:
|
||||
return {name: torch_to_mx(value) for name, value in torch_block.state_dict().items()}
|
||||
|
||||
|
||||
class MLXWanDiT:
|
||||
"""Experimental FP16 Wan/FastWan DiT forward path in MLX."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weights: dict[str, mx.array],
|
||||
blocks: list[MLXWanTransformerBlock],
|
||||
config: dict,
|
||||
*,
|
||||
compile: bool = False,
|
||||
) -> None:
|
||||
import os
|
||||
|
||||
self.weights = weights
|
||||
self.blocks = blocks
|
||||
self.config = config
|
||||
self.num_heads = int(config["num_attention_heads"])
|
||||
self.head_dim = int(config["attention_head_dim"])
|
||||
self.hidden_size = self.num_heads * self.head_dim
|
||||
self.ffn_dim = int(config["ffn_dim"])
|
||||
self.in_channels = int(config["in_channels"])
|
||||
self.out_channels = int(config["out_channels"])
|
||||
self.patch_size = tuple(config["patch_size"])
|
||||
self.freq_dim = int(config["freq_dim"])
|
||||
# Opt-in graph fusion. With fixed weights and static shapes, the whole
|
||||
# denoise-step forward is a pure function of (latents, timestep) -- a
|
||||
# good mx.compile target. Off by default so the eager path stays the
|
||||
# baseline; enable via constructor or FASTVIDEO_MLX_COMPILE=1 and verify
|
||||
# with the benchmark's SSIM ~= 1.0 check.
|
||||
self._enable_compile = compile or os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1"
|
||||
self._compiled_forward: Callable[..., Any] | None = None
|
||||
self._compiled_signature: tuple | None = None
|
||||
|
||||
def patch_embed(self, hidden_states):
|
||||
batch, channels, frames, height, width = hidden_states.shape
|
||||
pt, ph, pw = self.patch_size
|
||||
patch_dim = channels * pt * ph * pw
|
||||
x = hidden_states.reshape(batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw)
|
||||
x = x.transpose(0, 2, 4, 6, 1, 3, 5, 7).reshape(batch, -1, patch_dim)
|
||||
return linear(x, self.weights["patch_embedding.weight"], self.weights.get("patch_embedding.bias"))
|
||||
|
||||
def condition(self, timestep, encoder_hidden_states):
|
||||
t_freq = timestep_embedding(timestep, self.freq_dim).astype(
|
||||
weight_dtype(self.weights["condition_embedder.time_embedder.linear_1.weight"]))
|
||||
temb = linear(
|
||||
t_freq,
|
||||
self.weights["condition_embedder.time_embedder.linear_1.weight"],
|
||||
self.weights["condition_embedder.time_embedder.linear_1.bias"],
|
||||
)
|
||||
temb = silu(temb)
|
||||
temb = linear(
|
||||
temb,
|
||||
self.weights["condition_embedder.time_embedder.linear_2.weight"],
|
||||
self.weights["condition_embedder.time_embedder.linear_2.bias"],
|
||||
)
|
||||
timestep_proj = silu(temb)
|
||||
timestep_proj = linear(
|
||||
timestep_proj,
|
||||
self.weights["condition_embedder.time_proj.weight"],
|
||||
self.weights["condition_embedder.time_proj.bias"],
|
||||
).reshape(timestep.shape[0], 6, self.hidden_size)
|
||||
|
||||
encoder_hidden_states = linear(
|
||||
encoder_hidden_states,
|
||||
self.weights["condition_embedder.text_embedder.linear_1.weight"],
|
||||
self.weights["condition_embedder.text_embedder.linear_1.bias"],
|
||||
)
|
||||
encoder_hidden_states = gelu_tanh(encoder_hidden_states)
|
||||
encoder_hidden_states = linear(
|
||||
encoder_hidden_states,
|
||||
self.weights["condition_embedder.text_embedder.linear_2.weight"],
|
||||
self.weights["condition_embedder.text_embedder.linear_2.bias"],
|
||||
)
|
||||
return temb, timestep_proj, encoder_hidden_states
|
||||
|
||||
def output(self, hidden_states, temb, *, batch: int, frames: int, height: int, width: int):
|
||||
pt, ph, pw = self.patch_size
|
||||
post_patch_frames = frames // pt
|
||||
post_patch_height = height // ph
|
||||
post_patch_width = width // pw
|
||||
shift, scale = mx_split_two(self.weights["scale_shift_table"] + temb[:, None, :], axis=1)
|
||||
hidden_states = layer_norm(hidden_states, eps=float(self.config["eps"])) * (1.0 + scale) + shift
|
||||
hidden_states = hidden_states.astype(weight_dtype(self.weights["proj_out.weight"]))
|
||||
hidden_states = linear(hidden_states, self.weights["proj_out.weight"], self.weights["proj_out.bias"])
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch,
|
||||
post_patch_frames,
|
||||
post_patch_height,
|
||||
post_patch_width,
|
||||
pt,
|
||||
ph,
|
||||
pw,
|
||||
self.out_channels,
|
||||
)
|
||||
hidden_states = hidden_states.transpose(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
return hidden_states.reshape(batch, self.out_channels, frames, height, width)
|
||||
|
||||
def _forward(self, hidden_states, encoder_hidden_states, timestep, cos, sin):
|
||||
"""Pure forward used both eagerly and as the mx.compile target.
|
||||
|
||||
``cos``/``sin`` are passed as separate array args (rather than a tuple)
|
||||
so the function traces cleanly under mx.compile.
|
||||
"""
|
||||
batch, _, frames, height, width = hidden_states.shape
|
||||
freqs_cis = (cos, sin) if cos is not None else None
|
||||
hidden_states = self.patch_embed(hidden_states)
|
||||
temb, timestep_proj, encoder_hidden_states = self.condition(timestep, encoder_hidden_states)
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, freqs_cis=freqs_cis)
|
||||
return self.output(hidden_states, temb, batch=batch, frames=frames, height=height, width=width)
|
||||
|
||||
def __call__(self, hidden_states, encoder_hidden_states, timestep, freqs_cis):
|
||||
cos, sin = freqs_cis if freqs_cis is not None else (None, None)
|
||||
if self._enable_compile and cos is not None:
|
||||
import mlx.core as mx
|
||||
|
||||
# mx.compile keeps one traced graph per input signature, and each
|
||||
# graph pins its own materialization of the quantized weights. The
|
||||
# two-pass modes (--refine) call the DiT at a second resolution, so
|
||||
# keeping both graphs alive doubles resident DiT memory: 14B refine
|
||||
# peaked at 34.7 GiB instead of 20.8 GiB. Retire the previous graph
|
||||
# when the signature changes; the retrace costs far less than a
|
||||
# second copy of the weights.
|
||||
signature = (hidden_states.shape, encoder_hidden_states.shape, timestep.shape)
|
||||
if self._compiled_forward is not None and signature != self._compiled_signature:
|
||||
self._compiled_forward = None
|
||||
self._compiled_signature = None
|
||||
mx.clear_cache()
|
||||
if self._compiled_forward is None:
|
||||
self._compiled_forward = mx.compile(self._forward)
|
||||
self._compiled_signature = signature
|
||||
compiled_forward = self._compiled_forward
|
||||
try:
|
||||
return compiled_forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
|
||||
except Exception as exc: # noqa: BLE001 - some quant graphs may not trace; fall back to eager.
|
||||
logger.warning("mx.compile forward failed (%s); falling back to eager execution.", exc)
|
||||
self._enable_compile = False
|
||||
self._compiled_forward = None
|
||||
self._compiled_signature = None
|
||||
return self._forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
|
||||
|
||||
|
||||
def mx_split_two(x, *, axis: int):
|
||||
import mlx.core as mx
|
||||
|
||||
left, right = mx.split(x, 2, axis=axis)
|
||||
return left, right
|
||||
|
||||
|
||||
def _load_safetensor_value(handle, name: str):
|
||||
return handle.get_tensor(name)
|
||||
|
||||
|
||||
def _load_mx_array_from_safetensor(handle, name: str, dtype):
|
||||
"""Load a safetensors value and cast before creating the MLX array.
|
||||
|
||||
The FastWan Diffusers checkpoint is fp32. Creating an MLX array first and
|
||||
then casting it to fp16 briefly materializes a large fp32 MLX allocation.
|
||||
Casting the CPU tensor before crossing into MLX keeps the transient GPU-side
|
||||
footprint lower.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
import torch
|
||||
|
||||
tensor = handle.get_tensor(name)
|
||||
if dtype == mx.float16:
|
||||
tensor = tensor.to(torch.float16)
|
||||
elif dtype == mx.float32:
|
||||
tensor = tensor.to(torch.float32)
|
||||
elif dtype == mx.bfloat16:
|
||||
# NumPy has no bfloat16, so bridge through fp32 and cast on-device below.
|
||||
tensor = tensor.to(torch.float32)
|
||||
array = mx.array(tensor.numpy())
|
||||
del tensor
|
||||
if dtype is not None and array.dtype != dtype:
|
||||
array = array.astype(dtype)
|
||||
mx.eval(array)
|
||||
return array
|
||||
|
||||
|
||||
def _eval_loaded_weight(value) -> None:
|
||||
import mlx.core as mx
|
||||
|
||||
if isinstance(value, QuantizedMatrix):
|
||||
eval_args = [value.weight, value.scales]
|
||||
if value.biases is not None:
|
||||
eval_args.append(value.biases)
|
||||
mx.eval(*eval_args)
|
||||
else:
|
||||
mx.eval(value)
|
||||
|
||||
|
||||
# Diffusers-to-FastVideo key mapping for WanTransformerBlock weights.
|
||||
# Shared by both MLX and torch block loaders to keep mappings synchronized.
|
||||
_WAN_BLOCK_KEY_MAP = {
|
||||
"scale_shift_table": "scale_shift_table",
|
||||
"attn1.to_q.weight": "to_q.weight",
|
||||
"attn1.to_q.bias": "to_q.bias",
|
||||
"attn1.to_k.weight": "to_k.weight",
|
||||
"attn1.to_k.bias": "to_k.bias",
|
||||
"attn1.to_v.weight": "to_v.weight",
|
||||
"attn1.to_v.bias": "to_v.bias",
|
||||
"attn1.to_out.0.weight": "to_out.weight",
|
||||
"attn1.to_out.0.bias": "to_out.bias",
|
||||
"attn1.norm_q.weight": "norm_q.weight",
|
||||
"attn1.norm_k.weight": "norm_k.weight",
|
||||
"attn2.to_q.weight": "attn2.to_q.weight",
|
||||
"attn2.to_q.bias": "attn2.to_q.bias",
|
||||
"attn2.to_k.weight": "attn2.to_k.weight",
|
||||
"attn2.to_k.bias": "attn2.to_k.bias",
|
||||
"attn2.to_v.weight": "attn2.to_v.weight",
|
||||
"attn2.to_v.bias": "attn2.to_v.bias",
|
||||
"attn2.to_out.0.weight": "attn2.to_out.weight",
|
||||
"attn2.to_out.0.bias": "attn2.to_out.bias",
|
||||
"attn2.norm_q.weight": "attn2.norm_q.weight",
|
||||
"attn2.norm_k.weight": "attn2.norm_k.weight",
|
||||
"ffn.net.0.proj.weight": "ffn.fc_in.weight",
|
||||
"ffn.net.0.proj.bias": "ffn.fc_in.bias",
|
||||
"ffn.net.2.weight": "ffn.fc_out.weight",
|
||||
"ffn.net.2.bias": "ffn.fc_out.bias",
|
||||
"norm2.weight": "self_attn_residual_norm.norm.weight",
|
||||
"norm2.bias": "self_attn_residual_norm.norm.bias",
|
||||
}
|
||||
|
||||
|
||||
def mlx_block_weights_from_diffusers_safetensors(
|
||||
checkpoint_path: str | Path,
|
||||
*,
|
||||
block_index: int = 0,
|
||||
quantization: str | MLXQuantizationSpec | None = None,
|
||||
dtype=None,
|
||||
) -> dict[str, mx.array]:
|
||||
"""Load one Diffusers-format Wan block into the MLX dense-block key layout."""
|
||||
from safetensors import safe_open
|
||||
|
||||
prefix = f"blocks.{block_index}."
|
||||
key_map = _WAN_BLOCK_KEY_MAP
|
||||
|
||||
spec = MLXQuantizationSpec.from_name(quantization) if (quantization is None
|
||||
or isinstance(quantization, str)) else quantization
|
||||
ensure_quantization_supported(spec)
|
||||
matrix_targets = {target for target in key_map.values() if target.endswith(".weight") and "norm" not in target}
|
||||
weights = {}
|
||||
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
|
||||
available = set(handle.keys())
|
||||
for source_name, target_name in key_map.items():
|
||||
full = prefix + source_name
|
||||
if full not in available:
|
||||
# Biases are optional: e.g. Wan2.1-14B has bias-free attention/FFN.
|
||||
# The block forward already fetches biases via ``.get(...)``.
|
||||
if source_name.endswith(".bias"):
|
||||
continue
|
||||
raise KeyError(f"missing required block weight: {full}")
|
||||
array = _load_mx_array_from_safetensor(handle, full, dtype)
|
||||
loaded = quantize_matrix(array, spec) if target_name in matrix_targets else array
|
||||
_eval_loaded_weight(loaded)
|
||||
weights[target_name] = loaded
|
||||
del array
|
||||
return weights
|
||||
|
||||
|
||||
def mlx_dit_from_diffusers_safetensors(
|
||||
checkpoint_path: str | Path,
|
||||
config_path: str | Path,
|
||||
*,
|
||||
dtype: str = "fp16",
|
||||
num_blocks: int | None = None,
|
||||
quantization: str | MLXQuantizationSpec | None = None,
|
||||
compile: bool = False,
|
||||
) -> MLXWanDiT:
|
||||
import mlx.core as mx
|
||||
from safetensors import safe_open
|
||||
|
||||
config = json.loads(Path(config_path).read_text())
|
||||
total_blocks = int(config["num_layers"])
|
||||
if num_blocks is None:
|
||||
num_blocks = total_blocks
|
||||
cast_dtype = {"fp16": mx.float16, "bf16": mx.bfloat16, "fp32": mx.float32}[dtype]
|
||||
spec = MLXQuantizationSpec.from_name(quantization) if (quantization is None
|
||||
or isinstance(quantization, str)) else quantization
|
||||
ensure_quantization_supported(spec)
|
||||
|
||||
top_level_names = [
|
||||
"patch_embedding.weight",
|
||||
"patch_embedding.bias",
|
||||
"condition_embedder.time_embedder.linear_1.weight",
|
||||
"condition_embedder.time_embedder.linear_1.bias",
|
||||
"condition_embedder.time_embedder.linear_2.weight",
|
||||
"condition_embedder.time_embedder.linear_2.bias",
|
||||
"condition_embedder.time_proj.weight",
|
||||
"condition_embedder.time_proj.bias",
|
||||
"condition_embedder.text_embedder.linear_1.weight",
|
||||
"condition_embedder.text_embedder.linear_1.bias",
|
||||
"condition_embedder.text_embedder.linear_2.weight",
|
||||
"condition_embedder.text_embedder.linear_2.bias",
|
||||
"scale_shift_table",
|
||||
"proj_out.weight",
|
||||
"proj_out.bias",
|
||||
]
|
||||
weights = {}
|
||||
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
|
||||
available = set(handle.keys())
|
||||
for name in top_level_names:
|
||||
if name not in available:
|
||||
if name.endswith(".bias"):
|
||||
continue
|
||||
raise KeyError(f"missing required weight: {name}")
|
||||
array = _load_mx_array_from_safetensor(handle, name, cast_dtype)
|
||||
if name == "patch_embedding.weight":
|
||||
array = array.reshape(int(config["num_attention_heads"]) * int(config["attention_head_dim"]), -1)
|
||||
if name.endswith(".weight") and name not in {"scale_shift_table"}:
|
||||
loaded = quantize_matrix(array, spec)
|
||||
else:
|
||||
loaded = array
|
||||
_eval_loaded_weight(loaded)
|
||||
weights[name] = loaded
|
||||
del array
|
||||
|
||||
blocks = []
|
||||
for block_index in range(num_blocks):
|
||||
block_weights = mlx_block_weights_from_diffusers_safetensors(
|
||||
checkpoint_path,
|
||||
block_index=block_index,
|
||||
quantization=spec,
|
||||
dtype=cast_dtype,
|
||||
)
|
||||
block_weights = {
|
||||
name: (value if isinstance(value, QuantizedMatrix) else value.astype(cast_dtype))
|
||||
for name, value in block_weights.items()
|
||||
}
|
||||
for value in block_weights.values():
|
||||
_eval_loaded_weight(value)
|
||||
blocks.append(
|
||||
MLXWanTransformerBlock(
|
||||
block_weights,
|
||||
dim=int(config["num_attention_heads"]) * int(config["attention_head_dim"]),
|
||||
ffn_dim=int(config["ffn_dim"]),
|
||||
num_heads=int(config["num_attention_heads"]),
|
||||
eps=float(config["eps"]),
|
||||
))
|
||||
return MLXWanDiT(weights, blocks, config, compile=compile)
|
||||
|
||||
|
||||
def torch_block_state_from_diffusers_safetensors(
|
||||
checkpoint_path: str | Path,
|
||||
*,
|
||||
block_index: int = 0,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Load one Diffusers-format Wan block into FastVideo's dense block keys."""
|
||||
from safetensors import safe_open
|
||||
|
||||
prefix = f"blocks.{block_index}."
|
||||
key_map = _WAN_BLOCK_KEY_MAP
|
||||
|
||||
state = {}
|
||||
with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle:
|
||||
available = set(handle.keys())
|
||||
for source_name, target_name in key_map.items():
|
||||
full = prefix + source_name
|
||||
if full not in available:
|
||||
# Biases are optional: e.g. Wan2.1-14B has bias-free attention/FFN.
|
||||
if source_name.endswith(".bias"):
|
||||
continue
|
||||
raise KeyError(f"missing required block weight: {full}")
|
||||
state[target_name] = handle.get_tensor(full).float()
|
||||
return state
|
||||
@@ -1,154 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pixel-space spatial resampling for decoded MLX Wan frames.
|
||||
|
||||
Spatial fast mode denoises on a smaller latent grid and has to get back to
|
||||
the requested output size. That resize belongs *here* — after the VAE
|
||||
decode — and not in latent space.
|
||||
|
||||
A Wan latent cell is a learned code for an 8x8 (Wan2.1) or 16x16 (Wan2.2)
|
||||
pixel block, not a low-pass sample of the image. Linearly blending two
|
||||
adjacent codes does not produce the code of the blended blocks; it produces
|
||||
a vector the decoder was never trained on. The decoder answers with smeared,
|
||||
ringing texture laid over otherwise-correct structure — the silhouette
|
||||
survives, the detail turns to haze. Measured on Wan2.1-1.3B at 480x832, a
|
||||
2x bilinear latent upsample destroys 62% of the latent's high-frequency
|
||||
energy while leaving its overall magnitude intact, which is exactly the
|
||||
signature of that veil.
|
||||
|
||||
Resampling decoded RGB frames has none of that problem: an image *is* a
|
||||
sampled 2-D signal, so Lanczos/cubic interpolation is the operation it was
|
||||
defined for. The result is soft — it carries stage-1's real detail budget
|
||||
and no more — but it is clean and coherent.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
|
||||
import numpy as np
|
||||
|
||||
# Pixel-space interpolation kernels, best-quality first. ``lanczos`` is the
|
||||
# default: it holds edges better than cubic at 2x with no visible ringing on
|
||||
# decoder output, which is already band-limited.
|
||||
PIXEL_UPSAMPLE_MODES = ("lanczos", "cubic", "bilinear", "nearest")
|
||||
|
||||
DEFAULT_PIXEL_UPSAMPLE_MODE = "lanczos"
|
||||
|
||||
|
||||
def _interpolation_flag(mode: str) -> int:
|
||||
"""
|
||||
Map a pixel upsample mode name onto its OpenCV interpolation flag.
|
||||
|
||||
Parameters:
|
||||
mode (str): One of :data:`PIXEL_UPSAMPLE_MODES`.
|
||||
|
||||
Returns:
|
||||
int: The matching ``cv2.INTER_*`` flag.
|
||||
|
||||
Raises:
|
||||
ValueError: If the mode is not a supported pixel upsample mode.
|
||||
"""
|
||||
import cv2
|
||||
|
||||
flags = {
|
||||
"lanczos": cv2.INTER_LANCZOS4,
|
||||
"cubic": cv2.INTER_CUBIC,
|
||||
"bilinear": cv2.INTER_LINEAR,
|
||||
"nearest": cv2.INTER_NEAREST,
|
||||
}
|
||||
try:
|
||||
return flags[mode]
|
||||
except KeyError:
|
||||
raise ValueError(f"Unsupported pixel upsample mode: {mode!r} "
|
||||
f"(expected one of {', '.join(PIXEL_UPSAMPLE_MODES)})") from None
|
||||
|
||||
|
||||
def unsharp(frame: np.ndarray, amount: float) -> np.ndarray:
|
||||
"""Light unsharp mask, used to counter resampling / optical-flow softening.
|
||||
|
||||
Parameters:
|
||||
frame (np.ndarray): HxWx3 uint8 RGB frame.
|
||||
amount (float): Strength; ``0`` returns the frame unchanged.
|
||||
|
||||
Returns:
|
||||
np.ndarray: A new frame; the input is never modified in place.
|
||||
"""
|
||||
if amount <= 0.0:
|
||||
return frame
|
||||
import cv2
|
||||
|
||||
blur = cv2.GaussianBlur(frame, (0, 0), 1.0)
|
||||
return cv2.addWeighted(frame, 1.0 + amount, blur, -amount, 0)
|
||||
|
||||
|
||||
def upsample_frame(
|
||||
frame: np.ndarray,
|
||||
*,
|
||||
width: int,
|
||||
height: int,
|
||||
mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
sharpen: float = 0.0,
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Resample one decoded RGB frame to the target pixel size.
|
||||
|
||||
Parameters:
|
||||
frame (np.ndarray): HxWx3 uint8 RGB frame.
|
||||
width (int): Target width in pixels.
|
||||
height (int): Target height in pixels.
|
||||
mode (str): Interpolation kernel, one of :data:`PIXEL_UPSAMPLE_MODES`.
|
||||
sharpen (float): Unsharp strength applied after the resize.
|
||||
|
||||
Returns:
|
||||
np.ndarray: A new frame at ``height x width``; already-correct sizes
|
||||
are still passed through ``sharpen``.
|
||||
|
||||
Raises:
|
||||
ValueError: If the frame is not HxWx3, or the target size is not positive.
|
||||
"""
|
||||
import cv2
|
||||
|
||||
array = np.asarray(frame)
|
||||
if array.ndim != 3 or array.shape[2] != 3:
|
||||
raise ValueError(f"frame must have shape HxWx3, got {array.shape}")
|
||||
if width <= 0 or height <= 0:
|
||||
raise ValueError(f"target size must be positive, got {width}x{height}")
|
||||
if array.dtype != np.uint8:
|
||||
array = np.clip(array, 0, 255).astype(np.uint8)
|
||||
|
||||
if (array.shape[0], array.shape[1]) != (height, width):
|
||||
array = cv2.resize(array, (width, height), interpolation=_interpolation_flag(mode))
|
||||
return unsharp(array, sharpen)
|
||||
|
||||
|
||||
def upsample_frames(
|
||||
frames: Iterable[np.ndarray],
|
||||
*,
|
||||
width: int,
|
||||
height: int,
|
||||
mode: str = DEFAULT_PIXEL_UPSAMPLE_MODE,
|
||||
sharpen: float = 0.0,
|
||||
) -> list[np.ndarray]:
|
||||
"""
|
||||
Resample every decoded frame to the target pixel size.
|
||||
|
||||
Parameters:
|
||||
frames (Iterable[np.ndarray]): Decoded HxWx3 uint8 RGB frames.
|
||||
width (int): Target width in pixels.
|
||||
height (int): Target height in pixels.
|
||||
mode (str): Interpolation kernel, one of :data:`PIXEL_UPSAMPLE_MODES`.
|
||||
sharpen (float): Unsharp strength applied after each resize.
|
||||
|
||||
Returns:
|
||||
list[np.ndarray]: New frames at the target size, in input order.
|
||||
"""
|
||||
return [upsample_frame(frame, width=width, height=height, mode=mode, sharpen=sharpen) for frame in frames]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_PIXEL_UPSAMPLE_MODE",
|
||||
"PIXEL_UPSAMPLE_MODES",
|
||||
"unsharp",
|
||||
"upsample_frame",
|
||||
"upsample_frames",
|
||||
]
|
||||
@@ -1,243 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Memory-tier helpers for Apple Silicon MLX/MPS experiments.
|
||||
|
||||
macOS does not expose a perfect "pretend this machine only has 16 GB unified
|
||||
memory" switch. MLX can cap the allocator used by the Apple-native DiT path,
|
||||
and PyTorch MPS exposes process-level watermark environment variables for the
|
||||
hybrid prompt/decode stages. Applying both gives benchmark and generation
|
||||
entrypoints a practical, explicit way to exercise memory-tier presets.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import gc
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
GIB = 1024**3
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AppliedMemoryLimits:
|
||||
"""Memory limits applied for one Apple Silicon benchmark/generation process."""
|
||||
|
||||
mlx_memory_limit_gib: float | None = None
|
||||
mlx_cache_limit_gib: float | None = None
|
||||
mlx_disable_cache: bool = False
|
||||
mlx_wired_limit_gib: float | None = None
|
||||
torch_mps_high_watermark_ratio: float | None = None
|
||||
torch_mps_low_watermark_ratio: float | None = None
|
||||
applied_bytes: dict[str, int] = field(default_factory=dict)
|
||||
previous_bytes: dict[str, int] = field(default_factory=dict)
|
||||
errors: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
def as_metrics(self) -> dict[str, int | float | str | bool | None]:
|
||||
"""Flatten the configured memory limits, applied values, previous values, and errors into a metrics dictionary.
|
||||
|
||||
Returns:
|
||||
dict[str, int | float | str | bool | None]: Metrics keyed by limit names and their corresponding values.
|
||||
"""
|
||||
metrics: dict[str, int | float | str | bool | None] = {
|
||||
"mlx_memory_limit_gib": self.mlx_memory_limit_gib,
|
||||
"mlx_cache_limit_gib": self.mlx_cache_limit_gib,
|
||||
"mlx_disable_cache": self.mlx_disable_cache,
|
||||
"mlx_wired_limit_gib": self.mlx_wired_limit_gib,
|
||||
"torch_mps_high_watermark_ratio": self.torch_mps_high_watermark_ratio,
|
||||
"torch_mps_low_watermark_ratio": self.torch_mps_low_watermark_ratio,
|
||||
}
|
||||
for name, value in self.applied_bytes.items():
|
||||
metrics[f"{name}_bytes"] = value
|
||||
for name, value in self.previous_bytes.items():
|
||||
metrics[f"previous_{name}_bytes"] = value
|
||||
for name, error in self.errors.items():
|
||||
metrics[f"{name}_error"] = error
|
||||
return metrics
|
||||
|
||||
|
||||
def gib_to_bytes(value: float | None) -> int | None:
|
||||
"""
|
||||
Convert a positive memory limit from GiB to bytes.
|
||||
|
||||
Parameters:
|
||||
value (float | None): Memory limit in GiB, or `None` when unset.
|
||||
|
||||
Returns:
|
||||
int | None: The memory limit in bytes, or `None` when no limit is provided.
|
||||
|
||||
Raises:
|
||||
ValueError: If `value` is zero or negative.
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
if value <= 0:
|
||||
raise ValueError(f"Memory limit must be positive GiB, got {value}")
|
||||
return int(value * GIB)
|
||||
|
||||
|
||||
def cleanup_mlx(mx_module: Any | None = None) -> None:
|
||||
"""Collect unreachable MLX objects, then release their allocator cache."""
|
||||
if mx_module is None:
|
||||
import mlx.core as mx
|
||||
|
||||
mx_module = mx
|
||||
gc.collect()
|
||||
mx_module.clear_cache()
|
||||
|
||||
|
||||
def cleanup_torch_mps(torch_module: Any | None = None) -> None:
|
||||
"""Collect unreachable Torch objects, then release the MPS allocator cache."""
|
||||
if torch_module is None:
|
||||
import torch
|
||||
|
||||
torch_module = torch
|
||||
gc.collect()
|
||||
if torch_module.backends.mps.is_available():
|
||||
torch_module.mps.empty_cache()
|
||||
|
||||
|
||||
def _set_mps_env(name: str, value: float | None) -> float | None:
|
||||
"""Set a PyTorch MPS watermark environment variable.
|
||||
|
||||
Parameters:
|
||||
name (str): Name of the environment variable to set.
|
||||
value (float | None): Watermark ratio, or `None` to leave the variable unchanged.
|
||||
|
||||
Returns:
|
||||
float | None: The configured watermark ratio, or `None` when no value is provided.
|
||||
|
||||
Raises:
|
||||
ValueError: If `value` is negative.
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
if value < 0:
|
||||
raise ValueError(f"{name} must be non-negative, got {value}")
|
||||
os.environ[name] = str(value)
|
||||
return value
|
||||
|
||||
|
||||
def apply_memory_limits(
|
||||
*,
|
||||
mlx_memory_limit_gib: float | None = None,
|
||||
mlx_cache_limit_gib: float | None = None,
|
||||
mlx_disable_cache: bool = False,
|
||||
mlx_wired_limit_gib: float | None = None,
|
||||
torch_mps_high_watermark_ratio: float | None = None,
|
||||
torch_mps_low_watermark_ratio: float | None = None,
|
||||
mx_module: Any | None = None,
|
||||
) -> AppliedMemoryLimits:
|
||||
"""Apply optional MLX allocator limits and PyTorch MPS watermarks.
|
||||
|
||||
PyTorch reads MPS watermark variables when the MPS backend initializes, so
|
||||
call this before importing PyTorch. Specifying only a high watermark sets the
|
||||
low watermark to ``0.0``. MLX limit-setting failures are recorded in the
|
||||
result and do not prevent other limits from being applied.
|
||||
|
||||
Parameters:
|
||||
mlx_memory_limit_gib (float | None): Maximum MLX memory in GiB.
|
||||
mlx_cache_limit_gib (float | None): Maximum MLX cache size in GiB.
|
||||
mlx_disable_cache (bool): Whether to disable the MLX cache.
|
||||
mlx_wired_limit_gib (float | None): Maximum MLX wired memory in GiB.
|
||||
torch_mps_high_watermark_ratio (float | None): PyTorch MPS high watermark
|
||||
ratio.
|
||||
torch_mps_low_watermark_ratio (float | None): PyTorch MPS low watermark
|
||||
ratio.
|
||||
|
||||
Returns:
|
||||
AppliedMemoryLimits: Configured values, applied and previous MLX byte
|
||||
limits, MPS watermark values, and per-limit errors.
|
||||
"""
|
||||
if torch_mps_high_watermark_ratio is not None and torch_mps_low_watermark_ratio is None:
|
||||
torch_mps_low_watermark_ratio = 0.0
|
||||
|
||||
high = _set_mps_env("PYTORCH_MPS_HIGH_WATERMARK_RATIO", torch_mps_high_watermark_ratio)
|
||||
low = _set_mps_env("PYTORCH_MPS_LOW_WATERMARK_RATIO", torch_mps_low_watermark_ratio)
|
||||
|
||||
memory_bytes = gib_to_bytes(mlx_memory_limit_gib)
|
||||
cache_bytes = 0 if mlx_disable_cache else gib_to_bytes(mlx_cache_limit_gib)
|
||||
wired_bytes = gib_to_bytes(mlx_wired_limit_gib)
|
||||
|
||||
applied: dict[str, int] = {}
|
||||
previous: dict[str, int] = {}
|
||||
errors: dict[str, str] = {}
|
||||
if memory_bytes is not None or cache_bytes is not None or wired_bytes is not None:
|
||||
if mx_module is None:
|
||||
import mlx.core as mx
|
||||
|
||||
mx_module = mx
|
||||
|
||||
# Apply each limit independently; record failures without stopping.
|
||||
limits = [
|
||||
("mlx_memory_limit", memory_bytes, mx_module.set_memory_limit),
|
||||
("mlx_cache_limit", cache_bytes, mx_module.set_cache_limit),
|
||||
("mlx_wired_limit", wired_bytes, mx_module.set_wired_limit),
|
||||
]
|
||||
for name, value, setter in limits:
|
||||
if value is not None:
|
||||
try:
|
||||
previous[name] = int(setter(value))
|
||||
applied[name] = value
|
||||
except Exception as exc: # noqa: BLE001 - macOS/system-limit dependent.
|
||||
errors[name] = f"{type(exc).__name__}: {exc}"
|
||||
|
||||
return AppliedMemoryLimits(
|
||||
mlx_memory_limit_gib=mlx_memory_limit_gib,
|
||||
mlx_cache_limit_gib=mlx_cache_limit_gib,
|
||||
mlx_disable_cache=mlx_disable_cache,
|
||||
mlx_wired_limit_gib=mlx_wired_limit_gib,
|
||||
torch_mps_high_watermark_ratio=high,
|
||||
torch_mps_low_watermark_ratio=low,
|
||||
applied_bytes=applied,
|
||||
previous_bytes=previous,
|
||||
errors=errors,
|
||||
)
|
||||
|
||||
|
||||
def add_memory_limit_args(
|
||||
parser: argparse.ArgumentParser,
|
||||
*,
|
||||
mlx_memory_limit_gib: float | None = None,
|
||||
mlx_cache_limit_gib: float | None = None,
|
||||
mlx_disable_cache: bool = False,
|
||||
mlx_wired_limit_gib: float | None = None,
|
||||
torch_mps_high_watermark_ratio: float | None = None,
|
||||
torch_mps_low_watermark_ratio: float | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Add configurable Apple Silicon memory-limit options to an argument parser.
|
||||
|
||||
Parameters:
|
||||
parser (argparse.ArgumentParser): Parser to which the options are added.
|
||||
mlx_memory_limit_gib (float | None): Default MLX memory limit in GiB.
|
||||
mlx_cache_limit_gib (float | None): Default MLX cache limit in GiB.
|
||||
mlx_disable_cache (bool): Whether the cache limit defaults to zero.
|
||||
mlx_wired_limit_gib (float | None): Default MLX wired-memory limit in GiB.
|
||||
torch_mps_high_watermark_ratio (float | None): Default PyTorch MPS high-watermark ratio.
|
||||
torch_mps_low_watermark_ratio (float | None): Default PyTorch MPS low-watermark ratio.
|
||||
"""
|
||||
parser.add_argument("--mlx-memory-limit-gib",
|
||||
type=float,
|
||||
default=mlx_memory_limit_gib,
|
||||
help="Set MLX memory limit in GiB for memory-tier testing (DiT path).")
|
||||
parser.add_argument("--mlx-cache-limit-gib",
|
||||
type=float,
|
||||
default=mlx_cache_limit_gib,
|
||||
help="Set MLX cache limit in GiB. Use --mlx-disable-cache to force 0.")
|
||||
parser.add_argument("--mlx-disable-cache",
|
||||
action="store_true",
|
||||
default=mlx_disable_cache,
|
||||
help="Set MLX cache limit to 0 for stricter memory-tier tests.")
|
||||
parser.add_argument("--mlx-wired-limit-gib",
|
||||
type=float,
|
||||
default=mlx_wired_limit_gib,
|
||||
help="Set MLX wired-memory limit in GiB where supported by macOS/MLX.")
|
||||
parser.add_argument("--torch-mps-high-watermark-ratio",
|
||||
type=float,
|
||||
default=torch_mps_high_watermark_ratio,
|
||||
help="Set PYTORCH_MPS_HIGH_WATERMARK_RATIO before importing torch.")
|
||||
parser.add_argument("--torch-mps-low-watermark-ratio",
|
||||
type=float,
|
||||
default=torch_mps_low_watermark_ratio,
|
||||
help="Set PYTORCH_MPS_LOW_WATERMARK_RATIO before importing torch.")
|
||||
@@ -1,129 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Best-effort prompt-embedding cache shared by the MLX entrypoints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def fingerprint_digest(fingerprint: dict[str, object]) -> str:
|
||||
payload = json.dumps(fingerprint, sort_keys=True, separators=(",", ":"))
|
||||
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def text_encoder_fingerprint(model_root: Path) -> dict[str, object]:
|
||||
"""Return a cheap identity for the tokenizer and text-encoder files."""
|
||||
root = model_root.resolve()
|
||||
components = [path for name in ("tokenizer", "text_encoder") if (path := root / name).is_dir()]
|
||||
scan_roots = components or [root]
|
||||
files: list[list[object]] = []
|
||||
complete = True
|
||||
try:
|
||||
for scan_root in scan_roots:
|
||||
for path in sorted(scan_root.rglob("*")):
|
||||
try:
|
||||
if not path.is_file():
|
||||
continue
|
||||
stat = path.stat()
|
||||
files.append([
|
||||
path.relative_to(root).as_posix(),
|
||||
stat.st_size,
|
||||
stat.st_mtime_ns,
|
||||
stat.st_ctime_ns,
|
||||
])
|
||||
except OSError:
|
||||
complete = False
|
||||
except OSError:
|
||||
complete = False
|
||||
# ponytail: metadata avoids hashing multi-GB weights; use a model manifest
|
||||
# if supported workflows ever preserve size, mtime, and ctime while mutating.
|
||||
return {"root": str(root), "files": files, "complete": complete}
|
||||
|
||||
|
||||
def prompt_cache_meta_path(cache_path: Path) -> Path:
|
||||
return cache_path.with_suffix(cache_path.suffix + ".json")
|
||||
|
||||
|
||||
def _fingerprint_is_complete(fingerprint: dict[str, object]) -> bool:
|
||||
text_encoder = fingerprint.get("text_encoder")
|
||||
return not isinstance(text_encoder, dict) or text_encoder.get("complete") is not False
|
||||
|
||||
|
||||
def load_prompt_cache(
|
||||
cache_path: Path | None,
|
||||
fingerprint: dict[str, object],
|
||||
) -> np.ndarray | None:
|
||||
"""Load a matching cache entry, treating every cache failure as a miss."""
|
||||
if cache_path is None or not _fingerprint_is_complete(fingerprint):
|
||||
return None
|
||||
try:
|
||||
metadata = json.loads(prompt_cache_meta_path(cache_path).read_text())
|
||||
if not isinstance(metadata, dict):
|
||||
return None
|
||||
if metadata.get("fingerprint_sha256") != fingerprint_digest(fingerprint):
|
||||
return None
|
||||
payload = cache_path.read_bytes()
|
||||
if metadata.get("data_sha256") != hashlib.sha256(payload).hexdigest():
|
||||
return None
|
||||
array = np.load(io.BytesIO(payload), allow_pickle=False)
|
||||
return array if isinstance(array, np.ndarray) else None
|
||||
except (EOFError, OSError, UnicodeError, ValueError) as exc:
|
||||
logger.info("Prompt cache read skipped for %s: %s", cache_path, exc)
|
||||
return None
|
||||
|
||||
|
||||
def _atomic_write(path: Path, payload: bytes) -> None:
|
||||
temp_path: Path | None = None
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="wb",
|
||||
dir=path.parent,
|
||||
prefix=f".{path.name}.",
|
||||
suffix=".tmp",
|
||||
delete=False,
|
||||
) as handle:
|
||||
temp_path = Path(handle.name)
|
||||
handle.write(payload)
|
||||
temp_path.replace(path)
|
||||
finally:
|
||||
if temp_path is not None:
|
||||
temp_path.unlink(missing_ok=True)
|
||||
|
||||
|
||||
def save_prompt_cache(
|
||||
cache_path: Path | None,
|
||||
embeds: np.ndarray,
|
||||
fingerprint: dict[str, object],
|
||||
) -> bool:
|
||||
"""Atomically publish an integrity-bound cache entry when possible."""
|
||||
if cache_path is None or not _fingerprint_is_complete(fingerprint):
|
||||
return False
|
||||
try:
|
||||
cache_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
buffer = io.BytesIO()
|
||||
np.save(buffer, np.asarray(embeds), allow_pickle=False)
|
||||
payload = buffer.getvalue()
|
||||
metadata = (json.dumps(
|
||||
{
|
||||
"fingerprint_sha256": fingerprint_digest(fingerprint),
|
||||
"data_sha256": hashlib.sha256(payload).hexdigest(),
|
||||
"fingerprint": fingerprint,
|
||||
},
|
||||
indent=2) + "\n").encode("utf-8")
|
||||
# Publish data first. Until metadata follows, old metadata's data digest
|
||||
# makes the torn pair a harmless miss rather than a stale cache hit.
|
||||
_atomic_write(cache_path, payload)
|
||||
_atomic_write(prompt_cache_meta_path(cache_path), metadata)
|
||||
return True
|
||||
except (OSError, TypeError, ValueError) as exc:
|
||||
logger.info("Prompt cache write skipped for %s: %s", cache_path, exc)
|
||||
return False
|
||||
@@ -1,454 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Local prompt enrichment for the MLX Wan runtime (H3 Context-IR-style).
|
||||
|
||||
Wan's training captions are long and cinematic; short user prompts leave
|
||||
quality on the table. This module expands a raw prompt into Wan-style
|
||||
shot language **on device** — no remote API, no training.
|
||||
|
||||
Backends (first match wins):
|
||||
|
||||
1. **mlx-lm** — optional local LLM (``--enhance-prompt-model``).
|
||||
2. **template** — deterministic cinematic expansion (always available).
|
||||
|
||||
System-prompt contract matches the streaming server's enhancer defaults
|
||||
in ``fastvideo/entrypoints/streaming/prompt/enhancer.py`` so remote and
|
||||
local paths stay interchangeable.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Keep in lockstep with streaming PromptEnhancer defaults (enhance op).
|
||||
DEFAULT_ENHANCE_SYSTEM_PROMPT = ("You are a prompt enhancer for cinematic video generation. Given "
|
||||
"a user prompt, produce an enhanced prompt that is more vivid, "
|
||||
"specific, and concrete. Keep the subject intact; add lighting, "
|
||||
"camera, and motion detail. Reply with just the enhanced prompt.")
|
||||
|
||||
# Small default that fits 16 GB Macs alongside the 1.3B DiT when the user
|
||||
# opts into mlx-lm. Override with --enhance-prompt-model.
|
||||
DEFAULT_MLX_LM_MODEL = "mlx-community/Qwen2.5-0.5B-Instruct-4bit"
|
||||
|
||||
_CAMERA_CUES = (
|
||||
"cinematic",
|
||||
"camera",
|
||||
"lens",
|
||||
"shot",
|
||||
"bokeh",
|
||||
"dolly",
|
||||
"tracking",
|
||||
"close-up",
|
||||
"wide shot",
|
||||
"handheld",
|
||||
"steadicam",
|
||||
)
|
||||
_LIGHT_CUES = (
|
||||
"light",
|
||||
"lighting",
|
||||
"sun",
|
||||
"golden hour",
|
||||
"neon",
|
||||
"rim light",
|
||||
"softbox",
|
||||
"overcast",
|
||||
"moonlight",
|
||||
"volumetric",
|
||||
)
|
||||
_MOTION_CUES = (
|
||||
"moving",
|
||||
"motion",
|
||||
"walk",
|
||||
"run",
|
||||
"flies",
|
||||
"flying",
|
||||
"drifts",
|
||||
"sails",
|
||||
"flows",
|
||||
"pan",
|
||||
"tilt",
|
||||
"zoom",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EnhanceResult:
|
||||
"""Outcome of a prompt enrichment call."""
|
||||
|
||||
original: str
|
||||
enhanced: str
|
||||
backend: str
|
||||
elapsed_s: float
|
||||
model: str | None = None
|
||||
|
||||
@property
|
||||
def changed(self) -> bool:
|
||||
"""Indicates whether the enhanced prompt differs from the original after trimming surrounding whitespace.
|
||||
|
||||
Returns:
|
||||
bool: `True` if the prompts differ, `False` otherwise.
|
||||
"""
|
||||
return self.enhanced.strip() != self.original.strip()
|
||||
|
||||
|
||||
def _normalize_user_prompt(prompt: str) -> str:
|
||||
"""
|
||||
Normalize a user prompt for enhancement.
|
||||
|
||||
Parameters:
|
||||
prompt (str): User-provided prompt text.
|
||||
|
||||
Returns:
|
||||
str: The prompt with leading and trailing whitespace removed and internal whitespace collapsed.
|
||||
|
||||
Raises:
|
||||
ValueError: If the prompt is empty after whitespace normalization.
|
||||
"""
|
||||
text = " ".join(prompt.strip().split())
|
||||
if not text:
|
||||
raise ValueError("prompt must be non-empty")
|
||||
return text
|
||||
|
||||
|
||||
def _already_rich(prompt: str) -> bool:
|
||||
"""
|
||||
Determine whether a prompt already contains substantial camera and lighting detail.
|
||||
|
||||
Returns:
|
||||
bool: `true` if the prompt is at least 160 characters long and includes camera and lighting cues, `false` otherwise.
|
||||
"""
|
||||
lower = prompt.lower()
|
||||
has_camera = any(c in lower for c in _CAMERA_CUES)
|
||||
has_light = any(c in lower for c in _LIGHT_CUES)
|
||||
return len(prompt) >= 160 and has_camera and has_light
|
||||
|
||||
|
||||
def enhance_prompt_template(prompt: str) -> str:
|
||||
"""
|
||||
Expand a prompt with cinematic camera, lighting, motion, and visual-quality details.
|
||||
|
||||
Rich prompts are preserved, while thinner prompts receive deterministic enhancements
|
||||
without changing their subject.
|
||||
|
||||
Returns:
|
||||
str: The original or expanded prompt with normalized whitespace and punctuation.
|
||||
"""
|
||||
text = _normalize_user_prompt(prompt)
|
||||
if _already_rich(text):
|
||||
return text
|
||||
|
||||
lower = text.lower()
|
||||
parts = [text.rstrip(".")]
|
||||
|
||||
if not any(c in lower for c in _CAMERA_CUES):
|
||||
parts.append("shot on a 35mm anamorphic lens, gentle handheld micro-movement, "
|
||||
"shallow depth of field")
|
||||
if not any(c in lower for c in _LIGHT_CUES):
|
||||
parts.append("natural cinematic lighting with soft volumetric haze and subtle "
|
||||
"rim light separating subject from background")
|
||||
if not any(c in lower for c in _MOTION_CUES):
|
||||
parts.append("smooth continuous motion with grounded physics")
|
||||
|
||||
parts.append("highly detailed, coherent temporal continuity, film grain, "
|
||||
"color graded like a contemporary drama")
|
||||
enhanced = ", ".join(parts)
|
||||
# Single trailing period; collapse duplicate whitespace.
|
||||
enhanced = re.sub(r"\s+", " ", enhanced).strip()
|
||||
if not enhanced.endswith("."):
|
||||
enhanced += "."
|
||||
return enhanced
|
||||
|
||||
|
||||
def enhance_prompt_mlx_lm(
|
||||
prompt: str,
|
||||
*,
|
||||
model: str = DEFAULT_MLX_LM_MODEL,
|
||||
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
|
||||
max_tokens: int = 128,
|
||||
temp: float = 0.6,
|
||||
) -> str:
|
||||
"""
|
||||
Enhance a user prompt with a locally hosted mlx-lm instruction model.
|
||||
|
||||
Parameters:
|
||||
prompt (str): The prompt to enhance.
|
||||
model (str): The mlx-lm model identifier or path.
|
||||
system_prompt (str): Instructions that guide prompt enhancement.
|
||||
max_tokens (int): Maximum number of tokens to generate.
|
||||
temp (float): Sampling temperature for generation.
|
||||
|
||||
Returns:
|
||||
str: The enhanced prompt.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If mlx-lm is unavailable or produces an empty result.
|
||||
"""
|
||||
try:
|
||||
from mlx_lm import generate, load
|
||||
except ImportError as exc: # pragma: no cover - optional dep
|
||||
raise RuntimeError("mlx-lm is not installed. `uv pip install mlx-lm` or use "
|
||||
"--enhance-prompt-backend template.") from exc
|
||||
|
||||
text = _normalize_user_prompt(prompt)
|
||||
logger.info("[MLX enhance] loading %s", model)
|
||||
mlx_model, tokenizer = load(model)
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": system_prompt
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": text
|
||||
},
|
||||
]
|
||||
if hasattr(tokenizer, "apply_chat_template"):
|
||||
chat = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
else: # pragma: no cover - ancient tokenizers
|
||||
chat = f"{system_prompt}\n\nUser: {text}\nAssistant:"
|
||||
|
||||
raw = generate(
|
||||
mlx_model,
|
||||
tokenizer,
|
||||
prompt=chat,
|
||||
max_tokens=max_tokens,
|
||||
temp=temp,
|
||||
verbose=False,
|
||||
)
|
||||
enhanced = _clean_llm_output(raw, original=text)
|
||||
if not enhanced:
|
||||
raise RuntimeError("mlx-lm returned an empty enhance result")
|
||||
return enhanced
|
||||
|
||||
|
||||
def _clean_llm_output(raw: str, *, original: str) -> str:
|
||||
"""
|
||||
Clean generated prompt text and fall back to the original when the result is too short.
|
||||
|
||||
Parameters:
|
||||
raw (str): Raw text produced by the language model.
|
||||
original (str): Original prompt used as the fallback value.
|
||||
|
||||
Returns:
|
||||
str: Cleaned first paragraph of the generated text, or the original prompt when the generated text is too short.
|
||||
"""
|
||||
text = raw.strip()
|
||||
# Drop common prefatory phrases.
|
||||
for prefix in (
|
||||
"enhanced prompt:",
|
||||
"here's the enhanced prompt:",
|
||||
"here is the enhanced prompt:",
|
||||
"sure:",
|
||||
"sure,",
|
||||
):
|
||||
if text.lower().startswith(prefix):
|
||||
text = text[len(prefix):].strip()
|
||||
# Keep first non-empty paragraph only.
|
||||
para = text.split("\n\n")[0].strip()
|
||||
para = " ".join(para.split())
|
||||
if len(para) < max(12, len(original) // 4):
|
||||
return original
|
||||
return para
|
||||
|
||||
|
||||
def enhance_prompt(
|
||||
prompt: str,
|
||||
*,
|
||||
backend: str = "auto",
|
||||
model: str | None = None,
|
||||
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
|
||||
max_tokens: int = 128,
|
||||
) -> EnhanceResult:
|
||||
"""Enhance a prompt using the selected backend, falling back to a deterministic template when configured for automatic selection.
|
||||
|
||||
Parameters:
|
||||
prompt (str): The prompt to enhance.
|
||||
backend (str): The enhancement backend: ``"auto"``, ``"mlx-lm"``, or ``"template"``.
|
||||
model (str | None): The MLX language model to use.
|
||||
system_prompt (str): Instructions provided to the MLX language model.
|
||||
max_tokens (int): Maximum number of tokens generated by the MLX language model.
|
||||
|
||||
Returns:
|
||||
EnhanceResult: The original and enhanced prompts, selected backend, timing information, and model metadata.
|
||||
|
||||
Raises:
|
||||
ValueError: If the prompt is empty or the backend is unsupported.
|
||||
Exception: If the explicitly selected ``"mlx-lm"`` backend fails.
|
||||
"""
|
||||
text = _normalize_user_prompt(prompt)
|
||||
backend_norm = (backend or "auto").lower()
|
||||
if backend_norm not in {"auto", "mlx-lm", "template"}:
|
||||
raise ValueError(f"Unknown enhance backend: {backend}")
|
||||
|
||||
start = time.perf_counter()
|
||||
used_model: str | None = None
|
||||
|
||||
if backend_norm in {"auto", "mlx-lm"}:
|
||||
try:
|
||||
used_model = model or DEFAULT_MLX_LM_MODEL
|
||||
enhanced = enhance_prompt_mlx_lm(
|
||||
text,
|
||||
model=used_model,
|
||||
system_prompt=system_prompt,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
return EnhanceResult(
|
||||
original=text,
|
||||
enhanced=enhanced,
|
||||
backend="mlx-lm",
|
||||
elapsed_s=time.perf_counter() - start,
|
||||
model=used_model,
|
||||
)
|
||||
except Exception as exc:
|
||||
if backend_norm == "mlx-lm":
|
||||
raise
|
||||
logger.info(
|
||||
"[MLX enhance] mlx-lm unavailable (%s); using template backend",
|
||||
exc,
|
||||
)
|
||||
|
||||
enhanced = enhance_prompt_template(text)
|
||||
return EnhanceResult(
|
||||
original=text,
|
||||
enhanced=enhanced,
|
||||
backend="template",
|
||||
elapsed_s=time.perf_counter() - start,
|
||||
model=None,
|
||||
)
|
||||
|
||||
|
||||
def enhance_cache_path(
|
||||
prompt: str,
|
||||
*,
|
||||
backend: str,
|
||||
model: str | None,
|
||||
cache_dir: Path | None = None,
|
||||
) -> Path:
|
||||
"""Content-addressed cache file for an enhanced prompt string."""
|
||||
root = cache_dir or (Path.home() / ".cache" / "fastvideo" / "enhanced_prompts")
|
||||
key = "\0".join([prompt, backend, model or ""])
|
||||
digest = hashlib.sha256(key.encode("utf-8")).hexdigest()[:24]
|
||||
return root / f"{digest}.json"
|
||||
|
||||
|
||||
def load_or_enhance_prompt(
|
||||
prompt: str,
|
||||
*,
|
||||
backend: str = "auto",
|
||||
model: str | None = None,
|
||||
system_prompt: str = DEFAULT_ENHANCE_SYSTEM_PROMPT,
|
||||
max_tokens: int = 128,
|
||||
cache: bool = True,
|
||||
cache_dir: Path | None = None,
|
||||
) -> EnhanceResult:
|
||||
"""
|
||||
Enhance a prompt, reusing a cached result when available.
|
||||
|
||||
Parameters:
|
||||
prompt (str): The prompt to enhance.
|
||||
backend (str): Enhancement backend to use.
|
||||
model (str | None): Optional model identifier.
|
||||
system_prompt (str): System prompt for model-based enhancement.
|
||||
max_tokens (int): Maximum number of tokens generated by the model.
|
||||
cache (bool): Whether to read and write the on-disk cache.
|
||||
cache_dir (Path | None): Optional directory for cached results.
|
||||
|
||||
Returns:
|
||||
EnhanceResult: The enhanced prompt and backend metadata. Cached results are marked with the ``"cache"`` backend.
|
||||
"""
|
||||
text = _normalize_user_prompt(prompt)
|
||||
path = enhance_cache_path(text, backend=backend, model=model, cache_dir=cache_dir)
|
||||
if cache and path.is_file():
|
||||
try:
|
||||
payload = json.loads(path.read_text())
|
||||
return EnhanceResult(
|
||||
original=str(payload.get("original", text)),
|
||||
enhanced=str(payload["enhanced"]),
|
||||
# Mark cache hits explicitly so metrics/logs can distinguish
|
||||
# a free replay from a fresh template/mlx-lm call.
|
||||
backend="cache",
|
||||
elapsed_s=0.0,
|
||||
model=payload.get("model"),
|
||||
)
|
||||
except (OSError, KeyError, json.JSONDecodeError):
|
||||
pass
|
||||
|
||||
result = enhance_prompt(
|
||||
text,
|
||||
backend=backend,
|
||||
model=model,
|
||||
system_prompt=system_prompt,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
if cache:
|
||||
try:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"original": result.original,
|
||||
"enhanced": result.enhanced,
|
||||
"backend": result.backend,
|
||||
"model": result.model,
|
||||
},
|
||||
indent=2,
|
||||
))
|
||||
except OSError as exc: # pragma: no cover - cache is best-effort
|
||||
logger.info("[MLX enhance] cache write skipped: %s", exc)
|
||||
return result
|
||||
|
||||
|
||||
def enhance_result_as_metrics(result: EnhanceResult | None) -> dict[str, Any]:
|
||||
"""
|
||||
Convert prompt enhancement results into metrics fields.
|
||||
|
||||
Parameters:
|
||||
result (EnhanceResult | None): The enhancement result, or `None` when no enhancement was performed.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: A metrics mapping containing enhancement status, backend metadata, timing, and original and enhanced prompts.
|
||||
"""
|
||||
if result is None:
|
||||
return {
|
||||
"enhance_prompt": False,
|
||||
"enhance_backend": None,
|
||||
"enhance_model": None,
|
||||
"enhance_elapsed_s": None,
|
||||
"prompt_original": None,
|
||||
"prompt_enhanced": None,
|
||||
}
|
||||
return {
|
||||
"enhance_prompt": True,
|
||||
"enhance_backend": result.backend,
|
||||
"enhance_model": result.model,
|
||||
"enhance_elapsed_s": result.elapsed_s,
|
||||
"prompt_original": result.original,
|
||||
"prompt_enhanced": result.enhanced,
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_ENHANCE_SYSTEM_PROMPT",
|
||||
"DEFAULT_MLX_LM_MODEL",
|
||||
"EnhanceResult",
|
||||
"enhance_cache_path",
|
||||
"enhance_prompt",
|
||||
"enhance_prompt_mlx_lm",
|
||||
"enhance_prompt_template",
|
||||
"enhance_result_as_metrics",
|
||||
"load_or_enhance_prompt",
|
||||
]
|
||||
@@ -1,284 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MLX block-scaled quantization backends (affine INT8, MXFP8/4, NVFP4).
|
||||
|
||||
Isolated experiment module: probes which ``mx.quantize`` modes the installed
|
||||
MLX build supports and exposes a thin wrapper around native quantized matmul.
|
||||
Depends only on ``mlx.core`` and the standard library — do not import the rest
|
||||
of FastVideo from here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Final
|
||||
from collections.abc import Mapping
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
# Probe matrix side length: must be divisible by every mode's group size
|
||||
# (affine g64, mxfp* g32, nvfp4 g16).
|
||||
_PROBE_DIM: Final[int] = 64
|
||||
|
||||
|
||||
class QuantBackend(str, Enum):
|
||||
"""Named MLX quantization backends evaluated for M5 Neural Accelerators."""
|
||||
|
||||
AFFINE_INT8_G64 = "affine_int8_g64"
|
||||
MXFP8 = "mxfp8"
|
||||
MXFP4 = "mxfp4"
|
||||
NVFP4 = "nvfp4"
|
||||
|
||||
|
||||
BACKENDS: Final[tuple[str, ...]] = tuple(b.value for b in QuantBackend)
|
||||
|
||||
# Backend name -> kwargs for mx.quantize / mx.quantized_matmul.
|
||||
# Affine baseline matches FastVideo DiT load path (INT8, group size 64).
|
||||
# MX/NV block-scaled modes use MLX defaults (see mx.quantize docs).
|
||||
_BACKEND_KWARGS: Final[Mapping[str, Mapping[str, object]]] = {
|
||||
QuantBackend.AFFINE_INT8_G64.value: {
|
||||
"mode": "affine",
|
||||
"bits": 8,
|
||||
"group_size": 64,
|
||||
},
|
||||
QuantBackend.MXFP8.value: {
|
||||
"mode": "mxfp8",
|
||||
"bits": None,
|
||||
"group_size": None,
|
||||
},
|
||||
QuantBackend.MXFP4.value: {
|
||||
"mode": "mxfp4",
|
||||
"bits": None,
|
||||
"group_size": None,
|
||||
},
|
||||
QuantBackend.NVFP4.value: {
|
||||
"mode": "nvfp4",
|
||||
"bits": None,
|
||||
"group_size": None,
|
||||
},
|
||||
}
|
||||
|
||||
_SUPPORT_CACHE: dict[str, bool] = {}
|
||||
_SUPPORT_ERROR_CACHE: dict[str, str | None] = {}
|
||||
_BYTES_CACHE: dict[str, float] = {}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantizedWeight:
|
||||
"""Packed quantized weight plus scales/biases for one backend."""
|
||||
|
||||
weight: mx.array
|
||||
scales: mx.array
|
||||
biases: mx.array | None
|
||||
backend: str
|
||||
mode: str
|
||||
bits: int | None
|
||||
group_size: int | None
|
||||
# Original (rows, cols) of the fp weight, used for bytes-per-element.
|
||||
orig_shape: tuple[int, int]
|
||||
|
||||
|
||||
def _normalize_backend(backend: str) -> str:
|
||||
"""Normalize a quantization backend name and validate that it is supported.
|
||||
|
||||
Parameters:
|
||||
backend (str): Backend name to normalize.
|
||||
|
||||
Returns:
|
||||
str: The lowercase backend name without surrounding whitespace.
|
||||
|
||||
Raises:
|
||||
ValueError: If the backend name is unknown.
|
||||
"""
|
||||
name = backend.strip().lower()
|
||||
if name not in _BACKEND_KWARGS:
|
||||
known = ", ".join(BACKENDS)
|
||||
raise ValueError(f"Unknown quant backend {backend!r}. Expected one of: {known}")
|
||||
return name
|
||||
|
||||
|
||||
def _kwargs_for(backend: str) -> dict[str, object]:
|
||||
"""Return the MLX quantization arguments configured for a backend.
|
||||
|
||||
Parameters:
|
||||
backend (str): Backend name to resolve.
|
||||
|
||||
Returns:
|
||||
dict[str, object]: Quantization arguments for the normalized backend.
|
||||
"""
|
||||
return dict(_BACKEND_KWARGS[_normalize_backend(backend)])
|
||||
|
||||
|
||||
def support_error(backend: str) -> str | None:
|
||||
"""
|
||||
Check whether a quantization backend is supported by the current MLX runtime.
|
||||
|
||||
Parameters:
|
||||
backend (str): Quantization backend name.
|
||||
|
||||
Returns:
|
||||
str | None: An error description when the backend is unsupported, or `None` when supported.
|
||||
"""
|
||||
name = _normalize_backend(backend)
|
||||
if name in _SUPPORT_ERROR_CACHE:
|
||||
return _SUPPORT_ERROR_CACHE[name]
|
||||
|
||||
kwargs = _kwargs_for(name)
|
||||
try:
|
||||
w = mx.zeros((_PROBE_DIM, _PROBE_DIM), dtype=mx.float16)
|
||||
quantized = mx.quantize(
|
||||
w,
|
||||
group_size=kwargs["group_size"], # type: ignore[arg-type]
|
||||
bits=kwargs["bits"], # type: ignore[arg-type]
|
||||
mode=str(kwargs["mode"]),
|
||||
)
|
||||
w_q = quantized[0]
|
||||
scales = quantized[1]
|
||||
biases = quantized[2] if len(quantized) == 3 else None
|
||||
x = mx.zeros((1, _PROBE_DIM), dtype=mx.float16)
|
||||
y = mx.quantized_matmul(
|
||||
x,
|
||||
w_q,
|
||||
scales,
|
||||
biases,
|
||||
transpose=True,
|
||||
group_size=kwargs["group_size"], # type: ignore[arg-type]
|
||||
bits=kwargs["bits"], # type: ignore[arg-type]
|
||||
mode=str(kwargs["mode"]),
|
||||
)
|
||||
mx.eval(y)
|
||||
_SUPPORT_ERROR_CACHE[name] = None
|
||||
_SUPPORT_CACHE[name] = True
|
||||
except Exception as exc: # noqa: BLE001 - MLX raises varied types per mode/version.
|
||||
msg = f"{type(exc).__name__}: {exc}"
|
||||
_SUPPORT_ERROR_CACHE[name] = msg
|
||||
_SUPPORT_CACHE[name] = False
|
||||
return _SUPPORT_ERROR_CACHE[name]
|
||||
|
||||
|
||||
def is_supported(backend: str) -> bool:
|
||||
"""Return True if the installed MLX build can quantize/matmul with ``backend``."""
|
||||
name = _normalize_backend(backend)
|
||||
if name not in _SUPPORT_CACHE:
|
||||
support_error(name)
|
||||
return _SUPPORT_CACHE[name]
|
||||
|
||||
|
||||
def quantize_weight(w: mx.array, backend: str) -> QuantizedWeight:
|
||||
"""
|
||||
Quantize a two-dimensional weight matrix using the specified native MLX backend.
|
||||
|
||||
Parameters:
|
||||
w (mx.array): The two-dimensional weight matrix to quantize.
|
||||
backend (str): The quantization backend to use.
|
||||
|
||||
Returns:
|
||||
QuantizedWeight: The quantized weights and associated quantization metadata.
|
||||
|
||||
Raises:
|
||||
ValueError: If the backend is unknown, the weight is not two-dimensional,
|
||||
or its last dimension is not divisible by the backend's group size.
|
||||
RuntimeError: If the backend is unsupported by the installed MLX build.
|
||||
"""
|
||||
name = _normalize_backend(backend)
|
||||
err = support_error(name)
|
||||
if err is not None:
|
||||
mlx_version = getattr(mx, "__version__", "unknown")
|
||||
raise RuntimeError(f"Quant backend {name!r} is not supported by installed mlx "
|
||||
f"({mlx_version}): {err}")
|
||||
|
||||
if w.ndim != 2:
|
||||
raise ValueError(f"quantize_weight expects a 2D weight, got shape {tuple(w.shape)}")
|
||||
|
||||
rows, cols = int(w.shape[0]), int(w.shape[1])
|
||||
kwargs = _kwargs_for(name)
|
||||
group_size = kwargs["group_size"]
|
||||
# When group_size is None, MLX applies the mode default; only check when set.
|
||||
if isinstance(group_size, int) and cols % group_size != 0:
|
||||
raise ValueError(f"Weight last dim {cols} must be divisible by group_size={group_size} "
|
||||
f"for backend {name!r}")
|
||||
|
||||
quantized = mx.quantize(
|
||||
w,
|
||||
group_size=kwargs["group_size"], # type: ignore[arg-type]
|
||||
bits=kwargs["bits"], # type: ignore[arg-type]
|
||||
mode=str(kwargs["mode"]),
|
||||
)
|
||||
w_q = quantized[0]
|
||||
scales = quantized[1]
|
||||
biases = quantized[2] if len(quantized) == 3 else None
|
||||
eval_args = [w_q, scales] if biases is None else [w_q, scales, biases]
|
||||
mx.eval(*eval_args)
|
||||
|
||||
return QuantizedWeight(
|
||||
weight=w_q,
|
||||
scales=scales,
|
||||
biases=biases,
|
||||
backend=name,
|
||||
mode=str(kwargs["mode"]),
|
||||
bits=kwargs["bits"] if isinstance(kwargs["bits"], int) else None,
|
||||
group_size=group_size if isinstance(group_size, int) else None,
|
||||
orig_shape=(rows, cols),
|
||||
)
|
||||
|
||||
|
||||
def quantized_matmul(x: mx.array, qw: QuantizedWeight) -> mx.array:
|
||||
"""Compute ``x @ w.T`` in the quantized domain via ``mx.quantized_matmul``."""
|
||||
return mx.quantized_matmul(
|
||||
x,
|
||||
qw.weight,
|
||||
qw.scales,
|
||||
qw.biases,
|
||||
transpose=True,
|
||||
group_size=qw.group_size,
|
||||
bits=qw.bits,
|
||||
mode=qw.mode,
|
||||
)
|
||||
|
||||
|
||||
def _artifact_nbytes(qw: QuantizedWeight) -> int:
|
||||
"""
|
||||
Calculate the total storage size of a quantized weight artifact in bytes.
|
||||
|
||||
Parameters:
|
||||
qw (QuantizedWeight): Quantized weight artifact whose packed weights, scales, and optional biases are measured.
|
||||
|
||||
Returns:
|
||||
int: Total number of bytes used by the artifact's stored arrays.
|
||||
"""
|
||||
total = int(qw.weight.nbytes) + int(qw.scales.nbytes)
|
||||
if qw.biases is not None:
|
||||
total += int(qw.biases.nbytes)
|
||||
return total
|
||||
|
||||
|
||||
def bytes_per_weight(backend: str) -> float:
|
||||
"""
|
||||
Measure the effective storage cost of a quantized weight.
|
||||
|
||||
Parameters:
|
||||
backend (str): Quantization backend to measure.
|
||||
|
||||
Returns:
|
||||
float: Stored bytes per original weight element, including packed weights,
|
||||
scales, and optional biases.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the backend is unsupported.
|
||||
"""
|
||||
name = _normalize_backend(backend)
|
||||
if name in _BYTES_CACHE:
|
||||
return _BYTES_CACHE[name]
|
||||
|
||||
err = support_error(name)
|
||||
if err is not None:
|
||||
mlx_version = getattr(mx, "__version__", "unknown")
|
||||
raise RuntimeError(f"Cannot measure bytes_per_weight for unsupported backend {name!r} "
|
||||
f"(mlx {mlx_version}): {err}")
|
||||
|
||||
probe = mx.zeros((_PROBE_DIM, _PROBE_DIM), dtype=mx.float16)
|
||||
qw = quantize_weight(probe, name)
|
||||
n_elem = qw.orig_shape[0] * qw.orig_shape[1]
|
||||
value = _artifact_nbytes(qw) / float(n_elem)
|
||||
_BYTES_CACHE[name] = value
|
||||
return value
|
||||
@@ -1,690 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Two-pass spatial refine for the MLX Wan runtime (H3 / LTX-2 pattern).
|
||||
|
||||
Biggest quality lever on Apple Silicon without a new model or training:
|
||||
generate at base resolution, then run a second denoising pass with the
|
||||
*same* DiT at a higher resolution.
|
||||
|
||||
This is the MLX-side port of the CUDA refine template in
|
||||
``fastvideo/pipelines/basic/ltx2/stages/ltx2_refine.py`` and the H3
|
||||
"base + regenerate" pattern documented in
|
||||
``docs/design/mac_qad_two_product_strategy.md``:
|
||||
|
||||
1. :func:`plan_refine_resolutions` — split the request into stage-1
|
||||
(base) and stage-2 (target) pixel sizes, validating VAE / patch
|
||||
alignment the way :class:`LTX2RefineInitStage` does.
|
||||
2. :func:`upsample_latents_spatial` — 2× (or N×) spatial upsample of
|
||||
clean latents. Wan has no learned latent upsampler on Mac, so this
|
||||
is bilinear over the H×W plane (temporal axis untouched) — same
|
||||
role as LTX-2's ``upsample_video`` hand-off, without the learned
|
||||
residual.
|
||||
3. :func:`prepare_refine_latents` — upsample + re-noise the clean
|
||||
stage-1 latents to the stage-2 sigma so the second denoise has
|
||||
something to refine (mirrors :class:`LTX2UpsampleStage` +
|
||||
``apply_ltx2_gaussian_noiser``).
|
||||
4. :func:`run_two_pass_dmd` — orchestrate stage-1 denoise → refine
|
||||
hand-off → stage-2 denoise with the same model / prompt embeds.
|
||||
|
||||
No LoRA swap, no dedicated SR weights, no new training — pure pipeline
|
||||
work reusable by Wan2.1-14B and Wan2.2-5B on Apple Silicon.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from collections.abc import Callable, Sequence
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, add_noise, dmd_step
|
||||
|
||||
if TYPE_CHECKING: # pragma: no cover - typing only
|
||||
import mlx.core as mx
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Default stage-2 noise level when the caller does not supply a schedule.
|
||||
# Matches the first entry of LTX-2's STAGE_2_DISTILLED_SIGMA_VALUES in spirit
|
||||
# (start the refine denoise from a high-noise level) without hard-wiring the
|
||||
# LTX-2 distilled grid onto Wan's flow-match schedule.
|
||||
DEFAULT_REFINE_SIGMA = 0.909375
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RefinePlan:
|
||||
"""Resolved stage-1 / stage-2 geometry for a two-pass refine run."""
|
||||
|
||||
target_height: int
|
||||
target_width: int
|
||||
stage1_height: int
|
||||
stage1_width: int
|
||||
spatial_scale: int
|
||||
vae_spatial_compression: int
|
||||
vae_temporal_compression: int
|
||||
num_frames: int
|
||||
|
||||
@property
|
||||
def stage1_latent_height(self) -> int:
|
||||
"""Return the stage-1 latent height after VAE spatial compression."""
|
||||
return self.stage1_height // self.vae_spatial_compression
|
||||
|
||||
@property
|
||||
def stage1_latent_width(self) -> int:
|
||||
"""Return the stage-one latent width after VAE spatial compression."""
|
||||
return self.stage1_width // self.vae_spatial_compression
|
||||
|
||||
@property
|
||||
def stage2_latent_height(self) -> int:
|
||||
"""Calculate the target-resolution latent height.
|
||||
|
||||
Returns:
|
||||
int: The target height divided by the VAE spatial compression factor.
|
||||
"""
|
||||
return self.target_height // self.vae_spatial_compression
|
||||
|
||||
@property
|
||||
def stage2_latent_width(self) -> int:
|
||||
"""Return the target image width in latent-space units."""
|
||||
return self.target_width // self.vae_spatial_compression
|
||||
|
||||
@property
|
||||
def latent_frames(self) -> int:
|
||||
"""Calculate the number of latent frames after VAE temporal compression.
|
||||
|
||||
Returns:
|
||||
int: The compressed latent frame count.
|
||||
"""
|
||||
return (self.num_frames - 1) // self.vae_temporal_compression + 1
|
||||
|
||||
|
||||
def plan_refine_resolutions(
|
||||
*,
|
||||
height: int,
|
||||
width: int,
|
||||
num_frames: int,
|
||||
spatial_scale: int = 2,
|
||||
vae_spatial_compression: int = 8,
|
||||
vae_temporal_compression: int = 4,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
enabled: bool = True,
|
||||
mode_label: str = "Refine",
|
||||
) -> RefinePlan:
|
||||
"""
|
||||
Validate the requested dimensions and create the stage-1 and target-resolution refinement plan.
|
||||
|
||||
Parameters:
|
||||
height (int): Target image height in pixels.
|
||||
width (int): Target image width in pixels.
|
||||
num_frames (int): Number of frames in the input sequence.
|
||||
spatial_scale (int): Factor used to reduce spatial dimensions for stage 1.
|
||||
vae_spatial_compression (int): Spatial compression factor of the VAE.
|
||||
vae_temporal_compression (int): Temporal compression factor of the VAE.
|
||||
patch_size (tuple[int, int, int]): Temporal and spatial patch dimensions used to validate latent-grid alignment.
|
||||
enabled (bool): Whether to use two-pass refinement.
|
||||
mode_label (str): Name of the calling mode, used to prefix validation
|
||||
errors so ``--fast-spatial`` failures do not read as refine failures.
|
||||
|
||||
Returns:
|
||||
RefinePlan: The validated stage-1 and target-resolution plan.
|
||||
"""
|
||||
if height <= 0 or width <= 0:
|
||||
raise ValueError(f"height/width must be positive, got {height}x{width}")
|
||||
if spatial_scale < 1:
|
||||
raise ValueError(f"spatial_scale must be >= 1, got {spatial_scale}")
|
||||
if num_frames <= 0:
|
||||
raise ValueError(f"num_frames must be positive, got {num_frames}")
|
||||
if vae_spatial_compression < 1 or vae_temporal_compression < 1:
|
||||
raise ValueError("VAE compression factors must be positive")
|
||||
if height % vae_spatial_compression != 0 or width % vae_spatial_compression != 0:
|
||||
raise ValueError(f"height/width must be divisible by vae_spatial_compression={vae_spatial_compression} "
|
||||
f"(got {height}x{width}).")
|
||||
if (num_frames - 1) % vae_temporal_compression != 0:
|
||||
raise ValueError(f"num_frames must be 1 modulo vae_temporal_compression={vae_temporal_compression} "
|
||||
f"(got {num_frames}).")
|
||||
|
||||
if not enabled or spatial_scale == 1:
|
||||
plan = RefinePlan(
|
||||
target_height=height,
|
||||
target_width=width,
|
||||
stage1_height=height,
|
||||
stage1_width=width,
|
||||
spatial_scale=1,
|
||||
vae_spatial_compression=vae_spatial_compression,
|
||||
vae_temporal_compression=vae_temporal_compression,
|
||||
num_frames=num_frames,
|
||||
)
|
||||
_validate_plan(plan, patch_size=patch_size, mode_label=mode_label)
|
||||
return plan
|
||||
|
||||
if height % spatial_scale != 0 or width % spatial_scale != 0:
|
||||
raise ValueError(f"{mode_label} requires height/width divisible by spatial_scale={spatial_scale} "
|
||||
f"(got {height}x{width}).")
|
||||
|
||||
stage1_height = height // spatial_scale
|
||||
stage1_width = width // spatial_scale
|
||||
# Stage-1 must land on a VAE-aligned grid so the first denoise produces
|
||||
# valid latents; the LTX-2 init stage enforces the same constraint.
|
||||
if (stage1_height % vae_spatial_compression != 0 or stage1_width % vae_spatial_compression != 0):
|
||||
raise ValueError(f"{mode_label} requires height/width divisible by "
|
||||
f"{spatial_scale * vae_spatial_compression} "
|
||||
f"(got {height}x{width}, vae_spatial={vae_spatial_compression}).")
|
||||
|
||||
plan = RefinePlan(
|
||||
target_height=height,
|
||||
target_width=width,
|
||||
stage1_height=stage1_height,
|
||||
stage1_width=stage1_width,
|
||||
spatial_scale=spatial_scale,
|
||||
vae_spatial_compression=vae_spatial_compression,
|
||||
vae_temporal_compression=vae_temporal_compression,
|
||||
num_frames=num_frames,
|
||||
)
|
||||
_validate_plan(plan, patch_size=patch_size, mode_label=mode_label)
|
||||
logger.info(
|
||||
"[MLX refine] enabled: stage1=%dx%d stage2=%dx%d scale=%dx",
|
||||
stage1_width,
|
||||
stage1_height,
|
||||
width,
|
||||
height,
|
||||
spatial_scale,
|
||||
)
|
||||
return plan
|
||||
|
||||
|
||||
def _validate_plan(plan: RefinePlan, *, patch_size: tuple[int, int, int], mode_label: str = "Refine") -> None:
|
||||
"""
|
||||
Validate that both refinement stages have latent dimensions aligned to the patch grid.
|
||||
|
||||
Parameters:
|
||||
patch_size (tuple[int, int, int]): Temporal, height, and width patch dimensions.
|
||||
|
||||
Raises:
|
||||
ValueError: If a stage's spatial latent dimensions or the temporal latent
|
||||
dimension is not divisible by the corresponding patch dimension.
|
||||
"""
|
||||
pt, ph, pw = patch_size
|
||||
for label, lh, lw in (
|
||||
("stage1", plan.stage1_latent_height, plan.stage1_latent_width),
|
||||
("stage2", plan.stage2_latent_height, plan.stage2_latent_width),
|
||||
):
|
||||
if lh % ph != 0 or lw % pw != 0:
|
||||
raise ValueError(f"{mode_label} {label} latent grid {lh}x{lw} is not divisible by "
|
||||
f"patch spatial size {ph}x{pw}.")
|
||||
if plan.latent_frames % pt != 0:
|
||||
raise ValueError(f"{mode_label} latent_frames={plan.latent_frames} is not divisible by "
|
||||
f"patch temporal size {pt}.")
|
||||
|
||||
|
||||
def upsample_latents_spatial(
|
||||
latents: Any,
|
||||
*,
|
||||
scale: int = 2,
|
||||
mode: str = "bilinear",
|
||||
) -> Any:
|
||||
"""
|
||||
Upsample the spatial dimensions of 5-D latent arrays while preserving the batch, channel, and temporal dimensions.
|
||||
|
||||
Parameters:
|
||||
latents (Any): Latents with shape ``(B, C, T, H, W)``.
|
||||
scale (int): Integer factor for enlarging the spatial dimensions.
|
||||
mode (str): Interpolation mode, either ``"nearest"`` or ``"bilinear"``.
|
||||
|
||||
Returns:
|
||||
Any: Latents with shape ``(B, C, T, H * scale, W * scale)``.
|
||||
"""
|
||||
if scale < 1:
|
||||
raise ValueError(f"scale must be >= 1, got {scale}")
|
||||
if scale == 1:
|
||||
return latents
|
||||
|
||||
# Accept both mx.array and np.ndarray so unit tests can run without MLX.
|
||||
is_mlx = hasattr(latents, "dtype") and type(latents).__module__.startswith("mlx")
|
||||
if is_mlx:
|
||||
return _upsample_latents_mlx(latents, scale=scale, mode=mode)
|
||||
return _upsample_latents_numpy(np.asarray(latents), scale=scale, mode=mode)
|
||||
|
||||
|
||||
def _upsample_latents_numpy(
|
||||
latents: np.ndarray,
|
||||
*,
|
||||
scale: int,
|
||||
mode: str,
|
||||
) -> np.ndarray:
|
||||
"""Upsample 5-D latent arrays spatially using nearest-neighbor or bilinear interpolation.
|
||||
|
||||
Parameters:
|
||||
latents (np.ndarray): Latents with shape ``(B, C, T, H, W)``.
|
||||
scale (int): Spatial upsampling factor.
|
||||
mode (str): Interpolation mode, either ``"nearest"`` or ``"bilinear"``.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Spatially upsampled latents with preserved batch, channel, and temporal dimensions.
|
||||
|
||||
Raises:
|
||||
ValueError: If the latents are not five-dimensional or the interpolation mode is unsupported.
|
||||
"""
|
||||
if latents.ndim != 5:
|
||||
raise ValueError(f"Expected 5-D latents (B,C,T,H,W), got shape {latents.shape}")
|
||||
b, c, t, h, w = latents.shape
|
||||
if mode == "nearest":
|
||||
# (B,C,T,H,1,W,1) -> broadcast to (B,C,T,H,scale,W,scale) -> merge.
|
||||
out = np.repeat(np.repeat(latents, scale, axis=3), scale, axis=4)
|
||||
return out
|
||||
|
||||
if mode != "bilinear":
|
||||
raise ValueError(f"Unsupported upsample mode: {mode}")
|
||||
|
||||
# Bilinear over the spatial plane. Flatten (B,C,T) into a batch of 2-D
|
||||
# maps so a single vectorized gather covers every frame/channel.
|
||||
src = latents.reshape(b * c * t, h, w).astype(np.float32, copy=False)
|
||||
out_h, out_w = h * scale, w * scale
|
||||
# Map output pixel centers onto the input grid (align_corners=False).
|
||||
ys = (np.arange(out_h, dtype=np.float32) + 0.5) * (h / out_h) - 0.5
|
||||
xs = (np.arange(out_w, dtype=np.float32) + 0.5) * (w / out_w) - 0.5
|
||||
ys = np.clip(ys, 0.0, h - 1.0)
|
||||
xs = np.clip(xs, 0.0, w - 1.0)
|
||||
y0 = np.floor(ys).astype(np.int64)
|
||||
x0 = np.floor(xs).astype(np.int64)
|
||||
y1 = np.minimum(y0 + 1, h - 1)
|
||||
x1 = np.minimum(x0 + 1, w - 1)
|
||||
wy = (ys - y0.astype(np.float32))[:, None]
|
||||
wx = (xs - x0.astype(np.float32))[None, :]
|
||||
# Gather the four corners: shape (N, out_h, out_w).
|
||||
Ia = src[:, y0[:, None], x0[None, :]]
|
||||
Ib = src[:, y0[:, None], x1[None, :]]
|
||||
Ic = src[:, y1[:, None], x0[None, :]]
|
||||
Id = src[:, y1[:, None], x1[None, :]]
|
||||
wa = (1.0 - wy) * (1.0 - wx)
|
||||
wb = (1.0 - wy) * wx
|
||||
wc = wy * (1.0 - wx)
|
||||
wd = wy * wx
|
||||
out = wa * Ia + wb * Ib + wc * Ic + wd * Id
|
||||
return out.reshape(b, c, t, out_h, out_w).astype(latents.dtype, copy=False)
|
||||
|
||||
|
||||
def _upsample_latents_mlx(latents: mx.array, *, scale: int, mode: str) -> mx.array:
|
||||
"""
|
||||
Upsample MLX latent tensors along their spatial dimensions.
|
||||
|
||||
Parameters:
|
||||
latents (mx.array): A latent tensor with shape `(B, C, T, H, W)`.
|
||||
scale (int): The integer spatial upsampling factor.
|
||||
mode (str): The interpolation mode, such as `"nearest"` or `"bilinear"`.
|
||||
|
||||
Returns:
|
||||
mx.array: The spatially upsampled latent tensor with its original data type.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
# Route through NumPy for the interpolation math. Latent tensors at Mac
|
||||
# resolutions are small (e.g. 1×16×21×30×52 ≈ 1 MB) so the host hop is
|
||||
# cheaper than carrying a bespoke Metal bilinear kernel, and it keeps
|
||||
# the CPU-only unit tests and the MLX path on one implementation.
|
||||
np_latents = np.array(latents.astype(mx.float32))
|
||||
up = _upsample_latents_numpy(np_latents, scale=scale, mode=mode)
|
||||
return mx.array(up).astype(latents.dtype)
|
||||
|
||||
|
||||
def prepare_refine_latents(
|
||||
clean_latents: Any,
|
||||
*,
|
||||
scale: int = 2,
|
||||
sigma: float = DEFAULT_REFINE_SIGMA,
|
||||
noise: Any | None = None,
|
||||
add_noise_flag: bool = True,
|
||||
upsample_mode: str = "bilinear",
|
||||
seed: int | None = None,
|
||||
) -> Any:
|
||||
"""
|
||||
Upsample clean latents spatially and optionally mix them with Gaussian noise.
|
||||
|
||||
Parameters:
|
||||
clean_latents: The stage-1 latent tensor.
|
||||
sigma: Noise mixing factor between 0 and 1.
|
||||
noise: Optional noise tensor to mix with the upsampled latents.
|
||||
add_noise_flag: Whether to apply noise mixing.
|
||||
upsample_mode: Spatial interpolation mode.
|
||||
seed: Optional seed for generated noise.
|
||||
|
||||
Returns:
|
||||
The upsampled latents, optionally mixed with noise.
|
||||
|
||||
Raises:
|
||||
ValueError: If sigma is outside the range from 0 to 1.
|
||||
"""
|
||||
if sigma < 0.0 or sigma > 1.0:
|
||||
raise ValueError(f"sigma must be in [0, 1], got {sigma}")
|
||||
|
||||
upsampled = upsample_latents_spatial(clean_latents, scale=scale, mode=upsample_mode)
|
||||
if not add_noise_flag or sigma == 0.0:
|
||||
return upsampled
|
||||
|
||||
is_mlx = hasattr(upsampled, "dtype") and type(upsampled).__module__.startswith("mlx")
|
||||
if noise is None:
|
||||
noise = _draw_noise_like(upsampled, seed=seed, is_mlx=is_mlx)
|
||||
return add_noise(upsampled, noise, float(sigma))
|
||||
|
||||
|
||||
def refine_sigma_from_schedule(
|
||||
schedule: MLXDMDSchedule,
|
||||
timesteps: Sequence[float | int],
|
||||
) -> float:
|
||||
"""Derive the refinement noise level from the first refinement timestep.
|
||||
|
||||
Parameters:
|
||||
schedule (MLXDMDSchedule): Schedule used to map timesteps to noise levels.
|
||||
timesteps (Sequence[float | int]): Refinement timesteps, whose first value determines the sigma.
|
||||
|
||||
Returns:
|
||||
float: Sigma corresponding to the first refinement timestep.
|
||||
|
||||
Raises:
|
||||
ValueError: If `timesteps` is empty.
|
||||
"""
|
||||
if not timesteps:
|
||||
raise ValueError("timesteps must be non-empty to derive a refine sigma")
|
||||
return float(schedule.sigma_for(float(timesteps[0])))
|
||||
|
||||
|
||||
def default_refine_timesteps(
|
||||
schedule: MLXDMDSchedule,
|
||||
timesteps: Sequence[float | int],
|
||||
) -> list[float]:
|
||||
"""Derive stage-2 timesteps from the stage-1 DMD grid.
|
||||
|
||||
The stage-2 pass must start *below* full noise, otherwise the hand-off
|
||||
``(1 - sigma) * upsampled + sigma * noise`` weights stage 1 at zero and
|
||||
the refine pass silently becomes a plain full-resolution generation at
|
||||
twice the cost. FastWan's stage-1 grid opens at ``t=1000`` (``sigma``
|
||||
exactly 1.0), so reusing it verbatim — which is what happens when
|
||||
``--refine-dmd-denoising-steps`` is left unset — discards stage 1.
|
||||
|
||||
Dropping the leading full-noise entries keeps the pass on timesteps the
|
||||
distilled student was actually trained on (no off-grid ``t`` the DiT has
|
||||
never seen) while letting the stage-1 structure through.
|
||||
|
||||
Parameters:
|
||||
schedule (MLXDMDSchedule): Schedule used to map timesteps to noise levels.
|
||||
timesteps (Sequence[float | int]): The stage-1 DMD timestep grid.
|
||||
|
||||
Returns:
|
||||
list[float]: The stage-1 grid with leading full-noise timesteps removed.
|
||||
|
||||
Raises:
|
||||
ValueError: If every timestep in the grid is at full noise, leaving no
|
||||
usable refine step.
|
||||
"""
|
||||
steps = [float(step) for step in timesteps]
|
||||
first = 0
|
||||
while first < len(steps) and schedule.sigma_for(steps[first]) >= 1.0:
|
||||
first += 1
|
||||
if first == len(steps):
|
||||
raise ValueError(f"No usable refine timesteps in {steps}: every entry is at sigma >= 1 "
|
||||
"(full noise), which would discard the stage-1 result. Pass "
|
||||
"explicit stage-2 timesteps below the full-noise step.")
|
||||
return steps[first:]
|
||||
|
||||
|
||||
def run_dmd_loop(
|
||||
*,
|
||||
dit: Any,
|
||||
latents: Any,
|
||||
encoder_hidden_states: Any,
|
||||
freqs_cis: tuple[Any, Any],
|
||||
timesteps: Sequence[float | int],
|
||||
schedule: MLXDMDSchedule,
|
||||
mx_dtype: Any,
|
||||
seed: int | None = None,
|
||||
step_callback: Callable[[int, int], None] | None = None,
|
||||
label: str = "denoise",
|
||||
) -> Any:
|
||||
"""
|
||||
Denoise latents over the supplied timesteps using the DMD schedule.
|
||||
|
||||
Parameters:
|
||||
timesteps (Sequence[float | int]): Denoising timesteps in execution order.
|
||||
seed (int | None): Seed for reproducible intermediate noise generation.
|
||||
step_callback (Callable[[int, int], None] | None): Callback receiving the
|
||||
completed step number and total step count.
|
||||
label (str): Label used for progress output when no callback is provided.
|
||||
|
||||
Returns:
|
||||
Any: The denoised latents.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
renoise_rng = np.random.default_rng(seed) if seed is not None else None
|
||||
latents_out = latents
|
||||
n_steps = len(timesteps)
|
||||
for step_index, timestep in enumerate(timesteps):
|
||||
noise_input = latents_out
|
||||
ts_val = float(timestep)
|
||||
timestep_mx = mx.array([ts_val]).astype(mx.float32)
|
||||
noise_pred = dit(
|
||||
latents_out.astype(mx_dtype),
|
||||
encoder_hidden_states,
|
||||
timestep_mx,
|
||||
freqs_cis,
|
||||
)
|
||||
noise_input_f32 = noise_input.astype(mx.float32)
|
||||
pred_noise_f32 = noise_pred.astype(mx.float32)
|
||||
if step_index < n_steps - 1:
|
||||
next_ts: float | None = float(timesteps[step_index + 1])
|
||||
if renoise_rng is not None:
|
||||
renoise = mx.array(renoise_rng.standard_normal(tuple(noise_input_f32.shape)).astype(np.float32))
|
||||
else:
|
||||
renoise = mx.random.normal(noise_input_f32.shape).astype(mx.float32)
|
||||
else:
|
||||
next_ts, renoise = None, None
|
||||
latents_out = dmd_step(
|
||||
latents=noise_input_f32,
|
||||
noise_input_latent=noise_input_f32,
|
||||
pred_noise=pred_noise_f32,
|
||||
schedule=schedule,
|
||||
timestep=ts_val,
|
||||
next_timestep=next_ts,
|
||||
noise=renoise,
|
||||
).astype(mx_dtype)
|
||||
mx.eval(latents_out)
|
||||
if step_callback is not None:
|
||||
step_callback(step_index + 1, n_steps)
|
||||
else:
|
||||
print(f"{label} step {step_index + 1}/{n_steps} complete")
|
||||
return latents_out
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TwoPassResult:
|
||||
"""Outputs of :func:`run_two_pass_dmd`."""
|
||||
|
||||
latents: Any
|
||||
stage1_latents: Any
|
||||
plan: RefinePlan
|
||||
refine_sigma: float
|
||||
|
||||
|
||||
def run_two_pass_dmd(
|
||||
*,
|
||||
dit: Any,
|
||||
encoder_hidden_states: Any,
|
||||
noise_latents_stage1: Any,
|
||||
freqs_cis_stage1: tuple[Any, Any],
|
||||
freqs_cis_stage2: tuple[Any, Any] | None,
|
||||
plan: RefinePlan,
|
||||
schedule: MLXDMDSchedule,
|
||||
timesteps: Sequence[float | int],
|
||||
refine_timesteps: Sequence[float | int] | None = None,
|
||||
mx_dtype: Any,
|
||||
seed: int = 0,
|
||||
add_noise_flag: bool = True,
|
||||
upsample_mode: str = "bilinear",
|
||||
refine_sigma: float | None = None,
|
||||
step_callback: Callable[[str, int, int], None] | None = None,
|
||||
) -> TwoPassResult:
|
||||
"""
|
||||
Run base denoising and, when enabled, spatial refinement denoising.
|
||||
|
||||
Parameters:
|
||||
dit: DiT callable used for both denoising passes.
|
||||
encoder_hidden_states: Prompt embeddings shared across both passes.
|
||||
noise_latents_stage1: Initial stage-1 noise latents.
|
||||
freqs_cis_stage1: RoPE tables for the stage-1 resolution.
|
||||
freqs_cis_stage2: RoPE tables for the stage-2 resolution, required when refinement is enabled.
|
||||
plan: Refinement geometry and configuration.
|
||||
schedule: Flow-matching schedule used by both passes.
|
||||
timesteps: Stage-1 denoising timesteps.
|
||||
refine_timesteps: Stage-2 denoising timesteps. Uses `timesteps` when omitted.
|
||||
mx_dtype: MLX dtype used for DiT inputs and outputs.
|
||||
seed: Base seed for reproducible noise generation.
|
||||
add_noise_flag: Whether to add noise to the upsampled stage-1 latents.
|
||||
upsample_mode: Spatial upsampling mode, either `"bilinear"` or `"nearest"`.
|
||||
refine_sigma: Stage-2 starting noise level. Derived from the first refinement timestep when omitted.
|
||||
step_callback: Optional callback receiving the phase name, step index, and total step count.
|
||||
|
||||
Returns:
|
||||
TwoPassResult containing the final latents, stage-1 latents, refinement plan, and applied refinement sigma.
|
||||
|
||||
Raises:
|
||||
ValueError: If refinement is enabled without stage-2 RoPE tables, without refinement timesteps, or if upsampled latents do not match the planned stage-2 dimensions.
|
||||
"""
|
||||
stage1_cb = None
|
||||
stage2_cb = None
|
||||
if step_callback is not None:
|
||||
stage1_cb = lambda i, n: step_callback("stage1", i, n) # noqa: E731
|
||||
stage2_cb = lambda i, n: step_callback("stage2", i, n) # noqa: E731
|
||||
|
||||
stage1_latents = run_dmd_loop(
|
||||
dit=dit,
|
||||
latents=noise_latents_stage1,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
freqs_cis=freqs_cis_stage1,
|
||||
timesteps=timesteps,
|
||||
schedule=schedule,
|
||||
mx_dtype=mx_dtype,
|
||||
seed=seed,
|
||||
step_callback=stage1_cb,
|
||||
label="stage1 denoise",
|
||||
)
|
||||
|
||||
if plan.spatial_scale == 1:
|
||||
return TwoPassResult(
|
||||
latents=stage1_latents,
|
||||
stage1_latents=stage1_latents,
|
||||
plan=plan,
|
||||
refine_sigma=0.0,
|
||||
)
|
||||
|
||||
if freqs_cis_stage2 is None:
|
||||
raise ValueError("freqs_cis_stage2 is required when refine spatial_scale > 1")
|
||||
|
||||
if refine_timesteps is not None:
|
||||
stage2_timesteps = [float(step) for step in refine_timesteps]
|
||||
if not stage2_timesteps:
|
||||
raise ValueError("refine_timesteps must be non-empty when refine is enabled")
|
||||
else:
|
||||
# Not `list(timesteps)`: the stage-1 grid opens at full noise, which
|
||||
# would weight the stage-1 result at zero. See default_refine_timesteps.
|
||||
stage2_timesteps = default_refine_timesteps(schedule, timesteps)
|
||||
grid_sigma = refine_sigma_from_schedule(schedule, stage2_timesteps)
|
||||
sigma = float(refine_sigma) if refine_sigma is not None else grid_sigma
|
||||
if refine_sigma is not None and abs(sigma - grid_sigma) > 1e-6:
|
||||
# The loop tells the DiT `stage2_timesteps[0]`, which implies grid_sigma.
|
||||
# Overriding the hand-off noise level breaks that correspondence, so the
|
||||
# model is denoising from a level it was not told about. Useful for
|
||||
# exploring schedules that bottom out too high, but say so out loud.
|
||||
logger.warning(
|
||||
"[MLX refine] refine_sigma=%.4f overrides the schedule's %.4f for timestep %g; "
|
||||
"the DiT is told a timestep that no longer matches the noise it receives.",
|
||||
sigma,
|
||||
grid_sigma,
|
||||
stage2_timesteps[0],
|
||||
)
|
||||
|
||||
# A hand-off at sigma >= 1 is `0 * upsampled + 1 * noise`: stage 1 is
|
||||
# thrown away and refine degrades to a plain full-res run at 2x the cost.
|
||||
# Fail loudly rather than silently burning the first pass.
|
||||
if add_noise_flag and sigma >= 1.0:
|
||||
raise ValueError(f"Refine hand-off sigma={sigma:.4f} (from stage-2 timestep "
|
||||
f"{stage2_timesteps[0]:g}) discards the stage-1 result entirely: "
|
||||
"the upsampled latents are weighted (1 - sigma) = 0. Start the "
|
||||
"stage-2 grid below the full-noise timestep, or pass "
|
||||
"add_noise_flag=False to hand off the clean upsample.")
|
||||
|
||||
stage2_input = prepare_refine_latents(
|
||||
stage1_latents,
|
||||
scale=plan.spatial_scale,
|
||||
sigma=sigma,
|
||||
add_noise_flag=add_noise_flag,
|
||||
upsample_mode=upsample_mode,
|
||||
seed=seed + 1,
|
||||
)
|
||||
|
||||
# Shape guard: upsampled latents must match the stage-2 RoPE grid.
|
||||
expected_h = plan.stage2_latent_height
|
||||
expected_w = plan.stage2_latent_width
|
||||
got_h, got_w = int(stage2_input.shape[-2]), int(stage2_input.shape[-1])
|
||||
if got_h != expected_h or got_w != expected_w:
|
||||
raise ValueError(f"Refine upsample produced {got_h}x{got_w} latents, expected "
|
||||
f"{expected_h}x{expected_w} for target "
|
||||
f"{plan.target_height}x{plan.target_width}.")
|
||||
|
||||
logger.info(
|
||||
"[MLX refine] stage2 start: latent=%dx%d sigma=%.4f steps=%d",
|
||||
expected_w,
|
||||
expected_h,
|
||||
sigma,
|
||||
len(stage2_timesteps),
|
||||
)
|
||||
|
||||
stage2_latents = run_dmd_loop(
|
||||
dit=dit,
|
||||
latents=stage2_input,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
freqs_cis=freqs_cis_stage2,
|
||||
timesteps=stage2_timesteps,
|
||||
schedule=schedule,
|
||||
mx_dtype=mx_dtype,
|
||||
seed=seed + 2,
|
||||
step_callback=stage2_cb,
|
||||
label="stage2 refine",
|
||||
)
|
||||
return TwoPassResult(
|
||||
latents=stage2_latents,
|
||||
stage1_latents=stage1_latents,
|
||||
plan=plan,
|
||||
refine_sigma=sigma,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DEFAULT_REFINE_SIGMA",
|
||||
"RefinePlan",
|
||||
"TwoPassResult",
|
||||
"default_refine_timesteps",
|
||||
"plan_refine_resolutions",
|
||||
"prepare_refine_latents",
|
||||
"refine_sigma_from_schedule",
|
||||
"run_dmd_loop",
|
||||
"run_two_pass_dmd",
|
||||
"upsample_latents_spatial",
|
||||
]
|
||||
|
||||
|
||||
def _draw_noise_like(like: Any, *, seed: int | None, is_mlx: bool) -> Any:
|
||||
"""Generate Gaussian noise with the shape and array type of the input."""
|
||||
shape = tuple(int(s) for s in like.shape)
|
||||
if seed is not None:
|
||||
rng = np.random.default_rng(seed)
|
||||
noise_np = rng.standard_normal(shape).astype(np.float32)
|
||||
if is_mlx:
|
||||
import mlx.core as mx
|
||||
|
||||
return mx.array(noise_np).astype(mx.float32)
|
||||
return noise_np.astype(np.asarray(like).dtype, copy=False)
|
||||
if is_mlx:
|
||||
import mlx.core as mx
|
||||
|
||||
return mx.random.normal(shape).astype(mx.float32)
|
||||
return np.random.standard_normal(shape).astype(np.asarray(like).dtype, copy=False)
|
||||
@@ -1,137 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Small MLX RIFE wrapper for frame interpolation experiments.
|
||||
|
||||
The backend is the Apple-Silicon-native ``rife-mlx`` package, using the
|
||||
``mlx-community/RIFE-4.25`` weights. Frames are HWC RGB ``uint8`` arrays.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from functools import lru_cache
|
||||
|
||||
import numpy as np
|
||||
from huggingface_hub.utils import LocalEntryNotFoundError
|
||||
|
||||
|
||||
class RIFEBackendError(RuntimeError):
|
||||
"""Raised when the MLX RIFE backend cannot be loaded or run."""
|
||||
|
||||
|
||||
class RIFEWeightsUnavailableError(RIFEBackendError):
|
||||
"""Raised when uncached RIFE weights cannot be downloaded."""
|
||||
|
||||
|
||||
def aligned_keyframe_count(target_frames: int, factor: int, temporal_compression: int = 4) -> int:
|
||||
"""Return the smallest VAE-aligned keyframe count that RIFE can expand to the target."""
|
||||
if target_frames < 1:
|
||||
raise ValueError(f"target_frames must be >= 1, got {target_frames}")
|
||||
if factor < 1:
|
||||
raise ValueError(f"factor must be >= 1, got {factor}")
|
||||
if temporal_compression < 1:
|
||||
raise ValueError(f"temporal_compression must be >= 1, got {temporal_compression}")
|
||||
required_intervals = (target_frames - 1 + factor - 1) // factor
|
||||
aligned_intervals = ((required_intervals + temporal_compression - 1) // temporal_compression * temporal_compression)
|
||||
return aligned_intervals + 1
|
||||
|
||||
|
||||
def _require_hwc_rgb(frame: np.ndarray, index: int) -> np.ndarray:
|
||||
array = np.asarray(frame)
|
||||
if array.ndim != 3 or array.shape[2] != 3:
|
||||
raise ValueError(f"frame {index} must have shape HxWx3, got {array.shape}")
|
||||
if array.dtype != np.uint8:
|
||||
array = np.clip(array, 0, 255).astype(np.uint8)
|
||||
return np.ascontiguousarray(array)
|
||||
|
||||
|
||||
@lru_cache(maxsize=2)
|
||||
def load_model(version: str = "4.25", weights_dir: str | None = None):
|
||||
"""Load the MLX-native RIFE model.
|
||||
|
||||
``weights_dir`` is passed through to ``build_model`` in the vendored ``rife_mlx``.
|
||||
When it is ``None``, the package downloads/uses the Hugging Face
|
||||
``mlx-community/RIFE-4.25`` snapshot.
|
||||
"""
|
||||
try:
|
||||
from fastvideo.third_party.rife_mlx.utils.weights import build_model
|
||||
except ImportError:
|
||||
# Fall back to a separately installed upstream package, for anyone who
|
||||
# already has one in the environment.
|
||||
try:
|
||||
from rife_mlx.utils.weights import build_model
|
||||
except ImportError as exc:
|
||||
raise RIFEBackendError("MLX RIFE backend is unavailable. It ships vendored under "
|
||||
"fastvideo/third_party/rife_mlx, so this usually means MLX "
|
||||
"itself is missing: install with `uv pip install -e '.[mlx]'`.") from exc
|
||||
|
||||
try:
|
||||
return build_model(version, weights_dir=weights_dir)
|
||||
except LocalEntryNotFoundError as exc:
|
||||
raise RIFEWeightsUnavailableError(f"MLX RIFE {version} weights are unavailable: {exc}") from exc
|
||||
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
|
||||
raise RIFEBackendError(f"Failed to load MLX RIFE {version}: {exc}") from exc
|
||||
|
||||
|
||||
def interpolate_pair(
|
||||
frame_a: np.ndarray,
|
||||
frame_b: np.ndarray,
|
||||
timestep: float = 0.5,
|
||||
*,
|
||||
model=None,
|
||||
scale: float = 1.0,
|
||||
) -> np.ndarray:
|
||||
"""Interpolate one RGB frame between two input RGB frames."""
|
||||
if not 0.0 < timestep < 1.0:
|
||||
raise ValueError(f"timestep must be inside (0, 1), got {timestep}")
|
||||
img0 = _require_hwc_rgb(frame_a, 0)
|
||||
img1 = _require_hwc_rgb(frame_b, 1)
|
||||
if img0.shape != img1.shape:
|
||||
raise ValueError(f"frame shapes must match, got {img0.shape} and {img1.shape}")
|
||||
|
||||
if model is None:
|
||||
model = load_model()
|
||||
try:
|
||||
try:
|
||||
from fastvideo.third_party.rife_mlx.pipeline_mlx import interpolate_pair as _interpolate_pair
|
||||
except ImportError:
|
||||
from rife_mlx.pipeline_mlx import interpolate_pair as _interpolate_pair
|
||||
|
||||
return _interpolate_pair(model, img0, img1, timestep=timestep, scale=scale)
|
||||
except Exception as exc: # noqa: BLE001 - preserve exact backend failure.
|
||||
raise RIFEBackendError(f"MLX RIFE interpolation failed at timestep={timestep}: {exc}") from exc
|
||||
|
||||
|
||||
def interpolate(
|
||||
frames: list[np.ndarray] | Iterable[np.ndarray],
|
||||
factor: int = 2,
|
||||
*,
|
||||
model=None,
|
||||
scale: float = 1.0,
|
||||
) -> list[np.ndarray]:
|
||||
"""Return an Nx interpolated frame list.
|
||||
|
||||
For ``len(frames)=41`` and ``factor=2``, the output length is 81:
|
||||
``(41 - 1) * 2 + 1``. Original keyframes are preserved in order and RIFE
|
||||
fills ``factor - 1`` intermediate timesteps between each adjacent pair.
|
||||
"""
|
||||
frame_list = [_require_hwc_rgb(frame, idx) for idx, frame in enumerate(frames)]
|
||||
if factor < 1:
|
||||
raise ValueError(f"factor must be >= 1, got {factor}")
|
||||
if len(frame_list) < 2 or factor == 1:
|
||||
return [frame.copy() for frame in frame_list]
|
||||
|
||||
first_shape = frame_list[0].shape
|
||||
for idx, frame in enumerate(frame_list[1:], start=1):
|
||||
if frame.shape != first_shape:
|
||||
raise ValueError(f"all frames must have the same shape; frame 0={first_shape}, frame {idx}={frame.shape}")
|
||||
|
||||
if model is None:
|
||||
model = load_model()
|
||||
|
||||
out: list[np.ndarray] = []
|
||||
for left, right in zip(frame_list[:-1], frame_list[1:], strict=True):
|
||||
out.append(left)
|
||||
for step in range(1, factor):
|
||||
out.append(interpolate_pair(left, right, step / factor, model=model, scale=scale))
|
||||
out.append(frame_list[-1])
|
||||
return out
|
||||
@@ -1,139 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""On-device (MLX) DMD sampling for the FastWan runtime.
|
||||
|
||||
The hybrid proof-of-concept ran the FastWan DiT in MLX but bounced every
|
||||
denoising step back through torch/NumPy to run the DMD scheduler math
|
||||
(``MLX -> np.array -> torch (CPU) -> np.array -> MLX``). That host round-trip
|
||||
forces a full device sync per step and defeats MLX's lazy graph execution.
|
||||
|
||||
This module mirrors the exact DMD arithmetic from
|
||||
``fastvideo/models/utils.py::pred_noise_to_pred_video`` and
|
||||
``FlowMatchEulerDiscreteScheduler.add_noise`` while keeping every large tensor
|
||||
on the MLX device. The schedule lookup (``argmin`` over the ~1000-entry
|
||||
training schedule) is done once on the host in NumPy: it is tiny, it is the
|
||||
same value torch would compute, and it sidesteps the reduction-index quirk that
|
||||
affects ``argmin`` on the Metal/MPS backends (see the CPU fallbacks in
|
||||
``fastvideo/models/utils.py`` and ``scheduling_flow_match_euler_discrete.py``).
|
||||
|
||||
Because the DMD loop applies a single scalar timestep per step, ``sigma`` is a
|
||||
scalar and the update is a plain elementwise affine combination — no
|
||||
permute/flatten reshaping is required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
if TYPE_CHECKING: # pragma: no cover - typing only
|
||||
import mlx.core as mx
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MLXDMDSchedule:
|
||||
"""Host-side copy of a flow-match scheduler's ``(sigmas, timesteps)``.
|
||||
|
||||
Holds the full training schedule so a DMD timestep (e.g. one of
|
||||
``1000, 757, 522``) can be mapped to its flow-match ``sigma`` with the same
|
||||
nearest-timestep lookup the torch path uses.
|
||||
"""
|
||||
|
||||
sigmas: np.ndarray
|
||||
timesteps: np.ndarray
|
||||
|
||||
@classmethod
|
||||
def from_torch_scheduler(cls, scheduler: Any) -> MLXDMDSchedule:
|
||||
"""Snapshot ``scheduler.sigmas`` / ``scheduler.timesteps`` to NumPy.
|
||||
|
||||
Matches ``pred_noise_to_pred_video`` / ``add_noise``, which index the
|
||||
scheduler's *full* training schedule (not the per-inference subset).
|
||||
"""
|
||||
sigmas = scheduler.sigmas.detach().to("cpu").double().numpy()
|
||||
timesteps = scheduler.timesteps.detach().to("cpu").double().numpy()
|
||||
return cls(sigmas=np.asarray(sigmas), timesteps=np.asarray(timesteps))
|
||||
|
||||
def sigma_for(self, timestep: float) -> float:
|
||||
"""
|
||||
Find the sigma associated with the scheduled timestep nearest to the given timestep.
|
||||
|
||||
Parameters:
|
||||
timestep (float): Timestep for which to find the nearest scheduled sigma.
|
||||
|
||||
Returns:
|
||||
float: Sigma associated with the nearest scheduled timestep.
|
||||
"""
|
||||
idx = int(np.argmin(np.abs(self.timesteps - float(timestep))))
|
||||
return float(self.sigmas[idx])
|
||||
|
||||
|
||||
def pred_noise_to_pred_video(
|
||||
pred_noise: mx.array,
|
||||
noise_input_latent: mx.array,
|
||||
sigma: float,
|
||||
) -> mx.array:
|
||||
"""
|
||||
Compute the clean latent prediction from a flow-matching noise prediction.
|
||||
|
||||
Parameters:
|
||||
pred_noise (mx.array): Predicted noise.
|
||||
noise_input_latent (mx.array): Noised latent input.
|
||||
sigma (float): Noise level used for the prediction.
|
||||
|
||||
Returns:
|
||||
mx.array: Predicted clean latent.
|
||||
"""
|
||||
return noise_input_latent - sigma * pred_noise
|
||||
|
||||
|
||||
def add_noise(
|
||||
clean_latent: mx.array,
|
||||
noise: mx.array,
|
||||
sigma: float,
|
||||
) -> mx.array:
|
||||
"""Flow-match forward noising, mirroring the scheduler's ``add_noise``.
|
||||
|
||||
``sample = (1 - sigma) * clean_latent + sigma * noise``.
|
||||
"""
|
||||
return (1.0 - sigma) * clean_latent + sigma * noise
|
||||
|
||||
|
||||
def dmd_step(
|
||||
*,
|
||||
latents: mx.array,
|
||||
noise_input_latent: mx.array,
|
||||
pred_noise: mx.array,
|
||||
schedule: MLXDMDSchedule,
|
||||
timestep: float,
|
||||
next_timestep: float | None,
|
||||
noise: mx.array | None = None,
|
||||
) -> mx.array:
|
||||
"""
|
||||
Compute one DMD sampling update, optionally re-noising the clean latent prediction.
|
||||
|
||||
Args:
|
||||
latents: Retained for call-site compatibility and not used in the update.
|
||||
noise_input_latent: Noisy latent used to compute the clean prediction.
|
||||
pred_noise: Predicted noise or velocity.
|
||||
schedule: Flow-matching schedule used to map timesteps to sigmas.
|
||||
timestep: Current sampling timestep.
|
||||
next_timestep: Timestep for the next update, or `None` for the final step.
|
||||
noise: Fresh noise used for re-noising intermediate steps.
|
||||
|
||||
Returns:
|
||||
The re-noised latent for the next step or the clean latent prediction on
|
||||
the final step.
|
||||
|
||||
Raises:
|
||||
ValueError: If `next_timestep` is provided without `noise`.
|
||||
"""
|
||||
del latents # symmetry with the torch loop; not needed for the math.
|
||||
sigma = schedule.sigma_for(timestep)
|
||||
pred_video = pred_noise_to_pred_video(pred_noise, noise_input_latent, sigma)
|
||||
if next_timestep is None:
|
||||
return pred_video
|
||||
if noise is None:
|
||||
raise ValueError("dmd_step requires `noise` when `next_timestep` is set (re-noise step).")
|
||||
sigma_next = schedule.sigma_for(next_timestep)
|
||||
return add_noise(pred_video, noise, sigma_next)
|
||||
@@ -1,191 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Optional TAEHV decode helpers for Apple Silicon FastWan experiments.
|
||||
|
||||
The TAEHV module itself is vendored at ``fastvideo/third_party/taehv`` (MIT,
|
||||
madebyollin/taehv), so no source code is downloaded or executed at runtime.
|
||||
Only the ``taew2_1.pth`` checkpoint is fetched on demand, and its sha256 is
|
||||
verified before use.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import importlib.util
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
TAEW2_1_CHECKPOINT_URL = "https://raw.githubusercontent.com/madebyollin/taehv/main/taew2_1.pth"
|
||||
# sha256 of the upstream taew2_1.pth this module was validated against
|
||||
# (fetched 2026-07-02). If upstream publishes a new checkpoint, revalidate the
|
||||
# decode path and update this pin.
|
||||
TAEW2_1_CHECKPOINT_SHA256 = "d26151e76cdc2c9424bef988de874b33d9a53f30ef3060cd556c429c469c797e"
|
||||
# Wan2.2 5B (z_dim=48) — see madebyollin/taehv taew2_2.pth; prefer
|
||||
# ``fastvideo.mlx_runtime.wan_vae.ensure_taehv_checkpoint(z_dim=48)`` for new code.
|
||||
TAEW2_2_CHECKPOINT_URL = ("https://raw.githubusercontent.com/madebyollin/taehv/"
|
||||
"563f40bdc820ed86bcad72ea515ee48f06bd22ec/taew2_2.pth")
|
||||
|
||||
|
||||
def _default_cache_dir() -> Path:
|
||||
"""Return the default directory used to cache TAEHV checkpoints.
|
||||
|
||||
Returns:
|
||||
Path: The TAEHV checkpoint cache directory under the user's home directory.
|
||||
"""
|
||||
return Path.home() / ".cache" / "fastvideo" / "taehv"
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
"""
|
||||
Compute the SHA-256 digest of a file.
|
||||
|
||||
Parameters:
|
||||
path (Path): The file whose contents are hashed.
|
||||
|
||||
Returns:
|
||||
str: The file's SHA-256 digest in hexadecimal form.
|
||||
"""
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as handle:
|
||||
for chunk in iter(lambda: handle.read(1 << 20), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _verify_checkpoint(path: Path) -> None:
|
||||
"""Verify that a TAEW2.1 checkpoint matches the expected SHA-256 digest.
|
||||
|
||||
Parameters:
|
||||
path (Path): Path to the checkpoint file.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the checkpoint digest does not match the expected value.
|
||||
"""
|
||||
actual = _sha256(path)
|
||||
if actual != TAEW2_1_CHECKPOINT_SHA256:
|
||||
raise RuntimeError(f"TAEHV checkpoint at {path} failed sha256 verification "
|
||||
f"(expected {TAEW2_1_CHECKPOINT_SHA256}, got {actual}). "
|
||||
"Delete the file to re-download it, or pass --taehv-checkpoint-path "
|
||||
"pointing at a checkpoint you trust.")
|
||||
|
||||
|
||||
def ensure_taew2_1_checkpoint(checkpoint_path: Path | None = None) -> Path:
|
||||
"""
|
||||
Ensure the TAEW2.1 checkpoint is available locally.
|
||||
|
||||
A caller-provided path is treated as trusted and is only checked for existence.
|
||||
The module-managed cached checkpoint is verified against the pinned SHA-256 digest
|
||||
after downloading or before reuse.
|
||||
|
||||
Parameters:
|
||||
checkpoint_path (Path | None): Optional path to a caller-provided checkpoint.
|
||||
|
||||
Returns:
|
||||
Path: The available checkpoint path.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If a caller-provided checkpoint does not exist.
|
||||
RuntimeError: If a module-managed checkpoint fails verification.
|
||||
"""
|
||||
if checkpoint_path is not None:
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"TAEHV checkpoint not found: {checkpoint_path}")
|
||||
return checkpoint_path
|
||||
|
||||
checkpoint_path = _default_cache_dir() / "taew2_1.pth"
|
||||
if not checkpoint_path.exists():
|
||||
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
print(f"Downloading {TAEW2_1_CHECKPOINT_URL} -> {checkpoint_path}")
|
||||
import socket
|
||||
import tempfile
|
||||
# Download to a temporary file, verify, then atomically rename.
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="wb",
|
||||
dir=checkpoint_path.parent,
|
||||
prefix=".tmp_taew2_1_",
|
||||
suffix=".pth",
|
||||
delete=False,
|
||||
) as tmp_file:
|
||||
tmp_path = Path(tmp_file.name)
|
||||
try:
|
||||
old_timeout = socket.getdefaulttimeout()
|
||||
socket.setdefaulttimeout(300)
|
||||
try:
|
||||
urllib.request.urlretrieve(
|
||||
TAEW2_1_CHECKPOINT_URL,
|
||||
tmp_path, # noqa: S310 - pinned public artifact, hash-verified below.
|
||||
)
|
||||
finally:
|
||||
socket.setdefaulttimeout(old_timeout)
|
||||
_verify_checkpoint(tmp_path)
|
||||
tmp_path.replace(checkpoint_path)
|
||||
except Exception:
|
||||
tmp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
else:
|
||||
_verify_checkpoint(checkpoint_path)
|
||||
return checkpoint_path
|
||||
|
||||
|
||||
def _load_taehv_class(source_path: Path | None):
|
||||
"""Load the TAEHV class from the vendored implementation or a local source override.
|
||||
|
||||
Parameters:
|
||||
source_path (Path | None): Path to a local Python file defining `TAEHV`; `None` selects the vendored implementation.
|
||||
|
||||
Returns:
|
||||
The loaded `TAEHV` class.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If the specified source cannot be loaded.
|
||||
"""
|
||||
if source_path is None:
|
||||
from fastvideo.third_party.taehv import TAEHV
|
||||
|
||||
return TAEHV
|
||||
# Explicit local override for experimenting with a modified TAEHV; this is
|
||||
# a user-supplied file on disk, never something this module downloads.
|
||||
spec = importlib.util.spec_from_file_location("fastvideo_external_taehv", source_path)
|
||||
if spec is None or spec.loader is None:
|
||||
raise RuntimeError(f"Could not load TAEHV source from {source_path}")
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module.TAEHV
|
||||
|
||||
|
||||
def decode_latents_to_video_taehv(
|
||||
*,
|
||||
latents_np: np.ndarray,
|
||||
output_path: Path,
|
||||
fps: int,
|
||||
device,
|
||||
dtype,
|
||||
parallel: bool,
|
||||
source_path: Path | None = None,
|
||||
checkpoint_path: Path | None = None,
|
||||
) -> None:
|
||||
"""Decode Wan/FastWan diffusion latents with TAEW2.1 and export MP4.
|
||||
|
||||
TAEHV's Wan wrapper expects the diffusion latents directly, without applying
|
||||
the standard Wan VAE's `latents_mean` / `latents_std` shift.
|
||||
"""
|
||||
import torch
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
checkpoint_path = ensure_taew2_1_checkpoint(checkpoint_path)
|
||||
TAEHV = _load_taehv_class(source_path)
|
||||
taehv = TAEHV(str(checkpoint_path)).to(device=device, dtype=dtype)
|
||||
taehv.eval()
|
||||
|
||||
latents = torch.from_numpy(latents_np).to(device=device, dtype=dtype)
|
||||
with torch.no_grad():
|
||||
video_ntchw = taehv.decode_video(
|
||||
latents.transpose(1, 2),
|
||||
parallel=parallel,
|
||||
show_progress_bar=False,
|
||||
)
|
||||
video = video_ntchw.transpose(1, 2)
|
||||
video_np = video[0].permute(1, 2, 3, 0).float().cpu().numpy()
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
export_to_video(video_np, str(output_path), fps=fps)
|
||||
@@ -1,286 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Wan2.2-TI2V-5B dense MLX runtime — Track D.
|
||||
|
||||
The Wan2.2 TI2V-5B (FullAttn) differs from the ported Wan2.1-T2V only in:
|
||||
|
||||
- **Scale** (24 heads x 128, hidden 3072, ffn 14336) — pure config, block math
|
||||
identical, so the dense loader ``mlx_dit_from_diffusers_safetensors`` loads the
|
||||
weights unchanged and we re-wrap the blocks here.
|
||||
- **Per-token timestep conditioning** (``expand_timesteps=True``): the timestep is
|
||||
``[batch, seq_len]`` (a level per patch token — how TI2V keeps the conditioning
|
||||
image frame at t=0 while the video frames are noised). ``timestep_proj`` becomes
|
||||
``[batch, seq_len, 6, dim]`` and the block/output modulation is per-token
|
||||
(``[B, L, dim]``), a direct broadcast — this module implements exactly that.
|
||||
|
||||
I2V rides on the same forward: encode the image, replace the first latent frame,
|
||||
and set that frame's timestep to 0 (handled by the caller / sampler). See
|
||||
``docs/design/ti2v_5b_port_guide.md``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from collections.abc import Callable
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.mlx_runtime.fastwan import (
|
||||
MLXWanT2VCrossAttention,
|
||||
gelu_tanh,
|
||||
layer_norm,
|
||||
linear,
|
||||
mlx_dit_from_diffusers_safetensors,
|
||||
rms_norm,
|
||||
silu,
|
||||
timestep_embedding,
|
||||
weight_dtype,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import mlx.core as mx
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MLXWan22TransformerBlock:
|
||||
"""Dense Wan block with per-token (``[B, L, dim]``) timestep modulation."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weights: dict[str, mx.array],
|
||||
*,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
eps: float = 1e-6,
|
||||
):
|
||||
self.weights = weights
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.eps = eps
|
||||
self.attn2 = MLXWanT2VCrossAttention(weights, dim=dim, num_heads=num_heads, eps=eps)
|
||||
|
||||
def __call__(self, hidden_states, encoder_hidden_states, timestep_proj, cos, sin) -> mx.array:
|
||||
import mlx.core as mx
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
batch = hidden_states.shape[0]
|
||||
|
||||
# timestep_proj: [B, L, 6, dim] -> six per-token [B, L, dim] modulations.
|
||||
e = self.weights["scale_shift_table"][None] + timestep_proj.astype(mx.float32)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = [
|
||||
part.squeeze(2) for part in mx.split(e, 6, axis=2)
|
||||
]
|
||||
|
||||
# 1. Self-attention (dense, bidirectional) with per-token modulation.
|
||||
norm_hidden = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
|
||||
norm_hidden = (norm_hidden * (1.0 + scale_msa) + shift_msa).astype(orig_dtype)
|
||||
|
||||
query = linear(norm_hidden, self.weights["to_q.weight"], self.weights.get("to_q.bias"))
|
||||
key = linear(norm_hidden, self.weights["to_k.weight"], self.weights.get("to_k.bias"))
|
||||
value = linear(norm_hidden, self.weights["to_v.weight"], self.weights.get("to_v.bias"))
|
||||
|
||||
query = rms_norm(query, self.weights["norm_q.weight"],
|
||||
eps=self.eps).reshape(batch, -1, self.num_heads, self.head_dim)
|
||||
key = rms_norm(key, self.weights["norm_k.weight"], eps=self.eps).reshape(batch, -1, self.num_heads,
|
||||
self.head_dim)
|
||||
value = value.reshape(batch, -1, self.num_heads, self.head_dim)
|
||||
|
||||
from fastvideo.mlx_runtime.fastwan import apply_rotary_emb
|
||||
query = apply_rotary_emb(query, cos, sin, is_neox_style=False)
|
||||
key = apply_rotary_emb(key, cos, sin, is_neox_style=False)
|
||||
|
||||
attn = mx.fast.scaled_dot_product_attention(
|
||||
query.transpose(0, 2, 1, 3),
|
||||
key.transpose(0, 2, 1, 3),
|
||||
value.transpose(0, 2, 1, 3),
|
||||
scale=self.head_dim**-0.5,
|
||||
).transpose(0, 2, 1, 3)
|
||||
attn = attn.reshape(batch, -1, self.dim)
|
||||
attn = linear(attn, self.weights["to_out.weight"], self.weights.get("to_out.bias"))
|
||||
|
||||
hidden_states = hidden_states + (attn * gate_msa).astype(orig_dtype)
|
||||
norm_hidden = layer_norm(hidden_states.astype(mx.float32),
|
||||
weight=self.weights["self_attn_residual_norm.norm.weight"],
|
||||
bias=self.weights["self_attn_residual_norm.norm.bias"],
|
||||
eps=self.eps).astype(orig_dtype)
|
||||
|
||||
# 2. Cross-attention, then per-token shift/scale modulation.
|
||||
cross = self.attn2(norm_hidden, encoder_hidden_states)
|
||||
hidden_states = hidden_states + cross
|
||||
norm_hidden = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
|
||||
norm_hidden = (norm_hidden * (1.0 + c_scale_msa) + c_shift_msa).astype(orig_dtype)
|
||||
|
||||
# 3. Feed-forward with per-token gate.
|
||||
ff = linear(norm_hidden, self.weights["ffn.fc_in.weight"], self.weights.get("ffn.fc_in.bias"))
|
||||
ff = gelu_tanh(ff)
|
||||
ff = linear(ff, self.weights["ffn.fc_out.weight"], self.weights.get("ffn.fc_out.bias"))
|
||||
hidden_states = hidden_states + (ff * c_gate_msa).astype(orig_dtype)
|
||||
return hidden_states.astype(orig_dtype)
|
||||
|
||||
|
||||
class MLXWan22DiT:
|
||||
"""Wan2.2-TI2V-5B dense DiT with per-token timestep conditioning."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
weights: dict[str, mx.array],
|
||||
blocks: list[MLXWan22TransformerBlock],
|
||||
config: dict,
|
||||
*,
|
||||
compile: bool = False,
|
||||
) -> None:
|
||||
import os
|
||||
|
||||
self.weights = weights
|
||||
self.blocks = blocks
|
||||
self.config = config
|
||||
self.num_heads = int(config["num_attention_heads"])
|
||||
self.head_dim = int(config["attention_head_dim"])
|
||||
self.hidden_size = self.num_heads * self.head_dim
|
||||
self.freq_dim = int(config["freq_dim"])
|
||||
self.patch_size = tuple(config["patch_size"])
|
||||
self.out_channels = int(config["out_channels"])
|
||||
self.eps = float(config.get("eps", 1e-6))
|
||||
self._enable_compile = compile or os.environ.get("FASTVIDEO_MLX_COMPILE", "0") == "1"
|
||||
self._compiled_forward: Callable[..., Any] | None = None
|
||||
self._compiled_signature: tuple | None = None
|
||||
|
||||
def _patch_embed(self, hidden_states) -> mx.array:
|
||||
batch, channels, frames, height, width = hidden_states.shape
|
||||
pt, ph, pw = self.patch_size
|
||||
patch_dim = channels * pt * ph * pw
|
||||
x = hidden_states.reshape(batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw)
|
||||
x = x.transpose(0, 2, 4, 6, 1, 3, 5, 7).reshape(batch, -1, patch_dim)
|
||||
return linear(x, self.weights["patch_embedding.weight"], self.weights.get("patch_embedding.bias"))
|
||||
|
||||
def _condition(self, timestep, encoder_hidden_states) -> tuple:
|
||||
"""Per-token conditioning. ``timestep`` is ``[B, L]`` (one level per token)."""
|
||||
batch, seq = timestep.shape
|
||||
t_freq = timestep_embedding(timestep.reshape(-1), self.freq_dim).astype(
|
||||
weight_dtype(self.weights["condition_embedder.time_embedder.linear_1.weight"]))
|
||||
temb = linear(t_freq, self.weights["condition_embedder.time_embedder.linear_1.weight"],
|
||||
self.weights["condition_embedder.time_embedder.linear_1.bias"])
|
||||
temb = silu(temb)
|
||||
temb = linear(temb, self.weights["condition_embedder.time_embedder.linear_2.weight"],
|
||||
self.weights["condition_embedder.time_embedder.linear_2.bias"])
|
||||
timestep_proj = linear(silu(temb), self.weights["condition_embedder.time_proj.weight"],
|
||||
self.weights["condition_embedder.time_proj.bias"])
|
||||
timestep_proj = timestep_proj.reshape(batch, seq, 6, self.hidden_size)
|
||||
|
||||
ehs = linear(encoder_hidden_states, self.weights["condition_embedder.text_embedder.linear_1.weight"],
|
||||
self.weights["condition_embedder.text_embedder.linear_1.bias"])
|
||||
ehs = gelu_tanh(ehs)
|
||||
ehs = linear(ehs, self.weights["condition_embedder.text_embedder.linear_2.weight"],
|
||||
self.weights["condition_embedder.text_embedder.linear_2.bias"])
|
||||
temb_out = temb.reshape(batch, seq, self.hidden_size)
|
||||
return temb_out, timestep_proj, ehs
|
||||
|
||||
def _output(self, hidden_states, temb_out, *, batch, frames, height, width) -> mx.array:
|
||||
import mlx.core as mx
|
||||
|
||||
pt, ph, pw = self.patch_size
|
||||
post_pt, post_ph, post_pw = frames // pt, height // ph, width // pw
|
||||
# Per-token output modulation: scale_shift_table[1,2,dim] + temb[B,L,1,dim].
|
||||
e = self.weights["scale_shift_table"][None] + temb_out[:, :, None, :].astype(mx.float32)
|
||||
shift, scale = [part.squeeze(2) for part in mx.split(e, 2, axis=2)]
|
||||
norm = layer_norm(hidden_states.astype(mx.float32), eps=self.eps)
|
||||
norm = (norm * (1.0 + scale) + shift).astype(weight_dtype(self.weights["proj_out.weight"]))
|
||||
out = linear(norm, self.weights["proj_out.weight"], self.weights["proj_out.bias"])
|
||||
out = out.reshape(batch, post_pt, post_ph, post_pw, pt, ph, pw, self.out_channels)
|
||||
out = out.transpose(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
return out.reshape(batch, self.out_channels, frames, height, width)
|
||||
|
||||
def _forward(self, hidden_states, encoder_hidden_states, timestep, cos, sin) -> mx.array:
|
||||
batch, _, frames, height, width = hidden_states.shape
|
||||
hidden = self._patch_embed(hidden_states)
|
||||
temb_out, timestep_proj, ehs = self._condition(timestep, encoder_hidden_states)
|
||||
for block in self.blocks:
|
||||
hidden = block(hidden, ehs, timestep_proj, cos, sin)
|
||||
return self._output(hidden, temb_out, batch=batch, frames=frames, height=height, width=width)
|
||||
|
||||
def __call__(self, hidden_states, encoder_hidden_states, timestep, freqs_cis) -> mx.array:
|
||||
cos, sin = freqs_cis
|
||||
if self._enable_compile and cos is not None:
|
||||
import mlx.core as mx
|
||||
|
||||
# One traced graph per input signature, and each pins its own copy of
|
||||
# the quantized weights. --refine denoises at two resolutions, so
|
||||
# keeping both alive doubles resident DiT memory. Retire the previous
|
||||
# graph when the signature changes.
|
||||
signature = (hidden_states.shape, encoder_hidden_states.shape, timestep.shape)
|
||||
if self._compiled_forward is not None and signature != self._compiled_signature:
|
||||
self._compiled_forward = None
|
||||
self._compiled_signature = None
|
||||
mx.clear_cache()
|
||||
if self._compiled_forward is None:
|
||||
self._compiled_forward = mx.compile(self._forward)
|
||||
self._compiled_signature = signature
|
||||
compiled_forward = self._compiled_forward
|
||||
try:
|
||||
return compiled_forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
|
||||
except Exception as exc: # noqa: BLE001 - some graphs may not trace; fall back to eager.
|
||||
logger.warning(
|
||||
"Wan2.2 mx.compile forward failed (%s); falling back to eager execution.",
|
||||
exc,
|
||||
)
|
||||
self._enable_compile = False
|
||||
self._compiled_forward = None
|
||||
self._compiled_signature = None
|
||||
return self._forward(hidden_states, encoder_hidden_states, timestep, cos, sin)
|
||||
|
||||
|
||||
def mlx_wan22_dit_from_diffusers_safetensors(
|
||||
checkpoint_path: str | Path,
|
||||
config_path: str | Path,
|
||||
*,
|
||||
dtype: str = "fp16",
|
||||
num_blocks: int | None = None,
|
||||
quantization=None,
|
||||
compile: bool = False,
|
||||
) -> MLXWan22DiT:
|
||||
"""Load Wan2.2-TI2V-5B (FullAttn) into ``MLXWan22DiT`` via the dense loader."""
|
||||
dense = mlx_dit_from_diffusers_safetensors(checkpoint_path,
|
||||
config_path,
|
||||
dtype=dtype,
|
||||
num_blocks=num_blocks,
|
||||
quantization=quantization)
|
||||
inner_dim = int(dense.config["num_attention_heads"]) * int(dense.config["attention_head_dim"])
|
||||
blocks = [
|
||||
MLXWan22TransformerBlock(block.weights,
|
||||
dim=inner_dim,
|
||||
ffn_dim=int(dense.config["ffn_dim"]),
|
||||
num_heads=int(dense.config["num_attention_heads"]),
|
||||
eps=float(dense.config.get("eps", 1e-6))) for block in dense.blocks
|
||||
]
|
||||
return MLXWan22DiT(dense.weights, blocks, dense.config, compile=compile)
|
||||
|
||||
|
||||
def mlx_wan22_dit_from_mlx_checkpoint(
|
||||
checkpoint_dir: str | Path,
|
||||
*,
|
||||
compile: bool = False,
|
||||
) -> MLXWan22DiT:
|
||||
"""Rewrap a persisted MLX DiT checkpoint with Wan2.2 conditioning.
|
||||
|
||||
The generic checkpoint loader intentionally rebuilds ``MLXWanDiT`` because
|
||||
it is also used by the Wan2.1 runtime. Wan2.2 TI2V has the same weight
|
||||
layout but needs per-token timestep modulation, so callers must rewrap the
|
||||
loaded weights and blocks as :class:`MLXWan22DiT` before sampling.
|
||||
"""
|
||||
from fastvideo.mlx_runtime.checkpoint import load_mlx_dit_checkpoint
|
||||
|
||||
dense = load_mlx_dit_checkpoint(checkpoint_dir)
|
||||
inner_dim = int(dense.config["num_attention_heads"]) * int(dense.config["attention_head_dim"])
|
||||
blocks = [
|
||||
MLXWan22TransformerBlock(
|
||||
block.weights,
|
||||
dim=inner_dim,
|
||||
ffn_dim=int(dense.config["ffn_dim"]),
|
||||
num_heads=int(dense.config["num_attention_heads"]),
|
||||
eps=float(dense.config.get("eps", 1e-6)),
|
||||
) for block in dense.blocks
|
||||
]
|
||||
return MLXWan22DiT(dense.weights, blocks, dense.config, compile=compile)
|
||||
@@ -1,113 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Dense DMD sampling for MLXWan22DiT (Wan2.2 per-token timestep).
|
||||
|
||||
Matches the FastVideo pipeline's warped DMD schedule (``warp_denoising_step=True``,
|
||||
``dmd_denoising_steps=[1000,757,522]``, ``flow_shift=5.0`` for TI2V-5B) rather
|
||||
than treating raw step indices as continuous timesteps (a bug in early demos).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from collections.abc import Sequence
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.mlx_runtime.sampling import MLXDMDSchedule, dmd_step, pred_noise_to_pred_video
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import mlx.core as mx
|
||||
|
||||
from fastvideo.mlx_runtime.wan22 import MLXWan22DiT
|
||||
|
||||
|
||||
def build_wan22_dmd_schedule(
|
||||
dmd_denoising_steps: Sequence[int] | None = None,
|
||||
*,
|
||||
flow_shift: float = 5.0,
|
||||
warp_denoising_step: bool = True,
|
||||
) -> tuple[MLXDMDSchedule, list[float]]:
|
||||
"""
|
||||
Build the flow-matching schedule and continuous timesteps used for Wan2.2 DMD sampling.
|
||||
|
||||
Parameters:
|
||||
dmd_denoising_steps (Sequence[int] | None): Denoising step values to use; defaults to 1000, 757, and 522.
|
||||
flow_shift (float): Flow-matching shift applied when constructing the schedule.
|
||||
warp_denoising_step (bool): Whether to convert denoising steps to scheduler-warped continuous timesteps.
|
||||
|
||||
Returns:
|
||||
tuple[MLXDMDSchedule, list[float]]: The DMD schedule and corresponding continuous timesteps.
|
||||
"""
|
||||
import torch
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
steps = list(dmd_denoising_steps or [1000, 757, 522])
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=flow_shift)
|
||||
schedule = MLXDMDSchedule.from_torch_scheduler(scheduler)
|
||||
step_idx = torch.tensor(steps, dtype=torch.long)
|
||||
if warp_denoising_step:
|
||||
warped = torch.cat((scheduler.timesteps.cpu(), torch.tensor([0.0], dtype=torch.float32)))
|
||||
timesteps = [float(t) for t in warped[1000 - step_idx]]
|
||||
else:
|
||||
timesteps = [float(s) for s in steps]
|
||||
return schedule, timesteps
|
||||
|
||||
|
||||
def sample_wan22_dmd(
|
||||
model: MLXWan22DiT,
|
||||
encoder_hidden_states: mx.array,
|
||||
noise_latents: mx.array,
|
||||
freqs_cis: tuple,
|
||||
*,
|
||||
dmd_denoising_steps: Sequence[int] | None = None,
|
||||
flow_shift: float = 5.0,
|
||||
warp_denoising_step: bool = True,
|
||||
seed: int = 0,
|
||||
) -> mx.array:
|
||||
"""
|
||||
Generate clean video latents from noisy latents using iterative DMD denoising.
|
||||
|
||||
Parameters:
|
||||
noise_latents (mx.array): Initial noisy video latents.
|
||||
freqs_cis (tuple): Rotary positional frequency tensors used by the model.
|
||||
dmd_denoising_steps (Sequence[int] | None): DMD denoising steps, or the default schedule when omitted.
|
||||
flow_shift (float): Flow-matching schedule shift.
|
||||
warp_denoising_step (bool): Whether to warp the denoising timesteps.
|
||||
seed (int): Seed for reproducible intermediate re-noising.
|
||||
|
||||
Returns:
|
||||
mx.array: Denoised video latents.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
schedule, timesteps = build_wan22_dmd_schedule(dmd_denoising_steps,
|
||||
flow_shift=flow_shift,
|
||||
warp_denoising_step=warp_denoising_step)
|
||||
# NumPy RNG so re-noise is bit-reproducible across MLX / torch A/B dumps.
|
||||
renoise_rng = np.random.default_rng(seed)
|
||||
latents = noise_latents
|
||||
batch, _c, frames, height, width = latents.shape
|
||||
pt, ph, pw = model.patch_size
|
||||
tokens = (frames // pt) * (height // ph) * (width // pw)
|
||||
last = len(timesteps) - 1
|
||||
for i, t in enumerate(timesteps):
|
||||
ts = mx.full((batch, tokens), float(t), dtype=mx.float32)
|
||||
pred = model(latents.astype(mx.float16), encoder_hidden_states, ts, freqs_cis)
|
||||
ni = latents.astype(mx.float32)
|
||||
pn = pred.astype(mx.float32)
|
||||
if i < last:
|
||||
renoise = mx.array(renoise_rng.standard_normal(tuple(latents.shape)).astype(np.float32))
|
||||
latents = dmd_step(
|
||||
latents=ni,
|
||||
noise_input_latent=ni,
|
||||
pred_noise=pn,
|
||||
schedule=schedule,
|
||||
timestep=float(t),
|
||||
next_timestep=float(timesteps[i + 1]),
|
||||
noise=renoise,
|
||||
).astype(latents.dtype)
|
||||
else:
|
||||
latents = pred_noise_to_pred_video(pn, ni, schedule.sigma_for(float(t))).astype(latents.dtype)
|
||||
mx.eval(latents)
|
||||
return latents
|
||||
@@ -1,557 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Wan VAE decode helpers for Apple Silicon MLX inference.
|
||||
|
||||
Two decode backends:
|
||||
|
||||
1. **TAEHV (primary / fast)** — Tiny AutoEncoder (madebyollin/taehv). Fully
|
||||
MLX-native Conv2d path. ``taew2_1.pth`` for Wan2.1 (z_dim=16),
|
||||
``taew2_2.pth`` for Wan2.2 5B (z_dim=48, patch_size=2). Expected decode
|
||||
wall-clock ~seconds vs ~minutes for the full 3D VAE on MPS.
|
||||
|
||||
2. **Full AutoencoderKLWan (reference / quality)** — denormalize with
|
||||
``latents_mean`` / ``latents_std`` then torch decode (MPS preferred). Used
|
||||
for parity gates and when TAEHV is unavailable. A pure-MLX 3D-conv port of
|
||||
the residual Wan2.2 decoder is left as follow-up (causal feat-cache +
|
||||
residual up blocks are large); TAEHV covers the product latency path.
|
||||
|
||||
Diffusion latents from the DiT are **not** mean/std-normalized for TAEHV
|
||||
(matching ``taehv_decode.py``); full VAE decode **does** denormalize first
|
||||
(matching ``mlx_wan_prompt_to_video.decode_latents_to_video``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import urllib.request
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.mlx_runtime.memory import cleanup_mlx, cleanup_torch_mps
|
||||
|
||||
GIB = 1024**3
|
||||
|
||||
TAEW2_1_URL = "https://raw.githubusercontent.com/madebyollin/taehv/main/taew2_1.pth"
|
||||
TAEW2_2_URL = ("https://raw.githubusercontent.com/madebyollin/taehv/"
|
||||
"563f40bdc820ed86bcad72ea515ee48f06bd22ec/taew2_2.pth")
|
||||
# Validated 2026-07-02 / 2026-07-09 against upstream madebyollin/taehv.
|
||||
TAEW2_1_SHA256 = "d26151e76cdc2c9424bef988de874b33d9a53f30ef3060cd556c429c469c797e"
|
||||
TAEW2_2_SHA256 = "d053e216ca50e2bb837bbcd79b85f0366bea00e5938025572382a773b74c559a"
|
||||
|
||||
DecodeBackend = Literal["taehv", "taehv-torch", "wan-vae"]
|
||||
|
||||
|
||||
def _cache_dir() -> Path:
|
||||
"""Return the local directory used to cache TAEHV files."""
|
||||
return Path.home() / ".cache" / "fastvideo" / "taehv"
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
"""Compute the SHA-256 digest of a file.
|
||||
|
||||
Parameters:
|
||||
path (Path): Path to the file to hash.
|
||||
|
||||
Returns:
|
||||
str: Lowercase hexadecimal SHA-256 digest.
|
||||
"""
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as handle:
|
||||
for chunk in iter(lambda: handle.read(1 << 20), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _verify_checkpoint(path: Path, expected_digest: str) -> None:
|
||||
"""
|
||||
Verify a checkpoint's SHA-256 digest against a required expected digest.
|
||||
|
||||
Parameters:
|
||||
path (Path): Path to the checkpoint file.
|
||||
expected_digest (str): Expected lowercase SHA-256 digest.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If verification is enabled and the checkpoint digest does not match.
|
||||
"""
|
||||
if len(expected_digest) != 64 or any(char not in "0123456789abcdef" for char in expected_digest):
|
||||
raise ValueError("A valid lowercase SHA-256 digest is required for bundled TAEHV checkpoints")
|
||||
actual = _sha256(path)
|
||||
if actual != expected_digest:
|
||||
raise RuntimeError(f"TAEHV checkpoint at {path} failed sha256 verification "
|
||||
f"(expected {expected_digest}, got {actual}). "
|
||||
"Delete the file to re-download it.")
|
||||
|
||||
|
||||
def ensure_taehv_checkpoint(*, z_dim: int, checkpoint_path: Path | None = None) -> Path:
|
||||
"""
|
||||
Return a validated TAEHV checkpoint for the specified latent channel count.
|
||||
|
||||
Parameters:
|
||||
z_dim (int): Number of latent channels, supported values are 16 and 48.
|
||||
checkpoint_path (Path | None): Optional existing checkpoint path to validate and use.
|
||||
|
||||
Returns:
|
||||
Path: Path to the validated TAEHV checkpoint.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the supplied checkpoint path does not exist.
|
||||
ValueError: If no checkpoint is mapped to the specified latent channel count.
|
||||
"""
|
||||
if checkpoint_path is not None:
|
||||
if not checkpoint_path.exists():
|
||||
raise FileNotFoundError(f"TAEHV checkpoint not found: {checkpoint_path}")
|
||||
return checkpoint_path
|
||||
if z_dim == 16:
|
||||
name, url, expect = "taew2_1.pth", TAEW2_1_URL, TAEW2_1_SHA256
|
||||
elif z_dim == 48:
|
||||
name, url, expect = "taew2_2.pth", TAEW2_2_URL, TAEW2_2_SHA256
|
||||
else:
|
||||
raise ValueError(f"No TAEHV checkpoint mapped for z_dim={z_dim} (supported: 16, 48)")
|
||||
path = _cache_dir() / name
|
||||
if not path.exists():
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
print(f"Downloading {url} -> {path}")
|
||||
import socket
|
||||
import tempfile
|
||||
# Download to a temporary file, verify, then atomically rename.
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="wb",
|
||||
dir=path.parent,
|
||||
prefix=f".tmp_{name}_",
|
||||
suffix=".pth",
|
||||
delete=False,
|
||||
) as tmp_file:
|
||||
tmp_path = Path(tmp_file.name)
|
||||
try:
|
||||
old_timeout = socket.getdefaulttimeout()
|
||||
socket.setdefaulttimeout(300)
|
||||
try:
|
||||
urllib.request.urlretrieve(url, tmp_path) # noqa: S310 - public pinned artifact.
|
||||
finally:
|
||||
socket.setdefaulttimeout(old_timeout)
|
||||
_verify_checkpoint(tmp_path, expect)
|
||||
tmp_path.replace(path)
|
||||
except Exception:
|
||||
tmp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
else:
|
||||
_verify_checkpoint(path, expect)
|
||||
return path
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WanVAEConfigView:
|
||||
"""Minimal config fields needed for denormalize + spatial scale."""
|
||||
|
||||
z_dim: int
|
||||
latents_mean: tuple[float, ...]
|
||||
latents_std: tuple[float, ...]
|
||||
scale_factor_spatial: int = 8
|
||||
scale_factor_temporal: int = 4
|
||||
patch_size: int | None = None
|
||||
vae_dir: Path | None = None
|
||||
|
||||
@classmethod
|
||||
def from_vae_dir(cls, vae_dir: Path) -> WanVAEConfigView:
|
||||
"""
|
||||
Load Wan VAE configuration values from a directory.
|
||||
|
||||
Parameters:
|
||||
vae_dir (Path): Directory containing the VAE ``config.json`` file.
|
||||
|
||||
Returns:
|
||||
WanVAEConfigView: Configuration loaded from the VAE directory.
|
||||
"""
|
||||
cfg = json.loads((vae_dir / "config.json").read_text())
|
||||
return cls(
|
||||
z_dim=int(cfg["z_dim"]),
|
||||
latents_mean=tuple(float(x) for x in cfg["latents_mean"]),
|
||||
latents_std=tuple(float(x) for x in cfg["latents_std"]),
|
||||
scale_factor_spatial=int(cfg.get("scale_factor_spatial", 8)),
|
||||
scale_factor_temporal=int(cfg.get("scale_factor_temporal", 4)),
|
||||
patch_size=cfg.get("patch_size"),
|
||||
vae_dir=vae_dir,
|
||||
)
|
||||
|
||||
|
||||
def denormalize_latents_np(latents: np.ndarray, config: WanVAEConfigView) -> np.ndarray:
|
||||
"""
|
||||
Denormalize Wan VAE latent values using the configured means and standard deviations.
|
||||
|
||||
Parameters:
|
||||
latents (np.ndarray): Latent values in normalized form.
|
||||
config (WanVAEConfigView): Wan VAE latent statistics.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Denormalized latent values as float32.
|
||||
"""
|
||||
mean = np.asarray(config.latents_mean, dtype=np.float32).reshape(1, -1, 1, 1, 1)
|
||||
std = np.asarray(config.latents_std, dtype=np.float32).reshape(1, -1, 1, 1, 1)
|
||||
return latents.astype(np.float32) * std + mean
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# MLX TAEHV decoder (Conv2d stack — primary fully-MLX product path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mlx_conv2d(x: Any, weight: Any, bias: Any, *, stride: int = 1) -> Any:
|
||||
import mlx.core as mx
|
||||
|
||||
# x: NCHW, weight: OIHW
|
||||
y = mx.conv2d(x.transpose(0, 2, 3, 1), weight.transpose(0, 2, 3, 1), stride=stride, padding=1)
|
||||
y = y.transpose(0, 3, 1, 2)
|
||||
if bias is not None:
|
||||
y = y + bias.reshape(1, -1, 1, 1)
|
||||
return y
|
||||
|
||||
|
||||
def _mlx_conv2d_1x1(x: Any, weight: Any, bias: Any = None) -> Any:
|
||||
"""
|
||||
Applies a 1×1 convolution to an MLX tensor in channel-first layout.
|
||||
|
||||
Parameters:
|
||||
x (Any): Input tensor with shape [batch, channels, height, width].
|
||||
weight (Any): Convolution weights.
|
||||
bias (Any, optional): Optional output-channel bias.
|
||||
|
||||
Returns:
|
||||
Any: The convolved tensor with shape [batch, output_channels, height, width].
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
y = mx.conv2d(x.transpose(0, 2, 3, 1), weight.transpose(0, 2, 3, 1), stride=1, padding=0)
|
||||
y = y.transpose(0, 3, 1, 2)
|
||||
if bias is not None:
|
||||
y = y + bias.reshape(1, -1, 1, 1)
|
||||
return y
|
||||
|
||||
|
||||
def _load_torch_state(path: Path) -> dict[str, np.ndarray]:
|
||||
"""Load a PyTorch state dictionary as NumPy arrays.
|
||||
|
||||
Parameters:
|
||||
path (Path): Path to the PyTorch checkpoint.
|
||||
|
||||
Returns:
|
||||
dict[str, np.ndarray]: State dictionary with tensors converted to NumPy arrays.
|
||||
"""
|
||||
import torch
|
||||
|
||||
sd = torch.load(path, map_location="cpu", weights_only=True)
|
||||
return {k: v.detach().float().cpu().numpy() for k, v in sd.items()}
|
||||
|
||||
|
||||
class MLXTAEHVDecoder:
|
||||
"""Minimal MLX port of TAEHV ``decoder`` (parallel-over-time MemBlocks)."""
|
||||
|
||||
def __init__(self, checkpoint_path: Path, *, z_dim: int) -> None:
|
||||
"""Initialize the TAEHV decoder from a checkpoint for the specified latent dimensionality.
|
||||
|
||||
Parameters:
|
||||
checkpoint_path (Path): Path to the TAEHV checkpoint.
|
||||
z_dim (int): Number of latent channels, determining the decoder patch size.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
self.checkpoint_path = Path(checkpoint_path)
|
||||
self.latent_channels = z_dim
|
||||
# Derive patch_size from z_dim: 48 channels → patch_size=2, 16 → patch_size=1
|
||||
self.patch_size = 2 if z_dim == 48 else 1
|
||||
self.image_channels = 3
|
||||
self.frames_to_trim = 3 # TGrow strides (1,2,2) → 2**2 - 1 for w2.1/w2.2 defaults
|
||||
|
||||
sd = _load_torch_state(self.checkpoint_path)
|
||||
# Patch TGrow kernels like upstream TAEHV.patch_tgrow_layers.
|
||||
self.weights = {k: mx.array(v) for k, v in sd.items()}
|
||||
self._n_f = [256, 128, 64, 64]
|
||||
|
||||
def decode_ntchw(self, latents_ntchw: Any) -> Any:
|
||||
"""
|
||||
Decode latent video batches into clipped RGB frames.
|
||||
|
||||
Parameters:
|
||||
latents_ntchw (Any): Latents with shape ``[N, T, C, H, W]`` and the
|
||||
decoder's configured latent channel count.
|
||||
|
||||
Returns:
|
||||
Any: Decoded frames with shape ``[N, T_out, 3, H_out, W_out]`` and values
|
||||
clipped to the range ``[0, 1]``.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
x = latents_ntchw
|
||||
n, t, c, h, w = x.shape
|
||||
if c != self.latent_channels:
|
||||
raise ValueError(f"expected C={self.latent_channels}, got {c}")
|
||||
x = x.reshape(n * t, c, h, w)
|
||||
x = self._run_decoder_parallel(x, n=n)
|
||||
# Pixel-shuffle if patch_size > 1: (NT, 3*p*p, H, W) -> (NT, 3, H*p, W*p)
|
||||
if self.patch_size > 1:
|
||||
p = self.patch_size
|
||||
nt, c_out, hh, ww = x.shape
|
||||
x = x.reshape(nt, self.image_channels, p, p, hh, ww)
|
||||
x = x.transpose(0, 1, 4, 2, 5, 3).reshape(nt, self.image_channels, hh * p, ww * p)
|
||||
_, c_out, hh, ww = x.shape
|
||||
t_out = x.shape[0] // n
|
||||
x = x.reshape(n, t_out, c_out, hh, ww)
|
||||
if self.frames_to_trim > 0 and t_out > self.frames_to_trim:
|
||||
x = x[:, self.frames_to_trim:]
|
||||
return mx.clip(x, 0.0, 1.0)
|
||||
|
||||
def _run_decoder_parallel(self, x: Any, *, n: int) -> Any:
|
||||
"""
|
||||
Apply the TAEHV decoder stack to flattened batch and temporal frames while preserving temporal memory.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
w = self.weights
|
||||
|
||||
def memblock(base: int, xx: Any, past: Any) -> Any:
|
||||
"""
|
||||
Apply a temporal memory block to the current and past feature tensors.
|
||||
|
||||
Parameters:
|
||||
base (int): Decoder block index used to select the block weights.
|
||||
xx (Any): Current feature tensor.
|
||||
past (Any): Past feature tensor concatenated with the current features.
|
||||
|
||||
Returns:
|
||||
Any: Activated feature tensor produced by the memory block.
|
||||
"""
|
||||
cat = mx.concatenate([xx, past], axis=1)
|
||||
h = _mlx_conv2d(cat, w[f"decoder.{base}.conv.0.weight"], w.get(f"decoder.{base}.conv.0.bias"))
|
||||
h = mx.maximum(h, 0.0)
|
||||
h = _mlx_conv2d(h, w[f"decoder.{base}.conv.2.weight"], w.get(f"decoder.{base}.conv.2.bias"))
|
||||
h = mx.maximum(h, 0.0)
|
||||
h = _mlx_conv2d(h, w[f"decoder.{base}.conv.4.weight"], w.get(f"decoder.{base}.conv.4.bias"))
|
||||
skip_key = f"decoder.{base}.skip.weight"
|
||||
skip = _mlx_conv2d_1x1(xx, w[skip_key], None) if skip_key in w else xx
|
||||
return mx.maximum(h + skip, 0.0)
|
||||
|
||||
def upsample2(xx: Any) -> Any:
|
||||
"""
|
||||
Upsample a four-dimensional tensor by a factor of two along its spatial dimensions.
|
||||
|
||||
Parameters:
|
||||
xx (Any): Tensor with shape `(N, C, H, W)`.
|
||||
|
||||
Returns:
|
||||
Any: Tensor with shape `(N, C, 2H, 2W)` containing replicated spatial values.
|
||||
"""
|
||||
nt, c, h, ww = xx.shape
|
||||
xx = xx.reshape(nt, c, h, 1, ww, 1)
|
||||
xx = mx.broadcast_to(xx, (nt, c, h, 2, ww, 2))
|
||||
return xx.reshape(nt, c, h * 2, ww * 2)
|
||||
|
||||
def tgrow(base: int, xx: Any, stride: int) -> Any:
|
||||
wt = w[f"decoder.{base}.conv.weight"]
|
||||
out_ch = int(xx.shape[1]) * stride
|
||||
if int(wt.shape[0]) > out_ch:
|
||||
wt = wt[-out_ch:]
|
||||
y = _mlx_conv2d_1x1(xx, wt, None)
|
||||
if stride == 1:
|
||||
return y
|
||||
# TGrow.forward: (NT, C*stride, H, W) -> (NT*stride, C, H, W)
|
||||
nt, c, h, ww = y.shape
|
||||
c_in = c // stride
|
||||
y = y.reshape(nt, stride, c_in, h, ww).transpose(0, 1, 2, 3, 4)
|
||||
return y.reshape(nt * stride, c_in, h, ww)
|
||||
|
||||
def mem_past(xx: Any) -> Any:
|
||||
"""
|
||||
Build a temporal memory tensor containing a zero frame followed by the preceding frame at each time step.
|
||||
|
||||
Parameters:
|
||||
xx (Any): Flattened batch and temporal tensor with shape ``(batch * time, channels, height, width)``.
|
||||
|
||||
Returns:
|
||||
Any: Tensor with the same shape as ``xx`` containing the preceding frame for each temporal position.
|
||||
"""
|
||||
nt, c, h, ww = xx.shape
|
||||
t_cur = nt // n
|
||||
x_ = xx.reshape(n, t_cur, c, h, ww)
|
||||
# pad one zero frame at t=0, align past[t] = x[t-1]
|
||||
past = mx.concatenate([mx.zeros_like(x_[:, :1]), x_[:, :-1]], axis=1)
|
||||
return past.reshape(nt, c, h, ww)
|
||||
|
||||
# 0 Clamp, 1 conv, 2 ReLU
|
||||
x = mx.tanh(x / 3.0) * 3.0
|
||||
x = _mlx_conv2d(x, w["decoder.1.weight"], w.get("decoder.1.bias"))
|
||||
x = mx.maximum(x, 0.0)
|
||||
for mem_idx in (3, 4, 5):
|
||||
x = memblock(mem_idx, x, mem_past(x))
|
||||
x = upsample2(x)
|
||||
x = tgrow(7, x, 1)
|
||||
x = _mlx_conv2d(x, w["decoder.8.weight"], w.get("decoder.8.bias"))
|
||||
for mem_idx in (9, 10, 11):
|
||||
x = memblock(mem_idx, x, mem_past(x))
|
||||
x = upsample2(x)
|
||||
x = tgrow(13, x, 2)
|
||||
x = _mlx_conv2d(x, w["decoder.14.weight"], w.get("decoder.14.bias"))
|
||||
for mem_idx in (15, 16, 17):
|
||||
x = memblock(mem_idx, x, mem_past(x))
|
||||
x = upsample2(x)
|
||||
x = tgrow(19, x, 2)
|
||||
x = _mlx_conv2d(x, w["decoder.20.weight"], w.get("decoder.20.bias"))
|
||||
x = mx.maximum(x, 0.0)
|
||||
x = _mlx_conv2d(x, w["decoder.22.weight"], w.get("decoder.22.bias"))
|
||||
return x
|
||||
|
||||
|
||||
def decode_latents_taehv_mlx(
|
||||
latents_np: np.ndarray,
|
||||
*,
|
||||
z_dim: int | None = None,
|
||||
checkpoint_path: Path | None = None,
|
||||
) -> np.ndarray:
|
||||
"""
|
||||
Decode latent representations with the MLX TAEHV decoder.
|
||||
|
||||
Parameters:
|
||||
latents_np (np.ndarray): Latents arranged as [B, C, T, H, W].
|
||||
z_dim (int | None): Latent channel dimension used to select the decoder checkpoint.
|
||||
checkpoint_path (Path | None): Optional path to a TAEHV checkpoint.
|
||||
|
||||
Returns:
|
||||
np.ndarray: Decoded pixels arranged as [B, T, H, W, 3] with values in [0, 1].
|
||||
|
||||
Raises:
|
||||
ValueError: If `latents_np` does not have five dimensions.
|
||||
"""
|
||||
import mlx.core as mx
|
||||
|
||||
if latents_np.ndim != 5:
|
||||
raise ValueError(f"expected [B,C,T,H,W], got {latents_np.shape}")
|
||||
c = latents_np.shape[1]
|
||||
z = z_dim if z_dim is not None else c
|
||||
ckpt = ensure_taehv_checkpoint(z_dim=z, checkpoint_path=checkpoint_path)
|
||||
dec = MLXTAEHVDecoder(ckpt, z_dim=z)
|
||||
# NTCHW
|
||||
x = mx.array(latents_np.transpose(0, 2, 1, 3, 4).astype(np.float32))
|
||||
out = dec.decode_ntchw(x) # N T C H W
|
||||
mx.eval(out)
|
||||
arr = np.array(out)
|
||||
# B T H W C
|
||||
return arr.transpose(0, 1, 3, 4, 2)
|
||||
|
||||
|
||||
def decode_latents_wan_vae_torch(
|
||||
latents_np: np.ndarray,
|
||||
*,
|
||||
vae_dir: Path,
|
||||
device: str = "auto",
|
||||
dtype_name: str = "fp16",
|
||||
) -> np.ndarray:
|
||||
"""Full AutoencoderKLWan decode on torch (MPS/CPU) with mean/std denormalize.
|
||||
|
||||
Returns pixels ``[B, T, H, W, 3]`` float in [0, 1].
|
||||
"""
|
||||
import torch
|
||||
from diffusers import AutoencoderKLWan
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
if device == "auto":
|
||||
device = "mps" if torch.backends.mps.is_available() else "cpu"
|
||||
dtype = torch.float16 if dtype_name == "fp16" and device == "mps" else torch.float32
|
||||
config = WanVAEConfigView.from_vae_dir(vae_dir)
|
||||
vae = AutoencoderKLWan.from_pretrained(vae_dir, torch_dtype=dtype, local_files_only=True).to(device)
|
||||
vae.eval()
|
||||
latents = torch.from_numpy(latents_np.astype(np.float32)).to(device=device, dtype=dtype)
|
||||
mean = torch.tensor(config.latents_mean, device=device, dtype=dtype).view(1, -1, 1, 1, 1)
|
||||
inv_std = (1.0 / torch.tensor(config.latents_std, device=device, dtype=dtype)).view(1, -1, 1, 1, 1)
|
||||
latents = latents / inv_std + mean # matches prompt_to_video path
|
||||
with torch.no_grad():
|
||||
video = vae.decode(latents, return_dict=False)[0]
|
||||
video = VideoProcessor(vae_scale_factor=config.scale_factor_spatial).postprocess_video(video, output_type="np")
|
||||
return video # [B, T, H, W, 3]
|
||||
|
||||
|
||||
def decode_latents_to_video(
|
||||
latents_np: np.ndarray,
|
||||
output_path: Path,
|
||||
*,
|
||||
fps: int = 16,
|
||||
backend: DecodeBackend = "taehv",
|
||||
vae_dir: Path | None = None,
|
||||
z_dim: int | None = None,
|
||||
taehv_checkpoint: Path | None = None,
|
||||
torch_device: str = "auto",
|
||||
) -> dict[str, Any]:
|
||||
"""Decode latent video frames and export them as an MP4 file.
|
||||
|
||||
Parameters:
|
||||
latents_np (np.ndarray): Latent video representation to decode.
|
||||
output_path (Path): Destination path for the MP4 file.
|
||||
fps (int): Output video frame rate.
|
||||
backend (DecodeBackend): Decoder backend to use.
|
||||
vae_dir (Path | None): Directory containing the full Wan VAE when using
|
||||
the ``wan-vae`` backend.
|
||||
z_dim (int | None): Latent channel count for TAEHV decoding.
|
||||
taehv_checkpoint (Path | None): Optional TAEHV checkpoint path.
|
||||
torch_device (str): PyTorch device selection for PyTorch-based decoding.
|
||||
|
||||
Returns:
|
||||
dict[str, Any]: Decode time in seconds, backend name, output path, frame
|
||||
count, and video resolution.
|
||||
|
||||
Raises:
|
||||
ValueError: If the full VAE backend lacks ``vae_dir`` or the backend is
|
||||
unknown.
|
||||
"""
|
||||
import time
|
||||
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
t0 = time.perf_counter()
|
||||
if backend in ("taehv", "taehv-torch"):
|
||||
c = latents_np.shape[1] if z_dim is None else z_dim
|
||||
if backend == "taehv":
|
||||
video = decode_latents_taehv_mlx(latents_np, z_dim=c, checkpoint_path=taehv_checkpoint)
|
||||
else:
|
||||
# torch TAEHV (regression / parity reference)
|
||||
import torch
|
||||
from fastvideo.third_party.taehv import TAEHV
|
||||
|
||||
ckpt = ensure_taehv_checkpoint(z_dim=c, checkpoint_path=taehv_checkpoint)
|
||||
if torch_device == "auto":
|
||||
torch_device = "mps" if torch.backends.mps.is_available() else "cpu"
|
||||
dtype = torch.float16 if torch_device == "mps" else torch.float32
|
||||
model = TAEHV(str(ckpt)).to(device=torch_device, dtype=dtype).eval()
|
||||
lat = torch.from_numpy(latents_np).to(device=torch_device, dtype=dtype)
|
||||
with torch.no_grad():
|
||||
out = model.decode_video(lat.transpose(1, 2), parallel=True, show_progress_bar=False)
|
||||
video = out[0].permute(0, 2, 3, 1).float().cpu().numpy()[None, ...]
|
||||
# out is NTCHW -> need BTHWC; decode_video returns NTCHW for batch
|
||||
if video.ndim == 4:
|
||||
video = video[None]
|
||||
elif backend == "wan-vae":
|
||||
if vae_dir is None:
|
||||
raise ValueError("vae_dir required for wan-vae backend")
|
||||
video = decode_latents_wan_vae_torch(latents_np, vae_dir=vae_dir, device=torch_device)
|
||||
else:
|
||||
raise ValueError(f"unknown backend {backend}")
|
||||
|
||||
decode_s = time.perf_counter() - t0
|
||||
output_path = Path(output_path)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
# export_to_video expects list/array of frames HxWxC
|
||||
frames = video[0]
|
||||
frames = np.clip(frames, 0.0, 1.0)
|
||||
export_to_video(frames, str(output_path), fps=fps)
|
||||
if backend == "taehv":
|
||||
cleanup_mlx()
|
||||
else:
|
||||
if backend == "taehv-torch":
|
||||
del model, lat, out
|
||||
cleanup_torch_mps()
|
||||
return {
|
||||
"decode_s": decode_s,
|
||||
"backend": backend,
|
||||
"output_path": str(output_path),
|
||||
"num_frames": int(frames.shape[0]),
|
||||
"resolution": f"{frames.shape[2]}x{frames.shape[1]}" if frames.ndim == 4 else None,
|
||||
}
|
||||
@@ -1,223 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Chunked non-causal sliding-window self-attention for MLX scaling studies.
|
||||
|
||||
This module is intentionally standalone (``mlx.core`` + stdlib only) so it can
|
||||
be micro-benchmarked without pulling in the DiT / FastVideo stack.
|
||||
|
||||
Window policy
|
||||
-------------
|
||||
**Symmetric** sliding window (non-causal). For query index ``i`` the allowed
|
||||
key indices are:
|
||||
|
||||
sinks: ``j in [0, sink)`` (always visible to every query, if ``sink > 0``)
|
||||
local: ``j in [max(0, i - half), min(S, i + half + 1))``
|
||||
where ``half = window // 2``
|
||||
|
||||
so each query sees roughly ``window + 1`` local keys (plus any sinks outside
|
||||
that range). This is appropriate for a dense, bidirectional DiT denoise pass.
|
||||
|
||||
Implementation note (FLOPs)
|
||||
---------------------------
|
||||
A full-size additive attention mask still materialises an ``O(S^2)`` score
|
||||
matrix inside SDPA and does **not** reduce work. Instead we tile the sequence
|
||||
into query blocks and run ``mx.fast.scaled_dot_product_attention`` only against
|
||||
the union of keys that block needs (local slice ± sinks). That makes
|
||||
per-block work ``O(chunk * (window + sink) * D)`` and total work
|
||||
``O(S * (window + sink) * D)``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
|
||||
def _default_scale(head_dim: int, scale: float | None) -> float:
|
||||
"""
|
||||
Determine the attention scaling factor from an explicit value or head dimension.
|
||||
|
||||
Parameters:
|
||||
head_dim (int): The attention head dimension used to derive the default scale.
|
||||
scale (Optional[float]): An explicit scaling factor.
|
||||
|
||||
Returns:
|
||||
float: The explicit scale converted to a float, or the reciprocal square root of `head_dim`.
|
||||
|
||||
Raises:
|
||||
ValueError: If `scale` is not provided and `head_dim` is not positive.
|
||||
"""
|
||||
if scale is not None:
|
||||
return float(scale)
|
||||
if head_dim <= 0:
|
||||
raise ValueError(f"head_dim must be positive, got {head_dim}")
|
||||
return head_dim**-0.5
|
||||
|
||||
|
||||
def _validate_qkv(q: mx.array, k: mx.array, v: mx.array) -> tuple[int, int, int, int]:
|
||||
"""
|
||||
Validate compatible rank-4 query, key, and value tensors.
|
||||
|
||||
Parameters:
|
||||
q (mx.array): Query tensor shaped `(B, H, S, D)`.
|
||||
k (mx.array): Key tensor with the same shape as `q`.
|
||||
v (mx.array): Value tensor with the same shape as `q`.
|
||||
|
||||
Returns:
|
||||
tuple[int, int, int, int]: Batch size, head count, sequence length, and head dimension.
|
||||
|
||||
Raises:
|
||||
ValueError: If the tensors are not rank 4, do not have identical shapes, or have an empty sequence or head dimension.
|
||||
"""
|
||||
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
|
||||
raise ValueError(f"q/k/v must be rank-4 (B, H, S, D); got shapes "
|
||||
f"{q.shape}, {k.shape}, {v.shape}")
|
||||
b, h, s, d = q.shape
|
||||
if k.shape != (b, h, s, d) or v.shape != (b, h, s, d):
|
||||
raise ValueError(f"q/k/v shapes must match exactly; got q={q.shape}, k={k.shape}, v={v.shape}")
|
||||
if s == 0:
|
||||
raise ValueError("sequence length S must be > 0")
|
||||
if d == 0:
|
||||
raise ValueError("head dim D must be > 0")
|
||||
return b, h, s, d
|
||||
|
||||
|
||||
def full_attention(
|
||||
q: mx.array,
|
||||
k: mx.array,
|
||||
v: mx.array,
|
||||
scale: float | None = None,
|
||||
) -> mx.array:
|
||||
"""
|
||||
Compute dense scaled dot-product attention over the full sequence.
|
||||
|
||||
Parameters:
|
||||
scale (float, optional): Attention scaling factor. If omitted, uses the
|
||||
inverse square root of the head dimension.
|
||||
|
||||
Returns:
|
||||
mx.array: Attention output with shape ``(B, H, S, D)``.
|
||||
"""
|
||||
_, _, _, d = _validate_qkv(q, k, v)
|
||||
sc = _default_scale(d, scale)
|
||||
return mx.fast.scaled_dot_product_attention(q, k, v, scale=sc)
|
||||
|
||||
|
||||
def _concat_kv_slices(
|
||||
k: mx.array,
|
||||
v: mx.array,
|
||||
ranges: list[tuple[int, int]],
|
||||
) -> tuple[mx.array, mx.array]:
|
||||
"""Concatenate non-overlapping ``[start, end)`` key/value slices along seq."""
|
||||
if not ranges:
|
||||
raise ValueError("ranges must be non-empty")
|
||||
if len(ranges) == 1:
|
||||
s0, e0 = ranges[0]
|
||||
return k[:, :, s0:e0, :], v[:, :, s0:e0, :]
|
||||
k_parts = [k[:, :, s:e, :] for s, e in ranges]
|
||||
v_parts = [v[:, :, s:e, :] for s, e in ranges]
|
||||
return mx.concatenate(k_parts, axis=2), mx.concatenate(v_parts, axis=2)
|
||||
|
||||
|
||||
def _key_ranges_for_block(
|
||||
qs: int,
|
||||
qe: int,
|
||||
seq_len: int,
|
||||
half: int,
|
||||
sink: int,
|
||||
) -> list[tuple[int, int]]:
|
||||
"""Return ordered, non-overlapping key ranges for a query block ``[qs, qe)``.
|
||||
|
||||
Symmetric local window over every query in the block, plus global sinks
|
||||
``[0, sink)``. Overlap is merged into a single contiguous range when
|
||||
possible so we avoid double-counting sink tokens.
|
||||
"""
|
||||
local_start = max(0, qs - half)
|
||||
# Last query index is ``qe - 1``; its right edge is ``qe - 1 + half + 1 = qe + half``.
|
||||
local_end = min(seq_len, qe + half)
|
||||
if local_start >= local_end:
|
||||
# Degenerate (should not happen for valid qs < qe); fall back to sinks only.
|
||||
if sink > 0:
|
||||
return [(0, min(sink, seq_len))]
|
||||
raise ValueError(f"empty local key range for query block [{qs}, {qe})")
|
||||
|
||||
if sink <= 0:
|
||||
return [(local_start, local_end)]
|
||||
|
||||
sink_end = min(sink, seq_len)
|
||||
if local_start <= sink_end:
|
||||
# Sinks abut or overlap the local window — one contiguous slice from 0.
|
||||
return [(0, max(local_end, sink_end))]
|
||||
# Gap between sinks and local window: two slices, concat at SDPA time.
|
||||
return [(0, sink_end), (local_start, local_end)]
|
||||
|
||||
|
||||
def windowed_attention(
|
||||
q: mx.array,
|
||||
k: mx.array,
|
||||
v: mx.array,
|
||||
window: int,
|
||||
sink: int = 0,
|
||||
scale: float | None = None,
|
||||
*,
|
||||
chunk_size: int | None = None,
|
||||
) -> mx.array:
|
||||
"""
|
||||
Apply symmetric sliding-window self-attention with optional global sink positions.
|
||||
|
||||
Parameters:
|
||||
q (mx.array): Query tensor shaped `(B, H, S, D)`.
|
||||
k (mx.array): Key tensor shaped `(B, H, S, D)`.
|
||||
v (mx.array): Value tensor shaped `(B, H, S, D)`.
|
||||
window (int): Symmetric attention window width in tokens; must be at least 1.
|
||||
sink (int): Number of leading key positions available to every query; must
|
||||
be between 0 and the sequence length.
|
||||
scale (Optional[float]): Softmax scale. Defaults to `1 / sqrt(D)`.
|
||||
chunk_size (Optional[int]): Query block length used for chunked processing.
|
||||
Defaults to the smaller of `window` and 512.
|
||||
|
||||
Returns:
|
||||
mx.array: Attention output with the same shape as `q`.
|
||||
|
||||
Raises:
|
||||
ValueError: If the inputs or attention parameters are invalid.
|
||||
RuntimeError: If a query block has no available keys.
|
||||
"""
|
||||
_, _, seq_len, d = _validate_qkv(q, k, v)
|
||||
|
||||
if window < 1:
|
||||
raise ValueError(f"window must be >= 1, got {window}")
|
||||
if sink < 0:
|
||||
raise ValueError(f"sink must be >= 0, got {sink}")
|
||||
if sink > seq_len:
|
||||
raise ValueError(f"sink ({sink}) cannot exceed sequence length ({seq_len})")
|
||||
|
||||
sc = _default_scale(d, scale)
|
||||
half = window // 2
|
||||
|
||||
# When the requested window is at least the sequence length, every query can
|
||||
# see every key under a symmetric policy — fall back to one dense SDPA.
|
||||
# (Sinks are redundant once the full key set is used.)
|
||||
if window >= seq_len:
|
||||
return mx.fast.scaled_dot_product_attention(q, k, v, scale=sc)
|
||||
|
||||
chunk = min(window, 512) if chunk_size is None else int(chunk_size)
|
||||
if chunk < 1:
|
||||
raise ValueError(f"chunk_size must be >= 1, got {chunk}")
|
||||
chunk = min(chunk, seq_len)
|
||||
|
||||
outputs: list[mx.array] = []
|
||||
for qs in range(0, seq_len, chunk):
|
||||
qe = min(seq_len, qs + chunk)
|
||||
q_block = q[:, :, qs:qe, :]
|
||||
ranges = _key_ranges_for_block(qs, qe, seq_len, half, sink)
|
||||
k_block, v_block = _concat_kv_slices(k, v, ranges)
|
||||
if k_block.shape[2] == 0:
|
||||
raise RuntimeError(f"empty key set for query block [{qs}, {qe}) with window={window}, sink={sink}")
|
||||
query_positions = mx.arange(qs, qe)[:, None]
|
||||
key_positions = mx.array([position for start, end in ranges for position in range(start, end)])[None, :]
|
||||
mask = mx.abs(query_positions - key_positions) <= half
|
||||
if sink > 0:
|
||||
mask = mask | (key_positions < sink)
|
||||
out_block = mx.fast.scaled_dot_product_attention(q_block, k_block, v_block, scale=sc, mask=mask)
|
||||
outputs.append(out_block)
|
||||
|
||||
return mx.concatenate(outputs, axis=2)
|
||||
@@ -1,376 +0,0 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
"""Reusable native BigVGAN-v2 vocoder.
|
||||
|
||||
Adapted from NVIDIA BigVGAN-v2 and its alias-free activation implementation.
|
||||
The CUDA activation kernel is intentionally excluded; FastVideo uses the
|
||||
portable PyTorch path for deterministic loading and parity.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.nn import Conv1d, ConvTranspose1d
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
from torch.nn.utils.parametrize import remove_parametrizations
|
||||
|
||||
|
||||
class AttrDict(dict):
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self.__dict__ = self
|
||||
|
||||
|
||||
def get_padding(kernel_size: int, dilation: int = 1) -> int:
|
||||
return int((kernel_size * dilation - dilation) / 2)
|
||||
|
||||
|
||||
def init_weights(module: nn.Module, mean: float = 0.0, std: float = 0.01) -> None:
|
||||
if "Conv" in module.__class__.__name__:
|
||||
module.weight.data.normal_(mean, std)
|
||||
|
||||
|
||||
def kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> torch.Tensor:
|
||||
even = kernel_size % 2 == 0
|
||||
half_size = kernel_size // 2
|
||||
delta_f = 4 * half_width
|
||||
amplitude = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
|
||||
if amplitude > 50.0:
|
||||
beta = 0.1102 * (amplitude - 8.7)
|
||||
elif amplitude >= 21.0:
|
||||
beta = 0.5842 * (amplitude - 21) ** 0.4 + 0.07886 * (amplitude - 21.0)
|
||||
else:
|
||||
beta = 0.0
|
||||
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
|
||||
time = torch.arange(-half_size, half_size) + 0.5 if even else torch.arange(kernel_size) - half_size
|
||||
if cutoff == 0:
|
||||
kernel = torch.zeros_like(time)
|
||||
else:
|
||||
kernel = 2 * cutoff * window * torch.sinc(2 * cutoff * time)
|
||||
kernel /= kernel.sum()
|
||||
return kernel.view(1, 1, kernel_size)
|
||||
|
||||
|
||||
class LowPassFilter1d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
cutoff: float = 0.5,
|
||||
half_width: float = 0.6,
|
||||
stride: int = 1,
|
||||
padding: bool = True,
|
||||
padding_mode: str = "replicate",
|
||||
kernel_size: int = 12,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.kernel_size = kernel_size
|
||||
self.even = kernel_size % 2 == 0
|
||||
self.pad_left = kernel_size // 2 - int(self.even)
|
||||
self.pad_right = kernel_size // 2
|
||||
self.stride = stride
|
||||
self.padding = padding
|
||||
self.padding_mode = padding_mode
|
||||
self.register_buffer("filter", kaiser_sinc_filter1d(cutoff, half_width, kernel_size))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
_, channels, _ = x.shape
|
||||
if self.padding:
|
||||
x = F.pad(x, (self.pad_left, self.pad_right), mode=self.padding_mode)
|
||||
return F.conv1d(x, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels)
|
||||
|
||||
|
||||
class UpSample1d(nn.Module):
|
||||
def __init__(self, ratio: int = 2, kernel_size: int | None = None) -> None:
|
||||
super().__init__()
|
||||
self.ratio = ratio
|
||||
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
|
||||
self.stride = ratio
|
||||
self.pad = self.kernel_size // ratio - 1
|
||||
self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
|
||||
self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
|
||||
self.register_buffer(
|
||||
"filter", kaiser_sinc_filter1d(cutoff=0.5 / ratio, half_width=0.6 / ratio, kernel_size=self.kernel_size)
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
_, channels, _ = x.shape
|
||||
x = F.pad(x, (self.pad, self.pad), mode="replicate")
|
||||
x = self.ratio * F.conv_transpose1d(
|
||||
x, self.filter.expand(channels, -1, -1), stride=self.stride, groups=channels
|
||||
)
|
||||
return x[..., self.pad_left : -self.pad_right]
|
||||
|
||||
|
||||
class DownSample1d(nn.Module):
|
||||
def __init__(self, ratio: int = 2, kernel_size: int | None = None) -> None:
|
||||
super().__init__()
|
||||
self.ratio = ratio
|
||||
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
|
||||
self.lowpass = LowPassFilter1d(
|
||||
cutoff=0.5 / ratio, half_width=0.6 / ratio, stride=ratio, kernel_size=self.kernel_size
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.lowpass(x)
|
||||
|
||||
|
||||
class Activation1d(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
activation: nn.Module,
|
||||
up_ratio: int = 2,
|
||||
down_ratio: int = 2,
|
||||
up_kernel_size: int = 12,
|
||||
down_kernel_size: int = 12,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.up_ratio = up_ratio
|
||||
self.down_ratio = down_ratio
|
||||
self.act = activation
|
||||
self.upsample = UpSample1d(up_ratio, up_kernel_size)
|
||||
self.downsample = DownSample1d(down_ratio, down_kernel_size)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.downsample(self.act(self.upsample(x)))
|
||||
|
||||
|
||||
class Snake(nn.Module):
|
||||
def __init__(
|
||||
self, in_features: int, alpha: float = 1.0, alpha_trainable: bool = True, alpha_logscale: bool = False
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.in_features = in_features
|
||||
self.alpha_logscale = alpha_logscale
|
||||
initial = torch.zeros(in_features) * alpha if alpha_logscale else torch.ones(in_features) * alpha
|
||||
self.alpha = nn.Parameter(initial, requires_grad=alpha_trainable)
|
||||
self.no_div_by_zero = 1e-9
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
|
||||
if self.alpha_logscale:
|
||||
alpha = torch.exp(alpha)
|
||||
return x + (1.0 / (alpha + self.no_div_by_zero)) * torch.pow(
|
||||
torch.sin(x * alpha), 2)
|
||||
|
||||
|
||||
class SnakeBeta(nn.Module):
|
||||
def __init__(
|
||||
self, in_features: int, alpha: float = 1.0, alpha_trainable: bool = True, alpha_logscale: bool = False
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.in_features = in_features
|
||||
self.alpha_logscale = alpha_logscale
|
||||
initial = torch.zeros(in_features) * alpha if alpha_logscale else torch.ones(in_features) * alpha
|
||||
self.alpha = nn.Parameter(initial.clone(), requires_grad=alpha_trainable)
|
||||
self.beta = nn.Parameter(initial.clone(), requires_grad=alpha_trainable)
|
||||
self.no_div_by_zero = 1e-9
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
alpha = self.alpha.unsqueeze(0).unsqueeze(-1)
|
||||
beta = self.beta.unsqueeze(0).unsqueeze(-1)
|
||||
if self.alpha_logscale:
|
||||
alpha = torch.exp(alpha)
|
||||
beta = torch.exp(beta)
|
||||
return x + (1.0 / (beta + self.no_div_by_zero)) * torch.pow(
|
||||
torch.sin(x * alpha), 2)
|
||||
|
||||
|
||||
def _activation(name: str, channels: int, logscale: bool) -> Activation1d:
|
||||
if name == "snake":
|
||||
activation = Snake(channels, alpha_logscale=logscale)
|
||||
elif name == "snakebeta":
|
||||
activation = SnakeBeta(channels, alpha_logscale=logscale)
|
||||
else:
|
||||
raise ValueError(f"Unsupported BigVGAN activation: {name}")
|
||||
return Activation1d(activation)
|
||||
|
||||
|
||||
class AMPBlock1(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: AttrDict,
|
||||
channels: int,
|
||||
kernel_size: int = 3,
|
||||
dilation: tuple[int, ...] = (1, 3, 5),
|
||||
activation: str = "snake",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.convs1 = nn.ModuleList(
|
||||
[
|
||||
weight_norm(
|
||||
Conv1d(
|
||||
channels, channels, kernel_size, stride=1, dilation=rate, padding=get_padding(kernel_size, rate)
|
||||
)
|
||||
)
|
||||
for rate in dilation
|
||||
]
|
||||
)
|
||||
self.convs1.apply(init_weights)
|
||||
self.convs2 = nn.ModuleList(
|
||||
[
|
||||
weight_norm(
|
||||
Conv1d(channels, channels, kernel_size, stride=1, dilation=1, padding=get_padding(kernel_size, 1))
|
||||
)
|
||||
for _ in dilation
|
||||
]
|
||||
)
|
||||
self.convs2.apply(init_weights)
|
||||
self.activations = nn.ModuleList(
|
||||
[
|
||||
_activation(activation, channels, config.snake_logscale)
|
||||
for _ in range(len(self.convs1) + len(self.convs2))
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
acts1, acts2 = self.activations[::2], self.activations[1::2]
|
||||
for conv1, conv2, act1, act2 in zip(self.convs1, self.convs2, acts1, acts2, strict=True):
|
||||
residual = conv2(act2(conv1(act1(x))))
|
||||
x = residual + x
|
||||
return x
|
||||
|
||||
def remove_weight_norm(self) -> None:
|
||||
for layer in self.convs1:
|
||||
remove_parametrizations(layer, "weight")
|
||||
for layer in self.convs2:
|
||||
remove_parametrizations(layer, "weight")
|
||||
|
||||
|
||||
class AMPBlock2(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config: AttrDict,
|
||||
channels: int,
|
||||
kernel_size: int = 3,
|
||||
dilation: tuple[int, ...] = (1, 3, 5),
|
||||
activation: str = "snake",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.convs = nn.ModuleList(
|
||||
[
|
||||
weight_norm(
|
||||
Conv1d(
|
||||
channels, channels, kernel_size, stride=1, dilation=rate, padding=get_padding(kernel_size, rate)
|
||||
)
|
||||
)
|
||||
for rate in dilation
|
||||
]
|
||||
)
|
||||
self.convs.apply(init_weights)
|
||||
self.activations = nn.ModuleList([_activation(activation, channels, config.snake_logscale) for _ in self.convs])
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
for conv, activation in zip(self.convs, self.activations, strict=True):
|
||||
x = conv(activation(x)) + x
|
||||
return x
|
||||
|
||||
def remove_weight_norm(self) -> None:
|
||||
for layer in self.convs:
|
||||
remove_parametrizations(layer, "weight")
|
||||
|
||||
|
||||
class BigVGANV2(nn.Module):
|
||||
"""BigVGAN-v2 generator compatible with NVIDIA checkpoint keys."""
|
||||
|
||||
def __init__(self, config: dict[str, Any]) -> None:
|
||||
super().__init__()
|
||||
config = dict(config)
|
||||
config.pop("_class_name", None)
|
||||
config.pop("architectures", None)
|
||||
weight_norm_removed = bool(config.pop("weight_norm_removed", False))
|
||||
self.config = AttrDict(config)
|
||||
if self.config.get("use_cuda_kernel", False):
|
||||
raise ValueError("FastVideo BigVGANV2 supports only the portable PyTorch path")
|
||||
self.config["use_cuda_kernel"] = False
|
||||
self.num_kernels = len(self.config.resblock_kernel_sizes)
|
||||
self.num_upsamples = len(self.config.upsample_rates)
|
||||
self.conv_pre = weight_norm(Conv1d(self.config.num_mels, self.config.upsample_initial_channel, 7, 1, padding=3))
|
||||
if self.config.resblock == "1":
|
||||
block_class = AMPBlock1
|
||||
elif self.config.resblock == "2":
|
||||
block_class = AMPBlock2
|
||||
else:
|
||||
raise ValueError(f"Unsupported BigVGAN resblock: {self.config.resblock}")
|
||||
|
||||
self.ups = nn.ModuleList()
|
||||
for index, (rate, kernel) in enumerate(
|
||||
zip(self.config.upsample_rates, self.config.upsample_kernel_sizes, strict=True)
|
||||
):
|
||||
self.ups.append(
|
||||
nn.ModuleList(
|
||||
[
|
||||
weight_norm(
|
||||
ConvTranspose1d(
|
||||
self.config.upsample_initial_channel // (2**index),
|
||||
self.config.upsample_initial_channel // (2 ** (index + 1)),
|
||||
kernel,
|
||||
rate,
|
||||
padding=(kernel - rate) // 2,
|
||||
)
|
||||
)
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
self.resblocks = nn.ModuleList()
|
||||
for index in range(len(self.ups)):
|
||||
channels = self.config.upsample_initial_channel // (2 ** (index + 1))
|
||||
for kernel, dilation in zip(
|
||||
self.config.resblock_kernel_sizes, self.config.resblock_dilation_sizes, strict=True
|
||||
):
|
||||
self.resblocks.append(
|
||||
block_class(self.config, channels, kernel, tuple(dilation), activation=self.config.activation)
|
||||
)
|
||||
|
||||
channels = self.config.upsample_initial_channel // (2 ** len(self.ups))
|
||||
self.activation_post = _activation(self.config.activation, channels, self.config.snake_logscale)
|
||||
self.use_bias_at_final = self.config.get("use_bias_at_final", True)
|
||||
self.conv_post = weight_norm(Conv1d(channels, 1, 7, 1, padding=3, bias=self.use_bias_at_final))
|
||||
for upsampler in self.ups:
|
||||
upsampler.apply(init_weights)
|
||||
self.conv_post.apply(init_weights)
|
||||
self.use_tanh_at_final = self.config.get("use_tanh_at_final", True)
|
||||
if weight_norm_removed:
|
||||
self.remove_weight_norm()
|
||||
|
||||
def forward(self, mel: torch.Tensor) -> torch.Tensor:
|
||||
hidden = self.conv_pre(mel)
|
||||
for index in range(self.num_upsamples):
|
||||
for upsampler in self.ups[index]:
|
||||
hidden = upsampler(hidden)
|
||||
accumulated = None
|
||||
for kernel in range(self.num_kernels):
|
||||
block_output = self.resblocks[
|
||||
index * self.num_kernels + kernel](hidden)
|
||||
if accumulated is None:
|
||||
accumulated = block_output
|
||||
else:
|
||||
# Preserve BigVGAN's published sequential accumulation
|
||||
# order. Tree reduction drifts through later nonlinear
|
||||
# upsampling stages with the full checkpoint.
|
||||
accumulated += block_output
|
||||
assert accumulated is not None
|
||||
hidden = accumulated / self.num_kernels
|
||||
hidden = self.conv_post(self.activation_post(hidden))
|
||||
if self.use_tanh_at_final:
|
||||
return torch.tanh(hidden)
|
||||
return torch.clamp(hidden, min=-1.0, max=1.0)
|
||||
|
||||
def remove_weight_norm(self) -> None:
|
||||
try:
|
||||
for upsamplers in self.ups:
|
||||
for upsampler in upsamplers:
|
||||
remove_parametrizations(upsampler, "weight")
|
||||
for block in self.resblocks:
|
||||
block.remove_weight_norm()
|
||||
remove_parametrizations(self.conv_pre, "weight")
|
||||
remove_parametrizations(self.conv_post, "weight")
|
||||
except ValueError:
|
||||
# Idempotent for pipeline setup and converted checkpoints.
|
||||
return
|
||||
|
||||
|
||||
EntryClass = BigVGANV2
|
||||
@@ -1,827 +0,0 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
#
|
||||
# MIT License
|
||||
#
|
||||
# Copyright (c) 2024 Sony Research Inc.
|
||||
#
|
||||
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
# of this software and associated documentation files (the "Software"), to deal
|
||||
# in the Software without restriction, including without limitation the rights
|
||||
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
# copies of the Software, and to permit persons to whom the Software is
|
||||
# furnished to do so, subject to the following conditions:
|
||||
#
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
#
|
||||
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
# SOFTWARE.
|
||||
"""Native 1D audio VAE used by MMAudio."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from collections.abc import Iterable
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.models.loader.weight_utils import default_weight_loader
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
DATA_MEAN_80D = [
|
||||
-1.6058,
|
||||
-1.3676,
|
||||
-1.2520,
|
||||
-1.2453,
|
||||
-1.2078,
|
||||
-1.2224,
|
||||
-1.2419,
|
||||
-1.2439,
|
||||
-1.2922,
|
||||
-1.2927,
|
||||
-1.3170,
|
||||
-1.3543,
|
||||
-1.3401,
|
||||
-1.3836,
|
||||
-1.3907,
|
||||
-1.3912,
|
||||
-1.4313,
|
||||
-1.4152,
|
||||
-1.4527,
|
||||
-1.4728,
|
||||
-1.4568,
|
||||
-1.5101,
|
||||
-1.5051,
|
||||
-1.5172,
|
||||
-1.5623,
|
||||
-1.5373,
|
||||
-1.5746,
|
||||
-1.5687,
|
||||
-1.6032,
|
||||
-1.6131,
|
||||
-1.6081,
|
||||
-1.6331,
|
||||
-1.6489,
|
||||
-1.6489,
|
||||
-1.6700,
|
||||
-1.6738,
|
||||
-1.6953,
|
||||
-1.6969,
|
||||
-1.7048,
|
||||
-1.7280,
|
||||
-1.7361,
|
||||
-1.7495,
|
||||
-1.7658,
|
||||
-1.7814,
|
||||
-1.7889,
|
||||
-1.8064,
|
||||
-1.8221,
|
||||
-1.8377,
|
||||
-1.8417,
|
||||
-1.8643,
|
||||
-1.8857,
|
||||
-1.8929,
|
||||
-1.9173,
|
||||
-1.9379,
|
||||
-1.9531,
|
||||
-1.9673,
|
||||
-1.9824,
|
||||
-2.0042,
|
||||
-2.0215,
|
||||
-2.0436,
|
||||
-2.0766,
|
||||
-2.1064,
|
||||
-2.1418,
|
||||
-2.1855,
|
||||
-2.2319,
|
||||
-2.2767,
|
||||
-2.3161,
|
||||
-2.3572,
|
||||
-2.3954,
|
||||
-2.4282,
|
||||
-2.4659,
|
||||
-2.5072,
|
||||
-2.5552,
|
||||
-2.6074,
|
||||
-2.6584,
|
||||
-2.7107,
|
||||
-2.7634,
|
||||
-2.8266,
|
||||
-2.8981,
|
||||
-2.9673,
|
||||
]
|
||||
|
||||
DATA_STD_80D = [
|
||||
1.0291,
|
||||
1.0411,
|
||||
1.0043,
|
||||
0.9820,
|
||||
0.9677,
|
||||
0.9543,
|
||||
0.9450,
|
||||
0.9392,
|
||||
0.9343,
|
||||
0.9297,
|
||||
0.9276,
|
||||
0.9263,
|
||||
0.9242,
|
||||
0.9254,
|
||||
0.9232,
|
||||
0.9281,
|
||||
0.9263,
|
||||
0.9315,
|
||||
0.9274,
|
||||
0.9247,
|
||||
0.9277,
|
||||
0.9199,
|
||||
0.9188,
|
||||
0.9194,
|
||||
0.9160,
|
||||
0.9161,
|
||||
0.9146,
|
||||
0.9161,
|
||||
0.9100,
|
||||
0.9095,
|
||||
0.9145,
|
||||
0.9076,
|
||||
0.9066,
|
||||
0.9095,
|
||||
0.9032,
|
||||
0.9043,
|
||||
0.9038,
|
||||
0.9011,
|
||||
0.9019,
|
||||
0.9010,
|
||||
0.8984,
|
||||
0.8983,
|
||||
0.8986,
|
||||
0.8961,
|
||||
0.8962,
|
||||
0.8978,
|
||||
0.8962,
|
||||
0.8973,
|
||||
0.8993,
|
||||
0.8976,
|
||||
0.8995,
|
||||
0.9016,
|
||||
0.8982,
|
||||
0.8972,
|
||||
0.8974,
|
||||
0.8949,
|
||||
0.8940,
|
||||
0.8947,
|
||||
0.8936,
|
||||
0.8939,
|
||||
0.8951,
|
||||
0.8956,
|
||||
0.9017,
|
||||
0.9167,
|
||||
0.9436,
|
||||
0.9690,
|
||||
1.0003,
|
||||
1.0225,
|
||||
1.0381,
|
||||
1.0491,
|
||||
1.0545,
|
||||
1.0604,
|
||||
1.0761,
|
||||
1.0929,
|
||||
1.1089,
|
||||
1.1196,
|
||||
1.1176,
|
||||
1.1156,
|
||||
1.1117,
|
||||
1.1070,
|
||||
]
|
||||
|
||||
DATA_MEAN_128D = [
|
||||
-3.3462,
|
||||
-2.6723,
|
||||
-2.4893,
|
||||
-2.3143,
|
||||
-2.2664,
|
||||
-2.3317,
|
||||
-2.1802,
|
||||
-2.4006,
|
||||
-2.2357,
|
||||
-2.4597,
|
||||
-2.3717,
|
||||
-2.4690,
|
||||
-2.5142,
|
||||
-2.4919,
|
||||
-2.6610,
|
||||
-2.5047,
|
||||
-2.7483,
|
||||
-2.5926,
|
||||
-2.7462,
|
||||
-2.7033,
|
||||
-2.7386,
|
||||
-2.8112,
|
||||
-2.7502,
|
||||
-2.9594,
|
||||
-2.7473,
|
||||
-3.0035,
|
||||
-2.8891,
|
||||
-2.9922,
|
||||
-2.9856,
|
||||
-3.0157,
|
||||
-3.1191,
|
||||
-2.9893,
|
||||
-3.1718,
|
||||
-3.0745,
|
||||
-3.1879,
|
||||
-3.2310,
|
||||
-3.1424,
|
||||
-3.2296,
|
||||
-3.2791,
|
||||
-3.2782,
|
||||
-3.2756,
|
||||
-3.3134,
|
||||
-3.3509,
|
||||
-3.3750,
|
||||
-3.3951,
|
||||
-3.3698,
|
||||
-3.4505,
|
||||
-3.4509,
|
||||
-3.5089,
|
||||
-3.4647,
|
||||
-3.5536,
|
||||
-3.5788,
|
||||
-3.5867,
|
||||
-3.6036,
|
||||
-3.6400,
|
||||
-3.6747,
|
||||
-3.7072,
|
||||
-3.7279,
|
||||
-3.7283,
|
||||
-3.7795,
|
||||
-3.8259,
|
||||
-3.8447,
|
||||
-3.8663,
|
||||
-3.9182,
|
||||
-3.9605,
|
||||
-3.9861,
|
||||
-4.0105,
|
||||
-4.0373,
|
||||
-4.0762,
|
||||
-4.1121,
|
||||
-4.1488,
|
||||
-4.1874,
|
||||
-4.2461,
|
||||
-4.3170,
|
||||
-4.3639,
|
||||
-4.4452,
|
||||
-4.5282,
|
||||
-4.6297,
|
||||
-4.7019,
|
||||
-4.7960,
|
||||
-4.8700,
|
||||
-4.9507,
|
||||
-5.0303,
|
||||
-5.0866,
|
||||
-5.1634,
|
||||
-5.2342,
|
||||
-5.3242,
|
||||
-5.4053,
|
||||
-5.4927,
|
||||
-5.5712,
|
||||
-5.6464,
|
||||
-5.7052,
|
||||
-5.7619,
|
||||
-5.8410,
|
||||
-5.9188,
|
||||
-6.0103,
|
||||
-6.0955,
|
||||
-6.1673,
|
||||
-6.2362,
|
||||
-6.3120,
|
||||
-6.3926,
|
||||
-6.4797,
|
||||
-6.5565,
|
||||
-6.6511,
|
||||
-6.8130,
|
||||
-6.9961,
|
||||
-7.1275,
|
||||
-7.2457,
|
||||
-7.3576,
|
||||
-7.4663,
|
||||
-7.6136,
|
||||
-7.7469,
|
||||
-7.8815,
|
||||
-8.0132,
|
||||
-8.1515,
|
||||
-8.3071,
|
||||
-8.4722,
|
||||
-8.7418,
|
||||
-9.3975,
|
||||
-9.6628,
|
||||
-9.7671,
|
||||
-9.8863,
|
||||
-9.9992,
|
||||
-10.0860,
|
||||
-10.1709,
|
||||
-10.5418,
|
||||
-11.2795,
|
||||
-11.3861,
|
||||
]
|
||||
|
||||
DATA_STD_128D = [
|
||||
2.3804,
|
||||
2.4368,
|
||||
2.3772,
|
||||
2.3145,
|
||||
2.2803,
|
||||
2.2510,
|
||||
2.2316,
|
||||
2.2083,
|
||||
2.1996,
|
||||
2.1835,
|
||||
2.1769,
|
||||
2.1659,
|
||||
2.1631,
|
||||
2.1618,
|
||||
2.1540,
|
||||
2.1606,
|
||||
2.1571,
|
||||
2.1567,
|
||||
2.1612,
|
||||
2.1579,
|
||||
2.1679,
|
||||
2.1683,
|
||||
2.1634,
|
||||
2.1557,
|
||||
2.1668,
|
||||
2.1518,
|
||||
2.1415,
|
||||
2.1449,
|
||||
2.1406,
|
||||
2.1350,
|
||||
2.1313,
|
||||
2.1415,
|
||||
2.1281,
|
||||
2.1352,
|
||||
2.1219,
|
||||
2.1182,
|
||||
2.1327,
|
||||
2.1195,
|
||||
2.1137,
|
||||
2.1080,
|
||||
2.1179,
|
||||
2.1036,
|
||||
2.1087,
|
||||
2.1036,
|
||||
2.1015,
|
||||
2.1068,
|
||||
2.0975,
|
||||
2.0991,
|
||||
2.0902,
|
||||
2.1015,
|
||||
2.0857,
|
||||
2.0920,
|
||||
2.0893,
|
||||
2.0897,
|
||||
2.0910,
|
||||
2.0881,
|
||||
2.0925,
|
||||
2.0873,
|
||||
2.0960,
|
||||
2.0900,
|
||||
2.0957,
|
||||
2.0958,
|
||||
2.0978,
|
||||
2.0936,
|
||||
2.0886,
|
||||
2.0905,
|
||||
2.0845,
|
||||
2.0855,
|
||||
2.0796,
|
||||
2.0840,
|
||||
2.0813,
|
||||
2.0817,
|
||||
2.0838,
|
||||
2.0840,
|
||||
2.0917,
|
||||
2.1061,
|
||||
2.1431,
|
||||
2.1976,
|
||||
2.2482,
|
||||
2.3055,
|
||||
2.3700,
|
||||
2.4088,
|
||||
2.4372,
|
||||
2.4609,
|
||||
2.4731,
|
||||
2.4847,
|
||||
2.5072,
|
||||
2.5451,
|
||||
2.5772,
|
||||
2.6147,
|
||||
2.6529,
|
||||
2.6596,
|
||||
2.6645,
|
||||
2.6726,
|
||||
2.6803,
|
||||
2.6812,
|
||||
2.6899,
|
||||
2.6916,
|
||||
2.6931,
|
||||
2.6998,
|
||||
2.7062,
|
||||
2.7262,
|
||||
2.7222,
|
||||
2.7158,
|
||||
2.7041,
|
||||
2.7485,
|
||||
2.7491,
|
||||
2.7451,
|
||||
2.7485,
|
||||
2.7233,
|
||||
2.7297,
|
||||
2.7233,
|
||||
2.7145,
|
||||
2.6958,
|
||||
2.6788,
|
||||
2.6439,
|
||||
2.6007,
|
||||
2.4786,
|
||||
2.2469,
|
||||
2.1877,
|
||||
2.1392,
|
||||
2.0717,
|
||||
2.0107,
|
||||
1.9676,
|
||||
1.9140,
|
||||
1.7102,
|
||||
0.9101,
|
||||
0.7164,
|
||||
]
|
||||
|
||||
|
||||
_NORMALIZATION_EPSILON = 1e-4
|
||||
_SILU_DIVISOR = 0.596
|
||||
_RESIDUAL_BLEND = 0.3
|
||||
_RESIDUAL_DIVISOR = math.hypot(1.0 - _RESIDUAL_BLEND, _RESIDUAL_BLEND)
|
||||
|
||||
|
||||
def _rms_normalize(x: torch.Tensor, dim: int | tuple[int, ...]) -> torch.Tensor:
|
||||
dims = (dim,) if isinstance(dim, int) else dim
|
||||
element_count = math.prod(x.shape[axis] for axis in dims)
|
||||
l2_norm = torch.linalg.vector_norm(x, dim=dims, keepdim=True, dtype=torch.float32)
|
||||
rms = torch.add(_NORMALIZATION_EPSILON, l2_norm, alpha=element_count**-0.5)
|
||||
return x / rms.to(x.dtype)
|
||||
|
||||
|
||||
def _conv1d(in_channels: int, out_channels: int, kernel_size: int) -> nn.Conv1d:
|
||||
return nn.Conv1d(in_channels, out_channels, kernel_size, padding=kernel_size // 2, bias=False)
|
||||
|
||||
|
||||
def _conv1d_with_gain(conv: nn.Conv1d, x: torch.Tensor, gain: torch.Tensor | float) -> torch.Tensor:
|
||||
return F.conv1d(x, conv.weight * gain, stride=conv.stride, padding=conv.padding, dilation=conv.dilation,
|
||||
groups=conv.groups)
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution:
|
||||
def __init__(self, parameters: torch.Tensor, deterministic: bool = False) -> None:
|
||||
self.parameters = parameters
|
||||
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
|
||||
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
||||
self.deterministic = deterministic
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if deterministic:
|
||||
self.var = self.std = torch.zeros_like(self.mean)
|
||||
|
||||
def sample(self, generator: torch.Generator | None = None) -> torch.Tensor:
|
||||
noise = torch.empty_like(self.mean).normal_(generator=generator)
|
||||
return self.mean + self.std * noise
|
||||
|
||||
def mode(self) -> torch.Tensor:
|
||||
return self.mean
|
||||
|
||||
|
||||
class ResnetBlock1D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
in_dim: int,
|
||||
out_dim: int | None = None,
|
||||
conv_shortcut: bool = False,
|
||||
kernel_size: int = 3,
|
||||
use_norm: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = in_dim if out_dim is None else out_dim
|
||||
self.use_conv_shortcut = conv_shortcut
|
||||
self.use_norm = use_norm
|
||||
self.conv1 = _conv1d(in_dim, self.out_dim, kernel_size)
|
||||
self.conv2 = _conv1d(self.out_dim, self.out_dim, kernel_size)
|
||||
if self.in_dim != self.out_dim:
|
||||
if conv_shortcut:
|
||||
self.conv_shortcut = _conv1d(in_dim, self.out_dim, kernel_size)
|
||||
else:
|
||||
self.nin_shortcut = _conv1d(in_dim, self.out_dim, 1)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_norm:
|
||||
x = _rms_normalize(x, dim=1)
|
||||
hidden = self.conv1(F.silu(x) / _SILU_DIVISOR)
|
||||
hidden = self.conv2(F.silu(hidden) / _SILU_DIVISOR)
|
||||
if self.in_dim != self.out_dim:
|
||||
shortcut = self.conv_shortcut if self.use_conv_shortcut else self.nin_shortcut
|
||||
x = shortcut(x)
|
||||
return torch.lerp(x, hidden, _RESIDUAL_BLEND) / _RESIDUAL_DIVISOR
|
||||
|
||||
|
||||
class AttnBlock1D(nn.Module):
|
||||
def __init__(self, in_channels: int, num_heads: int = 1) -> None:
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.num_heads = num_heads
|
||||
self.qkv = _conv1d(in_channels, in_channels * 3, 1)
|
||||
self.proj_out = _conv1d(in_channels, in_channels, 1)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
qkv = self.qkv(x).reshape(x.shape[0], self.num_heads, -1, 3, x.shape[-1])
|
||||
query, key, value = _rms_normalize(qkv, dim=2).unbind(3)
|
||||
query = rearrange(query, "b h c l -> b h l c")
|
||||
key = rearrange(key, "b h c l -> b h l c")
|
||||
value = rearrange(value, "b h c l -> b h l c")
|
||||
hidden = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden = rearrange(hidden, "b h l c -> b (h c) l")
|
||||
return torch.lerp(x, self.proj_out(hidden), _RESIDUAL_BLEND) / _RESIDUAL_DIVISOR
|
||||
|
||||
|
||||
class Upsample1D(nn.Module):
|
||||
def __init__(self, in_channels: int, with_conv: bool) -> None:
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if with_conv:
|
||||
self.conv = _conv1d(in_channels, in_channels, 3)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = F.interpolate(x, scale_factor=2.0, mode="nearest-exact")
|
||||
return self.conv(x) if self.with_conv else x
|
||||
|
||||
|
||||
class Downsample1D(nn.Module):
|
||||
def __init__(self, in_channels: int, with_conv: bool) -> None:
|
||||
super().__init__()
|
||||
self.with_conv = with_conv
|
||||
if with_conv:
|
||||
self.conv1 = _conv1d(in_channels, in_channels, 1)
|
||||
self.conv2 = _conv1d(in_channels, in_channels, 1)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.with_conv:
|
||||
x = self.conv1(x)
|
||||
x = F.avg_pool1d(x, kernel_size=2, stride=2)
|
||||
return self.conv2(x) if self.with_conv else x
|
||||
|
||||
|
||||
class Encoder1D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dim: int,
|
||||
ch_mult: tuple[int, ...],
|
||||
num_res_blocks: int,
|
||||
attn_layers: list[int],
|
||||
down_layers: list[int],
|
||||
in_dim: int,
|
||||
embed_dim: int,
|
||||
resamp_with_conv: bool = True,
|
||||
double_z: bool = True,
|
||||
kernel_size: int = 3,
|
||||
clip_act: float = 256.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_layers = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.in_channels = in_dim
|
||||
self.clip_act = clip_act
|
||||
self.down_layers = down_layers
|
||||
self.attn_layers = attn_layers
|
||||
self.conv_in = _conv1d(in_dim, dim, kernel_size)
|
||||
|
||||
in_ch_mult = (1,) + ch_mult
|
||||
self.in_ch_mult = in_ch_mult
|
||||
self.down = nn.ModuleList()
|
||||
for level in range(self.num_layers):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_in = dim * in_ch_mult[level]
|
||||
block_out = dim * ch_mult[level]
|
||||
for _ in range(num_res_blocks):
|
||||
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, kernel_size=kernel_size, use_norm=True))
|
||||
block_in = block_out
|
||||
if level in attn_layers:
|
||||
attn.append(AttnBlock1D(block_in))
|
||||
down = nn.Module()
|
||||
down.block = block
|
||||
down.attn = attn
|
||||
if level in down_layers:
|
||||
down.downsample = Downsample1D(block_in, resamp_with_conv)
|
||||
self.down.append(down)
|
||||
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, kernel_size=kernel_size, use_norm=True)
|
||||
self.mid.attn_1 = AttnBlock1D(block_in)
|
||||
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, kernel_size=kernel_size, use_norm=True)
|
||||
output_dim = 2 * embed_dim if double_z else embed_dim
|
||||
self.conv_out = _conv1d(block_in, output_dim, kernel_size)
|
||||
self.learnable_gain = nn.Parameter(torch.zeros([]))
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
states = [self.conv_in(x)]
|
||||
for level in range(self.num_layers):
|
||||
for block_index in range(self.num_res_blocks):
|
||||
hidden = self.down[level].block[block_index](states[-1])
|
||||
if len(self.down[level].attn) > 0:
|
||||
hidden = self.down[level].attn[block_index](hidden)
|
||||
states.append(hidden.clamp(-self.clip_act, self.clip_act))
|
||||
if level in self.down_layers:
|
||||
states.append(self.down[level].downsample(states[-1]))
|
||||
hidden = self.mid.block_1(states[-1])
|
||||
hidden = self.mid.attn_1(hidden)
|
||||
hidden = self.mid.block_2(hidden).clamp(-self.clip_act, self.clip_act)
|
||||
return _conv1d_with_gain(self.conv_out, F.silu(hidden) / _SILU_DIVISOR, self.learnable_gain + 1)
|
||||
|
||||
|
||||
class Decoder1D(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dim: int,
|
||||
out_dim: int,
|
||||
ch_mult: tuple[int, ...],
|
||||
num_res_blocks: int,
|
||||
attn_layers: list[int],
|
||||
down_layers: list[int],
|
||||
in_dim: int,
|
||||
embed_dim: int,
|
||||
kernel_size: int = 3,
|
||||
resamp_with_conv: bool = True,
|
||||
clip_act: float = 256.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.ch = dim
|
||||
self.num_layers = len(ch_mult)
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.in_channels = in_dim
|
||||
self.clip_act = clip_act
|
||||
self.down_layers = [level + 1 for level in down_layers]
|
||||
block_in = dim * ch_mult[-1]
|
||||
self.conv_in = _conv1d(embed_dim, block_in, kernel_size)
|
||||
self.mid = nn.Module()
|
||||
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
|
||||
self.mid.attn_1 = AttnBlock1D(block_in)
|
||||
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
|
||||
|
||||
self.up = nn.ModuleList()
|
||||
for level in reversed(range(self.num_layers)):
|
||||
block = nn.ModuleList()
|
||||
attn = nn.ModuleList()
|
||||
block_out = dim * ch_mult[level]
|
||||
for _ in range(num_res_blocks + 1):
|
||||
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, use_norm=True))
|
||||
block_in = block_out
|
||||
if level in attn_layers:
|
||||
attn.append(AttnBlock1D(block_in))
|
||||
up = nn.Module()
|
||||
up.block = block
|
||||
up.attn = attn
|
||||
if level in self.down_layers:
|
||||
up.upsample = Upsample1D(block_in, resamp_with_conv)
|
||||
self.up.insert(0, up)
|
||||
|
||||
self.conv_out = _conv1d(block_in, out_dim, kernel_size)
|
||||
self.learnable_gain = nn.Parameter(torch.zeros([]))
|
||||
|
||||
def forward(self, z: torch.Tensor) -> torch.Tensor:
|
||||
hidden = self.conv_in(z)
|
||||
hidden = self.mid.block_1(hidden)
|
||||
hidden = self.mid.attn_1(hidden)
|
||||
hidden = self.mid.block_2(hidden).clamp(-self.clip_act, self.clip_act)
|
||||
for level in reversed(range(self.num_layers)):
|
||||
for block_index in range(self.num_res_blocks + 1):
|
||||
hidden = self.up[level].block[block_index](hidden)
|
||||
if len(self.up[level].attn) > 0:
|
||||
hidden = self.up[level].attn[block_index](hidden)
|
||||
hidden = hidden.clamp(-self.clip_act, self.clip_act)
|
||||
if level in self.down_layers:
|
||||
hidden = self.up[level].upsample(hidden)
|
||||
return _conv1d_with_gain(self.conv_out, F.silu(hidden) / _SILU_DIVISOR, self.learnable_gain + 1)
|
||||
|
||||
|
||||
class MMAudioVAE(nn.Module):
|
||||
"""MMAudio mel-spectrogram VAE for 16 kHz or 44.1 kHz audio."""
|
||||
|
||||
def __init__(self, mode: str | dict[str, Any] = "44k", need_encoder: bool = False) -> None:
|
||||
super().__init__()
|
||||
if isinstance(mode, dict):
|
||||
config = mode
|
||||
mode = config.get("mode", "44k")
|
||||
need_encoder = config.get("need_encoder", need_encoder)
|
||||
if mode == "16k":
|
||||
data_dim, embed_dim, hidden_dim = 80, 20, 384
|
||||
data_mean, data_std = DATA_MEAN_80D, DATA_STD_80D
|
||||
elif mode == "44k":
|
||||
data_dim, embed_dim, hidden_dim = 128, 40, 512
|
||||
data_mean, data_std = DATA_MEAN_128D, DATA_STD_128D
|
||||
else:
|
||||
raise ValueError(f"Unknown MMAudio VAE mode: {mode}")
|
||||
|
||||
self.mode = mode
|
||||
self.embed_dim = embed_dim
|
||||
self._weights_normalized = False
|
||||
self.register_buffer("data_mean", torch.tensor(data_mean, dtype=torch.float32).view(1, -1, 1))
|
||||
self.register_buffer("data_std", torch.tensor(data_std, dtype=torch.float32).view(1, -1, 1))
|
||||
if need_encoder:
|
||||
self.encoder = Encoder1D(
|
||||
dim=hidden_dim,
|
||||
ch_mult=(1, 2, 4),
|
||||
num_res_blocks=2,
|
||||
attn_layers=[3],
|
||||
down_layers=[0],
|
||||
in_dim=data_dim,
|
||||
embed_dim=embed_dim,
|
||||
)
|
||||
self.decoder = Decoder1D(
|
||||
dim=hidden_dim,
|
||||
ch_mult=(1, 2, 4),
|
||||
num_res_blocks=2,
|
||||
attn_layers=[3],
|
||||
down_layers=[0],
|
||||
in_dim=data_dim,
|
||||
out_dim=data_dim,
|
||||
embed_dim=embed_dim,
|
||||
)
|
||||
|
||||
def encode(self, mel: torch.Tensor, normalize_input: bool = True) -> DiagonalGaussianDistribution:
|
||||
self._require_normalized_weights()
|
||||
if not hasattr(self, "encoder"):
|
||||
raise RuntimeError("This MMAudio VAE was loaded decoder-only")
|
||||
if normalize_input:
|
||||
mel = self.normalize(mel)
|
||||
return DiagonalGaussianDistribution(self.encoder(mel))
|
||||
|
||||
def decode(self, latent: torch.Tensor, unnormalize_output: bool = True) -> torch.Tensor:
|
||||
self._require_normalized_weights()
|
||||
mel = self.decoder(latent)
|
||||
return self.unnormalize(mel) if unnormalize_output else mel
|
||||
|
||||
def forward(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
return self.decode(latent)
|
||||
|
||||
def normalize(self, mel: torch.Tensor) -> torch.Tensor:
|
||||
return (mel - self.data_mean) / self.data_std
|
||||
|
||||
def unnormalize(self, mel: torch.Tensor) -> torch.Tensor:
|
||||
return mel * self.data_std + self.data_mean
|
||||
|
||||
def _require_normalized_weights(self) -> None:
|
||||
if not self._weights_normalized:
|
||||
raise RuntimeError("call remove_weight_norm() before inference")
|
||||
|
||||
@torch.no_grad()
|
||||
def remove_weight_norm(self):
|
||||
for name, module in self.named_modules():
|
||||
if isinstance(module, nn.Conv1d):
|
||||
weight = _rms_normalize(module.weight.to(torch.float32), dim=(1, 2))
|
||||
weight = weight / math.sqrt(weight[0].numel())
|
||||
module.weight.copy_(weight.to(module.weight.dtype))
|
||||
logger.debug("Removed weight norm from %s", name)
|
||||
self._weights_normalized = True
|
||||
return self
|
||||
|
||||
def load_weights(
|
||||
self,
|
||||
weights: Iterable[tuple[str, torch.Tensor]],
|
||||
) -> set[str]:
|
||||
params = dict(self.named_parameters())
|
||||
loaded: set[str] = set()
|
||||
for name, tensor in weights:
|
||||
if name not in params:
|
||||
continue
|
||||
parameter = params[name]
|
||||
loader = getattr(parameter, "weight_loader", default_weight_loader)
|
||||
loader(parameter, tensor)
|
||||
loaded.add(name)
|
||||
return loaded
|
||||
|
||||
|
||||
EntryClass = MMAudioVAE
|
||||
@@ -10,7 +10,6 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention import DistributedAttention
|
||||
from fastvideo.attention.layer import DistributedAttention_VSA
|
||||
from fastvideo.attention.selector import get_attn_backend
|
||||
@@ -24,50 +23,12 @@ from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.minimax_h3_fusions import (
|
||||
HAVE_TRITON,
|
||||
fused_qknorm_rope,
|
||||
fused_residual_gate_rmsnorm_modulate,
|
||||
fused_rmsnorm_modulate,
|
||||
minimax_h3_swiglu,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.profiler import nvtx_range
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
MINIMAX_H3_MODALITY_NUM = 3
|
||||
_CFG = MiniMaxH3Config()
|
||||
_MINIMAX_H3_FUSION_NAMES = frozenset({"modulate", "qknorm_rope", "swiglu"})
|
||||
|
||||
|
||||
def _enabled_minimax_h3_fusions(value: str | None = None) -> frozenset[str]:
|
||||
"""Parse the independently switchable inference fusion set."""
|
||||
raw = envs.FASTVIDEO_MINIMAX_H3_FUSIONS if value is None else value
|
||||
normalized = raw.strip().lower()
|
||||
if normalized in {"", "0", "none"}:
|
||||
return frozenset()
|
||||
if normalized in {"1", "all"}:
|
||||
return _MINIMAX_H3_FUSION_NAMES
|
||||
enabled = frozenset(item.strip() for item in normalized.split(",") if item.strip())
|
||||
unknown = enabled - _MINIMAX_H3_FUSION_NAMES
|
||||
if unknown:
|
||||
supported = ",".join(sorted(_MINIMAX_H3_FUSION_NAMES))
|
||||
raise ValueError(f"Unknown MiniMax H3 fusion(s) {sorted(unknown)}; expected a subset of {supported}.")
|
||||
return enabled
|
||||
|
||||
|
||||
def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool:
|
||||
"""Triton kernels are inference-only and stay outside Dynamo capture.
|
||||
|
||||
The ``HAVE_TRITON`` check makes the eager fallback exact: on a CUDA build
|
||||
whose Triton failed to import, an enabled fusion falls back instead of
|
||||
hitting the strict wrappers' hard RuntimeError mid-forward.
|
||||
"""
|
||||
return (HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() and not torch.compiler.is_compiling())
|
||||
|
||||
|
||||
class MiniMaxH3RotaryPosEmbed(nn.Module):
|
||||
@@ -101,7 +62,6 @@ class MiniMaxH3FeedForward(nn.Module):
|
||||
ffn_dim: int,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
fuse_swiglu: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.fc_in = ReplicatedLinear(
|
||||
@@ -118,15 +78,11 @@ class MiniMaxH3FeedForward(nn.Module):
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.fc_out",
|
||||
)
|
||||
self.fuse_swiglu = fuse_swiglu
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states, _ = self.fc_in(hidden_states)
|
||||
if self.fuse_swiglu and _can_run_minimax_h3_fusion(hidden_states):
|
||||
hidden_states = minimax_h3_swiglu(hidden_states)
|
||||
else:
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
hidden_states = hidden_states * F.silu(gate)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
hidden_states = hidden_states * F.silu(gate)
|
||||
hidden_states, _ = self.fc_out(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
@@ -143,7 +99,6 @@ class MiniMaxH3Attention(nn.Module):
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...],
|
||||
quant_config: QuantizationConfig | None,
|
||||
prefix: str,
|
||||
fuse_qknorm_rope: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
@@ -179,7 +134,6 @@ class MiniMaxH3Attention(nn.Module):
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.to_out",
|
||||
)
|
||||
self.fuse_qknorm_rope = fuse_qknorm_rope
|
||||
# VSA carries a learned gate on its pooled-compression branch. The H3
|
||||
# checkpoint has no such weight, so the loader zero-initializes it
|
||||
# (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
|
||||
@@ -257,18 +211,11 @@ class MiniMaxH3Attention(nn.Module):
|
||||
query = query.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
|
||||
key = key.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
|
||||
value = value.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
|
||||
if (self.fuse_qknorm_rope and rotary_emb is not None and _can_run_minimax_h3_fusion(query)):
|
||||
cos, sin = rotary_emb
|
||||
cos = cos.to(query.dtype)
|
||||
sin = sin.to(query.dtype)
|
||||
query = fused_qknorm_rope(query, self.norm_q.weight, cos, sin, self.norm_q.eps)
|
||||
key = fused_qknorm_rope(key, self.norm_k.weight, cos, sin, self.norm_k.eps)
|
||||
else:
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
if rotary_emb is not None:
|
||||
query = self._apply_rotary_emb(query, rotary_emb)
|
||||
key = self._apply_rotary_emb(key, rotary_emb)
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
if rotary_emb is not None:
|
||||
query = self._apply_rotary_emb(query, rotary_emb)
|
||||
key = self._apply_rotary_emb(key, rotary_emb)
|
||||
|
||||
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
|
||||
# branch cannot express. Apply it above, then pass no RoPE here.
|
||||
@@ -450,9 +397,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
quant_config: QuantizationConfig | None,
|
||||
prefix: str,
|
||||
adaln_apply_silu: bool = True,
|
||||
fuse_modulate: bool = False,
|
||||
fuse_qknorm_rope: bool = False,
|
||||
fuse_swiglu: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
|
||||
@@ -464,7 +408,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
supported_attention_backends,
|
||||
quant_config,
|
||||
prefix=f"{prefix}.attn",
|
||||
fuse_qknorm_rope=fuse_qknorm_rope,
|
||||
)
|
||||
self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
|
||||
self.ff = MiniMaxH3FeedForward(
|
||||
@@ -472,7 +415,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
ffn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.ff",
|
||||
fuse_swiglu=fuse_swiglu,
|
||||
)
|
||||
self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
|
||||
time_embed_dim,
|
||||
@@ -481,7 +423,6 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
prefix=f"{prefix}.adaln_proj",
|
||||
apply_silu=adaln_apply_silu,
|
||||
)
|
||||
self.fuse_modulate = fuse_modulate
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -494,39 +435,19 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
t.to(hidden_states.dtype) for t in self.adaln_proj(temb))
|
||||
|
||||
use_modulate_fusion = self.fuse_modulate and _can_run_minimax_h3_fusion(hidden_states)
|
||||
if use_modulate_fusion:
|
||||
norm_hidden_states = fused_rmsnorm_modulate(
|
||||
hidden_states,
|
||||
self.norm1.weight,
|
||||
scale_msa,
|
||||
shift_msa,
|
||||
adaln_indices,
|
||||
self.norm1.eps,
|
||||
)
|
||||
else:
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
|
||||
residual = hidden_states
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
|
||||
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
|
||||
if use_modulate_fusion:
|
||||
hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate(
|
||||
hidden_states,
|
||||
attention_output,
|
||||
gate_msa,
|
||||
self.norm2.weight,
|
||||
scale_mlp,
|
||||
shift_mlp,
|
||||
adaln_indices,
|
||||
self.norm2.eps,
|
||||
)
|
||||
else:
|
||||
hidden_states = hidden_states + gate_msa.index_select(0, adaln_indices) * attention_output
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
|
||||
hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attention_output
|
||||
|
||||
residual = hidden_states
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
|
||||
feed_forward_output = self.ff(norm_hidden_states)
|
||||
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
|
||||
|
||||
class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
@@ -572,17 +493,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config, hf_config)
|
||||
arch = config.arch_config
|
||||
self.enabled_fusions = _enabled_minimax_h3_fusions()
|
||||
if self.enabled_fusions:
|
||||
if HAVE_TRITON:
|
||||
logger.info(
|
||||
"MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled and "
|
||||
"torch.compile-captured forwards fall back to eager).",
|
||||
",".join(sorted(self.enabled_fusions)))
|
||||
else:
|
||||
logger.warning(
|
||||
"FASTVIDEO_MINIMAX_H3_FUSIONS requested %s but Triton is unavailable; "
|
||||
"every forward stays on the eager path.", ",".join(sorted(self.enabled_fusions)))
|
||||
sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1
|
||||
if arch.num_attention_heads % sp_world_size:
|
||||
raise ValueError(f"MiniMax H3 attention heads ({arch.num_attention_heads}) must be divisible by "
|
||||
@@ -636,7 +546,7 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
"parameter, but factorized AdaLN weights are pinned to FP16 "
|
||||
"(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank "
|
||||
"checkpoint instead, then re-fit the basis with "
|
||||
"scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.")
|
||||
"tools/minimax_h3/fit_adaln_basis.py.")
|
||||
adaln_dim = self.adaln_rank or arch.time_embed_dim
|
||||
self.adaln_basis = ReplicatedLinear(
|
||||
arch.time_embed_dim,
|
||||
@@ -680,9 +590,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
config.quant_config,
|
||||
prefix=f"{config.prefix}.transformer_blocks.{index}",
|
||||
adaln_apply_silu=self.adaln_rank is None,
|
||||
fuse_modulate="modulate" in self.enabled_fusions,
|
||||
fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions,
|
||||
fuse_swiglu="swiglu" in self.enabled_fusions,
|
||||
) for index in range(arch.num_layers)
|
||||
])
|
||||
self.norm_out = MiniMaxH3AdaLayerNormOut(
|
||||
@@ -709,20 +616,6 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
)
|
||||
self.__post_init__()
|
||||
|
||||
def prepare_for_compile(self) -> None:
|
||||
"""Pipeline hook, called once right before torch.compile wraps the blocks.
|
||||
|
||||
Dynamo capture traces the eager branch of every fusion guard, so an
|
||||
enabled ``FASTVIDEO_MINIMAX_H3_FUSIONS`` set is silently inert inside
|
||||
compiled block forwards (H3 compiles per-block by default). Say so
|
||||
once instead of leaving the flag looking active.
|
||||
"""
|
||||
if self.enabled_fusions:
|
||||
logger.warning(
|
||||
"torch.compile is enabled for MiniMax H3, so the requested inference fusions (%s) are "
|
||||
"inert inside compiled block forwards; the compiled eager path runs instead.",
|
||||
",".join(sorted(self.enabled_fusions)))
|
||||
|
||||
def materialize_non_persistent_buffers(
|
||||
self,
|
||||
device: torch.device,
|
||||
@@ -841,17 +734,14 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
|
||||
rotary_emb = (rotary_cos, rotary_sin)
|
||||
|
||||
# The eager driver owns profiling markers while each block's compiled
|
||||
# forward owns the graph that the marker surrounds.
|
||||
for block_index, block in enumerate(self.transformer_blocks):
|
||||
with nvtx_range(f"minimax_h3.transformer_block.{block_index}"):
|
||||
packed_hidden_states = block(
|
||||
packed_hidden_states,
|
||||
temb,
|
||||
adaln_indices,
|
||||
rotary_emb,
|
||||
original_seq_len,
|
||||
)
|
||||
for block in self.transformer_blocks:
|
||||
packed_hidden_states = block(
|
||||
packed_hidden_states,
|
||||
temb,
|
||||
adaln_indices,
|
||||
rotary_emb,
|
||||
original_seq_len,
|
||||
)
|
||||
|
||||
packed_hidden_states = self.norm_out(
|
||||
packed_hidden_states,
|
||||
|
||||
@@ -1,20 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Inference-only MiniMax H3 fusions adapted from NVlabs/Sana Sol-Engine.
|
||||
|
||||
Source: https://github.com/NVlabs/Sana/tree/sol-engine/models/minimax_h3/GB200
|
||||
"""
|
||||
|
||||
from .modulation import (
|
||||
fused_residual_gate_rmsnorm_modulate,
|
||||
fused_rmsnorm_modulate,
|
||||
)
|
||||
from .qknorm_rope import HAVE_TRITON, fused_qknorm_rope
|
||||
from .swiglu import minimax_h3_swiglu
|
||||
|
||||
__all__ = [
|
||||
"HAVE_TRITON",
|
||||
"fused_qknorm_rope",
|
||||
"fused_residual_gate_rmsnorm_modulate",
|
||||
"fused_rmsnorm_modulate",
|
||||
"minimax_h3_swiglu",
|
||||
]
|
||||
@@ -1,302 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MiniMax H3 RMSNorm and row-indexed modulation fusions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
except ImportError as exc: # pragma: no cover - depends on the runtime image
|
||||
triton = None
|
||||
tl = None
|
||||
_TRITON_IMPORT_ERROR: ImportError | None = exc
|
||||
else:
|
||||
_TRITON_IMPORT_ERROR = None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"fused_residual_gate_rmsnorm_modulate",
|
||||
"fused_rmsnorm_modulate",
|
||||
]
|
||||
|
||||
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
|
||||
_rmsnorm_modulate_kernel = None
|
||||
_residual_gate_rmsnorm_modulate_kernel = None
|
||||
|
||||
|
||||
if triton is not None:
|
||||
|
||||
@triton.jit
|
||||
def _rmsnorm_modulate_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
weight_ptr,
|
||||
scale_ptr,
|
||||
shift_ptr,
|
||||
index_ptr,
|
||||
n_cols,
|
||||
n_index,
|
||||
eps,
|
||||
stride_x_row,
|
||||
stride_scale_row,
|
||||
stride_shift_row,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
cols = tl.arange(0, BLOCK)
|
||||
mask = cols < n_cols
|
||||
x_offsets = row * stride_x_row + cols
|
||||
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
|
||||
|
||||
x = tl.load(x_ptr + x_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
variance = tl.sum(x * x, axis=0) / n_cols
|
||||
normed = x * tl.math.rsqrt(variance + eps)
|
||||
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
scale = tl.load(
|
||||
scale_ptr + table_row * stride_scale_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
shift = tl.load(
|
||||
shift_ptr + table_row * stride_shift_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
output = normed * weight * (1.0 + scale) + shift
|
||||
tl.store(out_ptr + x_offsets, output.to(out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
@triton.jit
|
||||
def _residual_gate_rmsnorm_modulate_kernel(
|
||||
hidden_out_ptr,
|
||||
normed_out_ptr,
|
||||
residual_ptr,
|
||||
branch_ptr,
|
||||
gate_ptr,
|
||||
weight_ptr,
|
||||
scale_ptr,
|
||||
shift_ptr,
|
||||
index_ptr,
|
||||
n_cols,
|
||||
n_index,
|
||||
eps,
|
||||
stride_input_row,
|
||||
stride_gate_row,
|
||||
stride_scale_row,
|
||||
stride_shift_row,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
cols = tl.arange(0, BLOCK)
|
||||
mask = cols < n_cols
|
||||
input_offsets = row * stride_input_row + cols
|
||||
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
|
||||
|
||||
residual = tl.load(residual_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
branch = tl.load(branch_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
gate = tl.load(
|
||||
gate_ptr + table_row * stride_gate_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
hidden = residual + gate * branch
|
||||
tl.store(hidden_out_ptr + input_offsets, hidden.to(hidden_out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
variance = tl.sum(hidden * hidden, axis=0) / n_cols
|
||||
normed = hidden * tl.math.rsqrt(variance + eps)
|
||||
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
scale = tl.load(
|
||||
scale_ptr + table_row * stride_scale_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
shift = tl.load(
|
||||
shift_ptr + table_row * stride_shift_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
output = normed * weight * (1.0 + scale) + shift
|
||||
tl.store(normed_out_ptr + input_offsets, output.to(normed_out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
|
||||
def _validate_contract(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
tables: tuple[torch.Tensor, ...],
|
||||
index: torch.Tensor,
|
||||
eps: float,
|
||||
) -> None:
|
||||
if x.ndim < 2:
|
||||
raise ValueError(f"x must have shape (..., sequence_length, hidden_size), got {tuple(x.shape)}.")
|
||||
if x.numel() == 0 or x.shape[-1] == 0:
|
||||
raise ValueError("x must not be empty.")
|
||||
if x.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"x must use float16, bfloat16, or float32, got {x.dtype}.")
|
||||
hidden_size = x.shape[-1]
|
||||
sequence_length = x.shape[-2]
|
||||
if weight.shape != (hidden_size, ):
|
||||
raise ValueError(f"weight must have shape ({hidden_size},), got {tuple(weight.shape)}.")
|
||||
if weight.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"weight must use float16, bfloat16, or float32, got {weight.dtype}.")
|
||||
if index.ndim != 1 or index.numel() != sequence_length:
|
||||
raise ValueError(
|
||||
f"index must have shape ({sequence_length},) so it can wrap over batch rows, got {tuple(index.shape)}."
|
||||
)
|
||||
if index.dtype not in (torch.int32, torch.int64):
|
||||
raise TypeError(f"index must use int32 or int64, got {index.dtype}.")
|
||||
if not isinstance(eps, (float, int)) or isinstance(eps, bool) or not math.isfinite(eps) or eps <= 0:
|
||||
raise ValueError(f"eps must be a positive finite number, got {eps!r}.")
|
||||
|
||||
table_rows = tables[0].shape[0] if tables and tables[0].ndim == 2 else None
|
||||
for name, table in zip(("gate", "scale", "shift")[-len(tables):], tables, strict=True):
|
||||
if table.ndim != 2 or table.shape[1] != hidden_size:
|
||||
raise ValueError(f"{name} must have shape (table_rows, {hidden_size}), got {tuple(table.shape)}.")
|
||||
if table.shape[0] == 0 or table.shape[0] != table_rows:
|
||||
raise ValueError("all modulation tables must have the same non-zero row count.")
|
||||
if table.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"{name} must use float16, bfloat16, or float32, got {table.dtype}.")
|
||||
|
||||
tensors = (x, weight, *tables, index)
|
||||
if any(tensor.device != x.device for tensor in tensors[1:]):
|
||||
raise ValueError("x, weight, modulation tables, and index must be on the same device.")
|
||||
|
||||
|
||||
def _validate_residual_branch(residual: torch.Tensor, branch: torch.Tensor) -> None:
|
||||
if branch.shape != residual.shape:
|
||||
raise ValueError(f"branch must match residual shape {tuple(residual.shape)}, got {tuple(branch.shape)}.")
|
||||
if branch.dtype != residual.dtype:
|
||||
raise TypeError(f"branch dtype must match residual dtype {residual.dtype}, got {branch.dtype}.")
|
||||
if branch.device != residual.device:
|
||||
raise ValueError("branch and residual must be on the same device.")
|
||||
|
||||
|
||||
def _require_triton_cuda(x: torch.Tensor) -> None:
|
||||
if triton is None:
|
||||
detail = f": {_TRITON_IMPORT_ERROR}" if _TRITON_IMPORT_ERROR is not None else ""
|
||||
raise RuntimeError(f"MiniMax H3 modulation fusion requires Triton{detail}.")
|
||||
if x.device.type != "cuda":
|
||||
raise RuntimeError(f"MiniMax H3 modulation fusion requires CUDA tensors, got device {x.device}.")
|
||||
|
||||
|
||||
def _require_forward_only(*tensors: torch.Tensor) -> None:
|
||||
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in tensors):
|
||||
raise RuntimeError("MiniMax H3 modulation fusion is forward-only and does not support autograd.")
|
||||
|
||||
|
||||
def _next_power_of_two(value: int) -> int:
|
||||
return 1 << (value - 1).bit_length()
|
||||
|
||||
|
||||
def _num_warps(block_size: int) -> int:
|
||||
if block_size >= 8192:
|
||||
return 16
|
||||
if block_size >= 2048:
|
||||
return 8
|
||||
return 4
|
||||
|
||||
|
||||
def _row_addressable(table: torch.Tensor) -> torch.Tensor:
|
||||
return table if table.stride(-1) == 1 else table.contiguous()
|
||||
|
||||
|
||||
def fused_rmsnorm_modulate(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
index: torch.Tensor,
|
||||
eps: float,
|
||||
) -> torch.Tensor:
|
||||
"""Run RMSNorm and row-indexed modulation in one strict Triton kernel.
|
||||
|
||||
``index`` values must lie in ``[0, table_rows)``. Unlike eager
|
||||
``index_select``, the kernel does not raise on out-of-range values (a
|
||||
device-side bounds check would synchronize); callers are safe by
|
||||
construction (``timestep_indices * 3 + token_tags``, SP pads with 0).
|
||||
"""
|
||||
_validate_contract(x, weight, (scale, shift), index, eps)
|
||||
_require_forward_only(x, weight, scale, shift)
|
||||
_require_triton_cuda(x)
|
||||
|
||||
hidden_size = x.shape[-1]
|
||||
flat_x = x.reshape(-1, hidden_size).contiguous()
|
||||
weight = weight.contiguous()
|
||||
scale = _row_addressable(scale)
|
||||
shift = _row_addressable(shift)
|
||||
index = index.contiguous()
|
||||
output = torch.empty_like(flat_x)
|
||||
block_size = _next_power_of_two(hidden_size)
|
||||
_rmsnorm_modulate_kernel[(flat_x.shape[0], )](
|
||||
output,
|
||||
flat_x,
|
||||
weight,
|
||||
scale,
|
||||
shift,
|
||||
index,
|
||||
hidden_size,
|
||||
index.numel(),
|
||||
eps,
|
||||
flat_x.stride(0),
|
||||
scale.stride(0),
|
||||
shift.stride(0),
|
||||
BLOCK=block_size,
|
||||
num_warps=_num_warps(block_size),
|
||||
)
|
||||
return output.view_as(x)
|
||||
|
||||
|
||||
def fused_residual_gate_rmsnorm_modulate(
|
||||
residual: torch.Tensor,
|
||||
branch: torch.Tensor,
|
||||
gate: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
index: torch.Tensor,
|
||||
eps: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Fuse residual update, row-indexed gate, RMSNorm, and modulation.
|
||||
|
||||
``index`` values must lie in ``[0, table_rows)``; see
|
||||
:func:`fused_rmsnorm_modulate` for why the wrapper does not check them.
|
||||
"""
|
||||
_validate_residual_branch(residual, branch)
|
||||
_validate_contract(residual, weight, (gate, scale, shift), index, eps)
|
||||
_require_forward_only(residual, branch, gate, weight, scale, shift)
|
||||
_require_triton_cuda(residual)
|
||||
|
||||
hidden_size = residual.shape[-1]
|
||||
flat_residual = residual.reshape(-1, hidden_size).contiguous()
|
||||
flat_branch = branch.reshape(-1, hidden_size).contiguous()
|
||||
weight = weight.contiguous()
|
||||
gate = _row_addressable(gate)
|
||||
scale = _row_addressable(scale)
|
||||
shift = _row_addressable(shift)
|
||||
index = index.contiguous()
|
||||
hidden = torch.empty_like(flat_residual)
|
||||
modulated = torch.empty_like(flat_residual)
|
||||
block_size = _next_power_of_two(hidden_size)
|
||||
_residual_gate_rmsnorm_modulate_kernel[(flat_residual.shape[0], )](
|
||||
hidden,
|
||||
modulated,
|
||||
flat_residual,
|
||||
flat_branch,
|
||||
gate,
|
||||
weight,
|
||||
scale,
|
||||
shift,
|
||||
index,
|
||||
hidden_size,
|
||||
index.numel(),
|
||||
eps,
|
||||
flat_residual.stride(0),
|
||||
gate.stride(0),
|
||||
scale.stride(0),
|
||||
shift.stride(0),
|
||||
BLOCK=block_size,
|
||||
num_warps=_num_warps(block_size),
|
||||
)
|
||||
return hidden.view_as(residual), modulated.view_as(residual)
|
||||
@@ -1,174 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Fused per-head RMSNorm and partial rotary embedding for MiniMax H3."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
HAVE_TRITON = True
|
||||
except ImportError: # pragma: no cover - exercised only in environments without Triton
|
||||
triton = None
|
||||
tl = None
|
||||
HAVE_TRITON = False
|
||||
|
||||
|
||||
if HAVE_TRITON:
|
||||
|
||||
@triton.jit
|
||||
def _qknorm_partial_rope_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
weight_ptr,
|
||||
cos_ptr,
|
||||
sin_ptr,
|
||||
head_dim,
|
||||
rotary_dim,
|
||||
half_rotary_dim,
|
||||
num_heads,
|
||||
seq_len,
|
||||
eps,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
# int64, like the sibling kernels: with int32 program ids,
|
||||
# ``row * head_dim`` wraps once the flattened input reaches 2**31
|
||||
# elements (H3's 56 heads x 128 head_dim crosses that at
|
||||
# batch*seq >= 299_593 tokens per rank) and the loads/stores below
|
||||
# become out-of-bounds. ``seq_index`` inherits int64 from ``row``.
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
seq_index = (row // num_heads) % seq_len
|
||||
cols = tl.arange(0, BLOCK_SIZE)
|
||||
head_mask = cols < head_dim
|
||||
row_offset = row * head_dim
|
||||
|
||||
x = tl.load(x_ptr + row_offset + cols, mask=head_mask, other=0.0).to(tl.float32)
|
||||
variance = tl.sum(x * x, axis=0) / head_dim
|
||||
inv_rms = tl.math.rsqrt(variance + eps)
|
||||
weight = tl.load(weight_ptr + cols, mask=head_mask, other=0.0).to(tl.float32)
|
||||
normalized = x * inv_rms * weight
|
||||
|
||||
rotary_mask = cols < rotary_dim
|
||||
first_half = cols < half_rotary_dim
|
||||
partner_col = tl.where(first_half, cols + half_rotary_dim, cols - half_rotary_dim)
|
||||
partner_x = tl.load(
|
||||
x_ptr + row_offset + partner_col,
|
||||
mask=rotary_mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
partner_weight = tl.load(weight_ptr + partner_col, mask=rotary_mask, other=0.0).to(tl.float32)
|
||||
partner_normalized = partner_x * inv_rms * partner_weight
|
||||
rotated = tl.where(first_half, -partner_normalized, partner_normalized)
|
||||
|
||||
table_offset = seq_index * rotary_dim + cols
|
||||
cos = tl.load(cos_ptr + table_offset, mask=rotary_mask, other=1.0).to(tl.float32)
|
||||
sin = tl.load(sin_ptr + table_offset, mask=rotary_mask, other=0.0).to(tl.float32)
|
||||
rotary_output = normalized * cos + rotated * sin
|
||||
output = tl.where(rotary_mask, rotary_output, normalized)
|
||||
tl.store(out_ptr + row_offset + cols, output.to(out_ptr.dtype.element_ty), mask=head_mask)
|
||||
|
||||
|
||||
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
|
||||
|
||||
|
||||
def _validate_inputs(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
eps: float,
|
||||
) -> tuple[int, int, int, int, int]:
|
||||
for name, tensor in (("x", x), ("weight", weight), ("cos", cos), ("sin", sin)):
|
||||
if not isinstance(tensor, torch.Tensor):
|
||||
raise TypeError(f"{name} must be a torch.Tensor, got {type(tensor).__name__}")
|
||||
|
||||
if x.ndim != 4:
|
||||
raise ValueError(f"x must have shape (batch, seq, heads, head_dim), got {tuple(x.shape)}")
|
||||
batch, seq_len, num_heads, head_dim = x.shape
|
||||
if min(batch, seq_len, num_heads, head_dim) <= 0:
|
||||
raise ValueError(f"x dimensions must all be positive, got {tuple(x.shape)}")
|
||||
if weight.shape != (head_dim, ):
|
||||
raise ValueError(f"weight must have shape ({head_dim},), got {tuple(weight.shape)}")
|
||||
if cos.ndim != 2:
|
||||
raise ValueError(f"cos must have shape (seq, rotary_dim), got {tuple(cos.shape)}")
|
||||
if sin.shape != cos.shape:
|
||||
raise ValueError(f"sin must match cos shape {tuple(cos.shape)}, got {tuple(sin.shape)}")
|
||||
if cos.shape[0] != seq_len:
|
||||
raise ValueError(f"cos/sin sequence length must be {seq_len}, got {cos.shape[0]}")
|
||||
|
||||
rotary_dim = cos.shape[1]
|
||||
if rotary_dim <= 0:
|
||||
raise ValueError(f"rotary_dim must be positive, got {rotary_dim}")
|
||||
if rotary_dim > head_dim:
|
||||
raise ValueError(f"rotary_dim must not exceed head_dim, got rotary_dim={rotary_dim}, head_dim={head_dim}")
|
||||
if rotary_dim % 2:
|
||||
raise ValueError(f"rotary_dim must be even, got {rotary_dim}")
|
||||
|
||||
if x.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"x dtype must be float16, bfloat16, or float32, got {x.dtype}")
|
||||
for name, tensor in (("weight", weight), ("cos", cos), ("sin", sin)):
|
||||
if tensor.dtype != x.dtype:
|
||||
raise TypeError(f"{name} dtype must match x dtype {x.dtype}, got {tensor.dtype}")
|
||||
if tensor.device != x.device:
|
||||
raise ValueError(f"{name} device must match x device {x.device}, got {tensor.device}")
|
||||
|
||||
if not isinstance(eps, (float, int)) or not math.isfinite(float(eps)) or eps <= 0:
|
||||
raise ValueError(f"eps must be a positive finite number, got {eps!r}")
|
||||
return batch, seq_len, num_heads, head_dim, rotary_dim
|
||||
|
||||
|
||||
def fused_qknorm_rope(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
eps: float,
|
||||
) -> torch.Tensor:
|
||||
"""Run per-head RMSNorm and partial RoPE in one Sol-Engine-style kernel.
|
||||
|
||||
RMSNorm reduction and RoPE arithmetic stay in FP32 registers until the
|
||||
final store. Triton's reduction order and the absence of eager's BF16
|
||||
intermediate materializations can produce small, expected rounding drift.
|
||||
|
||||
Row offsets are computed in int64, so inputs beyond 2**31 total elements
|
||||
(about 300k tokens per rank at H3's 56 heads x 128 head_dim) address
|
||||
correctly.
|
||||
"""
|
||||
batch, seq_len, num_heads, head_dim, rotary_dim = _validate_inputs(x, weight, cos, sin, eps)
|
||||
if not weight.is_contiguous():
|
||||
raise ValueError("weight must be contiguous")
|
||||
if not cos.is_contiguous() or not sin.is_contiguous():
|
||||
raise ValueError("cos and sin must be contiguous (seq, rotary_dim) tables")
|
||||
if not x.is_cuda:
|
||||
raise RuntimeError("fused_qknorm_rope requires CUDA tensors")
|
||||
if not HAVE_TRITON:
|
||||
raise RuntimeError("fused_qknorm_rope requires Triton")
|
||||
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in (x, weight, cos, sin)):
|
||||
raise RuntimeError("fused_qknorm_rope is inference-only and does not implement autograd")
|
||||
|
||||
flat_x = x.reshape(-1, head_dim).contiguous()
|
||||
flat_out = torch.empty_like(flat_x)
|
||||
block_size = 1 << (head_dim - 1).bit_length()
|
||||
_qknorm_partial_rope_kernel[(flat_x.shape[0], )](
|
||||
flat_out,
|
||||
flat_x,
|
||||
weight,
|
||||
cos,
|
||||
sin,
|
||||
head_dim,
|
||||
rotary_dim,
|
||||
rotary_dim // 2,
|
||||
num_heads,
|
||||
seq_len,
|
||||
eps,
|
||||
BLOCK_SIZE=block_size,
|
||||
num_warps=4,
|
||||
)
|
||||
return flat_out.view(batch, seq_len, num_heads, head_dim)
|
||||
|
||||
|
||||
__all__ = ["HAVE_TRITON", "fused_qknorm_rope"]
|
||||
@@ -1,104 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MiniMax H3's value-first packed SwiGLU fusion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
HAVE_TRITON = True
|
||||
except ImportError: # pragma: no cover - exercised only in environments without Triton
|
||||
triton = None
|
||||
tl = None
|
||||
HAVE_TRITON = False
|
||||
|
||||
|
||||
def _validate_input(x: torch.Tensor) -> int:
|
||||
if x.ndim == 0:
|
||||
raise ValueError("MiniMax H3 SwiGLU expects at least one dimension")
|
||||
|
||||
packed_width = x.shape[-1]
|
||||
if packed_width == 0 or packed_width % 2 != 0:
|
||||
raise ValueError(
|
||||
"MiniMax H3 SwiGLU expects a positive even last dimension containing packed (value, gate) halves, "
|
||||
f"got {packed_width}"
|
||||
)
|
||||
if not x.is_floating_point():
|
||||
raise TypeError(f"MiniMax H3 SwiGLU expects a floating-point tensor, got {x.dtype}")
|
||||
return packed_width // 2
|
||||
|
||||
|
||||
if HAVE_TRITON:
|
||||
|
||||
@triton.jit
|
||||
def _minimax_h3_swiglu_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
ffn_dim,
|
||||
stride_in_row,
|
||||
stride_out_row,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
cols = tl.arange(0, BLOCK_SIZE)
|
||||
mask = cols < ffn_dim
|
||||
|
||||
value = tl.load(x_ptr + row * stride_in_row + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
gate = tl.load(x_ptr + row * stride_in_row + ffn_dim + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
# Match Sol-Engine: keep the complete SwiGLU expression in FP32 and
|
||||
# convert only the final output store.
|
||||
out = value * (gate * tl.sigmoid(gate))
|
||||
tl.store(out_ptr + row * stride_out_row + cols, out.to(out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
else:
|
||||
_minimax_h3_swiglu_kernel = None
|
||||
|
||||
|
||||
def _num_warps(block_size: int) -> int:
|
||||
if block_size >= 8192:
|
||||
return 16
|
||||
if block_size >= 2048:
|
||||
return 8
|
||||
return 4
|
||||
|
||||
|
||||
def minimax_h3_swiglu(x: torch.Tensor) -> torch.Tensor:
|
||||
"""Run the forward-only Triton fusion over an H3 ``(..., 2 * ffn_dim)`` input.
|
||||
|
||||
This is intentionally a strict kernel wrapper: callers own fallback policy and
|
||||
must only invoke it for a supported CUDA inference path.
|
||||
"""
|
||||
ffn_dim = _validate_input(x)
|
||||
if not x.is_cuda:
|
||||
raise ValueError("MiniMax H3 fused SwiGLU requires a CUDA tensor")
|
||||
if x.dtype not in (torch.float16, torch.bfloat16, torch.float32):
|
||||
raise TypeError(f"MiniMax H3 fused SwiGLU supports float16, bfloat16, and float32, got {x.dtype}")
|
||||
if torch.is_grad_enabled() and x.requires_grad:
|
||||
raise RuntimeError("MiniMax H3 fused SwiGLU is forward-only and does not implement autograd")
|
||||
if _minimax_h3_swiglu_kernel is None:
|
||||
raise RuntimeError("MiniMax H3 fused SwiGLU requires Triton")
|
||||
|
||||
packed_width = x.shape[-1]
|
||||
flat = x.reshape(-1, packed_width).contiguous()
|
||||
output_shape = (*x.shape[:-1], ffn_dim)
|
||||
if flat.shape[0] == 0:
|
||||
return torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
|
||||
out = torch.empty((flat.shape[0], ffn_dim), dtype=x.dtype, device=x.device)
|
||||
block_size = triton.next_power_of_2(ffn_dim)
|
||||
_minimax_h3_swiglu_kernel[(flat.shape[0],)](
|
||||
out,
|
||||
flat,
|
||||
ffn_dim,
|
||||
flat.stride(0),
|
||||
out.stride(0),
|
||||
BLOCK_SIZE=block_size,
|
||||
num_warps=_num_warps(block_size),
|
||||
)
|
||||
return out.view(output_shape)
|
||||
|
||||
|
||||
__all__ = ["HAVE_TRITON", "minimax_h3_swiglu"]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user