Compare commits

...
Author SHA1 Message Date
William Lin be1df9fec5 [ci][diagnostic]: capture GameCraft FA4 candidate on GB200 2026-08-26 22:51:48 +00:00
KyleNeverGivesUp e9bbaca07d [perf] Disable every offload path on unified memory, unblocking MiniMax H3 generation on one GB10 (#1715) 2026-08-26 15:32:56 -07:00
KyleNeverGivesUp 9bfa585448 [perf]: stop holding the whole checkpoint during DiT load, unblocking MiniMax H3 on one GB10 (#1714) 2026-08-26 15:23:51 -07:00
Raghav K b2062556a9 [perf] VSA Triton: widen the autotune num_stages range (the optimum was outside it) (#1706) 2026-08-26 12:30:56 -07:00
KyleNeverGivesUp c9c5585758 [perf]: MiniMax H3 on GB10 - skip text encoder CPU offload on unified memory (5m49s to 30ms) (#1710) 2026-08-26 12:03:14 -07:00
William Lin 9212f4f218 [ci] make GPU validation change-aware (#1747) 2026-08-25 21:26:25 -07:00
Aryan KumarandAryan Kumar 6388db815b [bugfix] FastMetal-QAD MLX support: refuse CUDA QAD trees, use packed mlx_dit config, stream loads (#1736) (#1758)
Co-authored-by: Aryan Kumar <aryan5v@users.noreply.github.com>
2026-08-25 14:51:40 -07:00
lpc0220 7a4285189f [kernel] Route block-sparse VSA to the sm_100a forward behind FASTVIDEO_VSA_SM100A (opt-in) (#1754) 2026-08-24 15:26:56 -07:00
William Lin a837fe841a [docs] Update FastH3 README (#1749) 2026-08-23 05:05:58 -07:00
William Lin f9e3680f11 [perf] Align FastH3 optimized inference profile (#1748) 2026-08-23 02:01:03 -07:00
William Lin 98f761ec45 [bugfix] validation: inherit the trained denoising ladder (#1738) 2026-08-22 23:09:26 -07:00
Shao Duan c041318f2c [perf] Add fused NVLink all-to-all for Ulysses (#1740) 2026-08-22 23:06:24 -07:00
William Lin 604e0205a4 [perf] Keep odd MiniMax-H3 VSA tiles on sm100a (#1745) 2026-08-22 18:29:30 -07:00
William Lin 13213395b4 [perf] Parallelize MiniMax-H3 VAE over sequence ranks (#1744) 2026-08-22 18:05:31 -07:00
William Lin 46afee5998 [bugfix] Classify MiniMax-H3 inference controls in schema inventory (#1743) 2026-08-22 17:17:02 -07:00
William Lin c488fa1211 [perf] Add opt-in packed-varlen FA4 for MiniMax-H3 (#1742) 2026-08-22 17:16:47 -07:00
William Lin d3cff517cd [perf] Add opt-in regional fullgraph compile for DiT inference (#1741) 2026-08-22 12:14:39 -07:00
Junda Su 2f3d407406 [perf] Optimize MiniMax H3 VAE decoding (#1734) 2026-08-21 14:57:32 -07:00
Kaiqin Kong bcffa4026e [perf] Optimize MiniMax-H3 text encoder memory (#1732) 2026-08-21 14:57:06 -07:00
William Lin 6d6a10be7a [feat] FastVideo-Minimax-FastH3-Preview few-step example + 64-token-tile VSA-H3 inference path (#1731) 2026-08-21 12:40:09 -05:00
Kaiqin Kong 73dd105f3d [perf] Add opt-in MiniMax-H3 Sol-Engine fusions (#1735) 2026-08-21 12:39:28 -05:00
191 changed files with 14952 additions and 1552 deletions
+99
View File
@@ -0,0 +1,99 @@
---
name: ci-runner
description: Work on FastVideo's Slurm-only, change-aware GPU CI lanes, static Buildkite graph, trusted ci-runner policy, lane scripts, and GB200 validation.
---
# Slinky Slurm CI lanes
FastVideo's `ci-runner` Buildkite queue is the control plane for all active
GPU CI. A private host-owned dispatcher leases GPUs from the Slinky Slurm tray
and runs the immutable PR SHA inside an isolated Enroot container. Buildkite
pipeline upload and Slurm submission occur on the login plane; every test
payload executes on Slurm compute.
The files under `fastvideo/tests/modal/` and `.buildkite/scripts/pr_test.sh`
are dormant rollback code. Never add an active Buildkite or slash-command
route to them. `pr_test.sh` must continue to reject Buildkite invocations.
The private operator bundle is deliberately outside this repository because
it contains site paths and credentials. See
`docs/contributing/ci_architecture.md`; this skill covers the repository half
and the coordination contract with that bundle.
## Invariants
- `.buildkite/pipeline.yml` contains exactly one static step for every active
GPU lane. Each step pins a unique key and label, a 90-minute timeout, the
trusted `/opt/fastvideo-ci-runner/run-ci` command (`run-unit` is the one
compatibility wrapper), step-level internal `TEST_TYPE`, and
`queue: "ci-runner"`.
- Active CI contains no `pr_test.sh` command, Modal invocation, default queue,
Buildkite plugin, `soft_fail`, or job-controlled artifact glob.
- The six Fastcheck lanes use `:microscope:` labels. Full-Suite-only lanes use
`:test_tube:` or `:bar_chart:` so direct reruns update the right aggregate.
- SSIM and vanilla training request all four GPUs. Keep both in the
`fastvideo/slinky/whole-tray` Buildkite concurrency group with a limit of one
so the second job does not consume an agent or command timeout while waiting
for the same tray.
- `/test full` schedules all twenty lanes. `/merge`, `ready`, and new pushes to
ready PRs use the trusted base-branch planner in
`.github/scripts/plan_merge_ci.py`: automatic Fastcheck remains the universal
six-lane baseline, and the merge build adds only path-relevant integration
lanes. Unknown source/build paths fail closed to all fourteen additive lanes.
The trusted uploader still normalizes and validates the complete static graph
before Buildkite evaluates its plan conditions.
- Focused merge builds may pass allowlisted golden-gate and SSIM test basenames.
The private host validates the lane plan and basenames before staging them,
and the in-container scripts validate them again. Direct `/test ssim`,
explicit `/test full`, and the weekly main-branch schedule run the complete
SSIM matrix.
- The trusted uploader serves exactly three entry pipelines:
`pr-fastcheck` for automatic PR builds, `ci` for slash-command/ready-label
API builds, and `fastvideo-performance-lane` for the weekly schedule. Keep
incoming GitHub webhook processing disabled on `ci` so it cannot duplicate
`pr-fastcheck` on every PR update.
- Test payloads live in `.buildkite/scripts/unit_test.sh` or executable
`.buildkite/scripts/lanes/<lane>.sh`. Backend policy (GPU count, extras,
secrets, kernel build, artifacts) stays in the agent-owned lane table.
- Tests must preserve an inherited `MASTER_PORT`. Packed containers share the
tray network namespace, so the private runner assigns a distinct port range
per GPU lease and the SSIM scheduler assigns task offsets within its range.
- The ARM64 runner image includes the pinned FA4 CuTe overlay validated on
GB200. Keep SSIM at `FASTVIDEO_FA4=1` because its references were seeded with
FA4; keep lanes with FA2 baselines at `FASTVIDEO_FA4=0`. A runner image change
must revalidate both the FA4 import and an actual GB200 forward kernel.
- `fastvideo/tests/ssim/ci_runner.py` is the active four-GPU SSIM scheduler.
New SSIM files are discovered through `REQUIRED_GPUS` and
`*_MODEL_TO_PARAMS`; do not wire them through the dormant Modal scheduler.
- The host policy fail-closes unknown tuples. A repository-side lane change is
inert until the operator updates the private lane table and uploader policy
in the same rollout.
## Adding or changing a lane
1. Read the closest `AGENTS.md` and the domain-specific testing guide.
2. Add or update the executable lane payload under `.buildkite/scripts/`.
Keep it deterministic and free of host-specific paths or credential fetches.
3. Add the static pipeline step and canonical `/test <name>` mapping. Keep the
`<name>-ci` alias only when compatibility requires it.
4. Add its source/test path ownership to `.github/scripts/plan_merge_ci.py`.
Prefer the narrowest correctness-preserving lane set; leave unknown paths
fail-closed. Extend `fastvideo/tests/contract/test_ci_test_collection.py`,
`test_merge_ci_plan.py`, and focused CPU-only scheduler/policy tests.
5. Coordinate the private lane row: GPU count (1-4), wall time, script, scope
pairs, step key, command, HF cache/token, tracking mode, extras, attention
backend policy, kernel policy, and artifact relay. Active training lanes
keep W&B offline and do not stage a W&B credential.
6. Update the trusted pipeline-uploader schema. A mismatch must reject the
pipeline rather than silently skip a lane.
7. Run `pre-commit run --files <changed paths>`, the planner's representative
diff matrix, contract tests, private driver tests, and a real GB200 canary.
Multi-GPU, hardware-reference, training, performance, and SSIM changes need
their own target-hardware evidence.
## Rollback
Rollback the Slurm routing/configuration change or pause the `ci-runner` queue.
Do not silently reactivate Modal. A manual Modal experiment requires the
explicit local opt-in documented in `ci_architecture.md`; returning it to
production CI needs a separate reviewed decision.
@@ -1,6 +1,6 @@
---
name: reseed-ssim-references
description: Re-seed HF reference videos for a single existing SSIM test on Modal L40S. Always backs up current refs locally first, regenerates on Modal, pauses for the user to eyeball before-vs-after quality, then overwrites the targeted `<model_id>` subtree on `FastVideo/ssim-reference-videos` with `--force`. Use when an intentional code change (model port fix, attention backend swap, kernel upgrade, hyperparameter change) has invalidated existing refs and they need to be regenerated. Pairs with `seed-ssim-references`, which is for first-time seeding only.
description: Re-seed HF reference videos for a single existing SSIM test on Modal L40S. Always backs up current refs locally first, regenerates on Modal, pauses for the user to eyeball before-vs-after quality, then overwrites the targeted model subtree on `FastVideo/ssim-reference-videos` with `--force`. Use when an intentional code change (model port fix, attention backend swap, kernel upgrade, hyperparameter change) has invalidated existing refs and they need to be regenerated. Pairs with `seed-ssim-references`, which is for first-time seeding only.
---
# Re-seed SSIM Reference Videos
@@ -13,7 +13,7 @@ on HF — the old refs are overwritten — so the skill always:
1. Confirms intent with a one-liner the user has to type.
2. Downloads the existing refs as a local, timestamped backup.
3. Regenerates on Modal L40S (same code path that CI uses).
3. Regenerates through the manual legacy Modal L40S maintenance path.
4. Pauses for a side-by-side eyeball of backup vs new mp4s.
5. Uploads with `--force`, scoped to the single `--model-id`.
6. Reminds the user to keep the backup until the PR lands.
@@ -51,8 +51,9 @@ harder to recover from than failing closed.
Hardcoded:
- Modal GPU: **L40S** (matches CI; re-seeding from another SKU produces refs
that L40S CI cannot match).
- Modal GPU: **L40S**. This is a manual reference-maintenance target, not the
active Slurm CI compute path; changing the SKU also changes the historical
`L40S_reference_videos` contract.
- Quality tier: **`default`**. `full_quality` is a separate, deliberate
operation.
- HF repo: `FastVideo/ssim-reference-videos` (override via
+4 -2
View File
@@ -35,7 +35,8 @@ The skill is run **manually**, once per new test. Before invoking it, the user
has already sanity-tested the new test locally — it launches `VideoGenerator`
and writes an artefact without crashing (the missing-reference assertion at
the end is expected). The skill does not re-test locally; it goes straight
to Modal L40S (which is what CI uses).
to the manual legacy Modal L40S reference-maintenance target. Active CI runs
on the Slinky Slurm cluster and only consumes the resulting references.
## When to use
@@ -61,7 +62,8 @@ Prompt the user for it if they didn't supply it.
Everything else is fixed:
- Modal runner GPU: **L40S** (hardcoded in `fastvideo/tests/modal/ssim_test.py`).
- Modal maintenance GPU: **L40S** (hardcoded in
`fastvideo/tests/modal/ssim_test.py`; this is not the active CI compute path).
- Device folder: `L40S_reference_videos`.
- Quality tier: `default` (the tier CI runs). The `full_quality` tier is not
seeded by this skill.
+448 -528
View File
@@ -1,7 +1,8 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
# Buildkite only launches Modal; remote jobs initialize their own submodules.
# Slurm workers clone the immutable commit and initialize submodules inside
# their isolated container. The Buildkite login-plane checkout is a no-op.
BUILDKITE_GIT_SUBMODULES: false
notify:
@@ -10,539 +11,458 @@ notify:
if: build.env("TEST_SCOPE") == "fastcheck" || build.env("TEST_SCOPE") == null
- github_commit_status:
context: "full-suite-passed"
if: build.env("TEST_SCOPE") == "full"
if: build.env("TEST_SCOPE") == "full" || build.env("TEST_SCOPE") == "merge"
- github_commit_status:
context: "direct-test-completed"
if: build.env("TEST_SCOPE") == "direct"
- github_commit_status:
context: "scheduled-ssim-passed"
if: build.env("TEST_SCOPE") == "scheduled"
# This is the complete active GPU CI surface. Every command is a trusted host
# dispatcher, and every test payload executes inside the Slinky Slurm tray.
# fastvideo/tests/modal remains available only for an explicit manual rollback;
# no active pipeline or slash-command route invokes it.
steps:
# ============================================================
# Direct test: triggered by /test <name> slash command.
# Labels match fastcheck/full-suite counterparts so the GitHub
# check status overwrites the original failed check.
# Only ONE step executes per build (gated by TEST_TYPE).
# ============================================================
- label: ":microscope: Encoder Tests"
key: "encoder"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,encoder,/) ||
build.env("TEST_SCOPE") == "fastcheck" ||
build.env("TEST_SCOPE") == null ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "encoder" || build.env("TEST_TYPE") == "encoder_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "encoder_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
# --- Fastcheck-scope direct tests ---
- label: ":microscope: Encoder Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "encoder"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: VAE Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "vae"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: Transformer Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "transformer"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: Kernel Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "kernel_tests"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":vertical_traffic_light: Golden-Gate Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "golden_gate"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: Unit Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "unit_test"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: DreamVerse App Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "dreamverse_app"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: VAE Tests"
key: "vae"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,vae,/) ||
build.env("TEST_SCOPE") == "fastcheck" ||
build.env("TEST_SCOPE") == null ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "vae" || build.env("TEST_TYPE") == "vae_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "vae_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
# --- Full-suite-scope direct tests ---
- label: ":bar_chart: SSIM Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "ssim"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
- exit_status: 1
limit: 2
agents:
queue: "default"
- label: ":test_tube: LoRA Inference Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "inference_lora"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: LoRA Extraction Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Training Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Distillation DMD Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "distillation_dmd"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Self-Forcing Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "self_forcing"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: LoRA Training Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training_lora"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
- exit_status: 1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Training Tests VSA"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training_vsa"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
- exit_status: 1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Inference Tests VMoBA"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "inference_vmoba"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Performance Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "performance"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: API Server Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "api_server"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Train Framework Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "train_framework"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Eval Metrics Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "eval"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":microscope: Transformer Tests"
key: "transformer"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,transformer,/) ||
build.env("TEST_SCOPE") == "fastcheck" ||
build.env("TEST_SCOPE") == null ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "transformer" || build.env("TEST_TYPE") == "transformer_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "transformer_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
# ============================================================
# Fastcheck: Runs on every PR (~10-15 min parallel)
# Core component validation: encoders, VAEs, transformers,
# CUDA kernels, and unit tests.
# ============================================================
- label: "Trigger Fastcheck"
if: build.env("TEST_SCOPE") == "fastcheck" || build.env("TEST_SCOPE") == null
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
plugins:
- monorepo-diff#v1.4.0:
diff: 'git fetch origin "${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}" && git diff --name-only "origin/${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}...HEAD"'
watch:
- path:
- "fastvideo/models/encoders/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/encoders/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: ":microscope: Encoder Tests"
env:
- TEST_TYPE=encoder
agents:
queue: "default"
- path:
- "fastvideo/models/vaes/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/vaes/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: ":microscope: VAE Tests"
env:
- TEST_TYPE=vae
agents:
queue: "default"
- path:
- "fastvideo/models/dits/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/layers/**"
- "fastvideo/attention/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":microscope: Transformer Tests"
env:
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":microscope: Kernel Tests"
env:
- TEST_TYPE=kernel_tests
agents:
queue: "default"
- path:
- "fastvideo/**"
- ".buildkite/**"
- ".github/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":microscope: Unit Tests"
env:
- TEST_TYPE=unit_test
agents:
queue: "default"
- path:
- "apps/dreamverse/**"
- "pyproject.toml"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: ":microscope: DreamVerse App Tests"
env:
- TEST_TYPE=dreamverse_app
agents:
queue: "default"
- label: ":microscope: Kernel Tests"
key: "kernel-tests"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,kernel-tests,/) ||
build.env("TEST_SCOPE") == "fastcheck" ||
build.env("TEST_SCOPE") == null ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "kernel_tests" || build.env("TEST_TYPE") == "kernel_tests_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "kernel_tests_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
# ============================================================
# Full Suite: Runs when TEST_SCOPE=full
# Triggered by adding the 'ready' label (via ci-trigger-full-suite.yml)
# or on-demand via /test full slash command.
# Includes integration tests, SSIM regression, training pipelines,
# and performance benchmarks.
# ============================================================
- label: "Trigger Full Suite"
if: build.env("TEST_SCOPE") == "full"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
plugins:
- monorepo-diff#v1.4.0:
diff: 'git fetch origin "${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}" && git diff --name-only "origin/${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-main}...HEAD"'
watch:
- path:
- "fastvideo/**/*.py"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: ":bar_chart: SSIM Tests"
env:
- TEST_TYPE=ssim
retry:
automatic:
- exit_status: 1
limit: 2
agents:
queue: "default"
- path:
- "fastvideo/tests/lora/**"
- "fastvideo/models/loader/**"
- "fastvideo/tests/transformers/**"
- "fastvideo/pipelines/**"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 20m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Inference Tests"
env:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "scripts/lora_extraction/**"
- "fastvideo/tests/lora_extraction/**"
- "fastvideo/models/loader/**"
- "fastvideo/training/training_utils.py"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 25m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Training Tests"
env:
- TEST_TYPE=training
agents:
queue: "default"
- path:
- "fastvideo/training/*distillation_pipeline.py"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 25m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Distillation DMD Tests"
env:
- TEST_TYPE=distillation_dmd
agents:
queue: "default"
- path:
- "fastvideo/training/*self_forcing_distillation_pipeline.py"
- "fastvideo/tests/training/self-forcing/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Self-Forcing Tests"
env:
- TEST_TYPE=self_forcing
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 25m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Training Tests"
env:
- TEST_TYPE=training_lora
retry:
automatic:
- exit_status: 1
limit: 2
agents:
queue: "default"
- path:
- "fastvideo/**"
- "fastvideo-kernel/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Training Tests VSA"
env:
- TEST_TYPE=training_vsa
retry:
automatic:
- exit_status: 1
limit: 2
agents:
queue: "default"
- path:
- "fastvideo-kernel/**"
- "fastvideo/attention/backends/vmoba.py"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Inference Tests VMoBA"
env:
- TEST_TYPE=inference_vmoba
agents:
queue: "default"
- path:
- "fastvideo/models/dits/**"
- "fastvideo/pipelines/**"
- "fastvideo/attention/**"
- "fastvideo/layers/**"
- "fastvideo/worker/**"
- "fastvideo/entrypoints/**"
- "fastvideo/performance/**"
- "fastvideo/tests/performance/**"
- ".buildkite/performance-benchmarks/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Performance Tests"
env:
- TEST_TYPE=performance
agents:
queue: "default"
- path:
- "fastvideo/entrypoints/openai/**"
- "fastvideo/entrypoints/cli/serve.py"
- "fastvideo/tests/entrypoints/test_openai_api_integration.py"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: ":test_tube: API Server Tests"
env:
- TEST_TYPE=api_server
agents:
queue: "default"
- path:
- "fastvideo/train/**"
- "fastvideo/tests/train/models/**"
- "fastvideo/tests/train/fixtures/**"
- "fastvideo/models/dits/**"
- "fastvideo/models/loader/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Train Framework Tests"
env:
- TEST_TYPE=train_framework
agents:
queue: "default"
- path:
- "fastvideo/eval/**"
- "fastvideo/tests/eval/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: ":test_tube: Eval Metrics Tests"
env:
- TEST_TYPE=eval
agents:
queue: "default"
- label: ":microscope: Unit Tests"
key: "unit"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,unit,/) ||
build.env("TEST_SCOPE") == "fastcheck" ||
build.env("TEST_SCOPE") == null ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "unit_test" || build.env("TEST_TYPE") == "unit_test_ci"))
command: "/opt/fastvideo-ci-runner/run-unit"
timeout_in_minutes: 90
env:
TEST_TYPE: "unit_test_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":microscope: DreamVerse App Tests"
key: "dreamverse"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,dreamverse,/) ||
build.env("TEST_SCOPE") == "fastcheck" ||
build.env("TEST_SCOPE") == null ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "dreamverse_app" || build.env("TEST_TYPE") == "dreamverse_app_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "dreamverse_app_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Golden-Gate Tests"
key: "golden-gate"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,golden-gate,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "golden_gate" || build.env("TEST_TYPE") == "golden_gate_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "golden_gate_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":bar_chart: SSIM Tests"
key: "ssim"
if: |
build.env("TEST_SCOPE") == "full" ||
build.env("TEST_SCOPE") == "scheduled" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,ssim,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "ssim" || build.env("TEST_TYPE") == "ssim_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
concurrency: 1
concurrency_group: "fastvideo/slinky/whole-tray"
env:
TEST_TYPE: "ssim_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
- exit_status: 1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: LoRA Inference Tests"
key: "lora-inference"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,lora-inference,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "inference_lora" || build.env("TEST_TYPE") == "inference_lora_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "inference_lora_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: LoRA Extraction Tests"
key: "lora-extraction"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,lora-extraction,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "lora_extraction" || build.env("TEST_TYPE") == "lora_extraction_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "lora_extraction_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Training Tests"
key: "training"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,training,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "training" || build.env("TEST_TYPE") == "training_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
concurrency: 1
concurrency_group: "fastvideo/slinky/whole-tray"
env:
TEST_TYPE: "training_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Distillation DMD Tests"
key: "distillation"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,distillation,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "distillation_dmd" || build.env("TEST_TYPE") == "distillation_dmd_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "distillation_dmd_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Self-Forcing Tests"
key: "self-forcing"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,self-forcing,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "self_forcing" || build.env("TEST_TYPE") == "self_forcing_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "self_forcing_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: LoRA Training Tests"
key: "lora-training"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,lora-training,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "training_lora" || build.env("TEST_TYPE") == "training_lora_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "training_lora_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
- exit_status: 1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Training Tests VSA"
key: "training-vsa"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,training-vsa,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "training_vsa" || build.env("TEST_TYPE") == "training_vsa_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "training_vsa_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
- exit_status: 1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Inference Tests VMoBA"
key: "inference-vmoba"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,inference-vmoba,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "inference_vmoba" || build.env("TEST_TYPE") == "inference_vmoba_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "inference_vmoba_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Performance Tests"
key: "performance"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,performance,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "performance" || build.env("TEST_TYPE") == "performance_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "performance_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: API Server Tests"
key: "api-server"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,api-server,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "api_server" || build.env("TEST_TYPE") == "api_server_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "api_server_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Train Framework Tests"
key: "train-framework"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,train-framework,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "train_framework" || build.env("TEST_TYPE") == "train_framework_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "train_framework_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
- label: ":test_tube: Eval Metrics Tests"
key: "eval"
if: |
build.env("TEST_SCOPE") == "full" ||
(build.env("TEST_SCOPE") == "merge" &&
build.env("MERGE_TEST_PLAN") =~ /,eval,/) ||
(build.env("TEST_SCOPE") == "direct" &&
(build.env("TEST_TYPE") == "eval" || build.env("TEST_TYPE") == "eval_ci"))
command: "/opt/fastvideo-ci-runner/run-ci"
timeout_in_minutes: 90
env:
TEST_TYPE: "eval_ci"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "ci-runner"
@@ -0,0 +1,148 @@
#!/usr/bin/env python3
"""Emit and recover a small GameCraft MP4 through a Buildkite job log."""
from __future__ import annotations
import argparse
import base64
import binascii
import hashlib
import html
import json
import re
import sys
from pathlib import Path
MARKER = "FV_GAMECRAFT_FA4_MP4"
CHUNK_BYTES = 3 * 1024
MAX_MEDIA_BYTES = 1024 * 1024
def _sha256(data: bytes) -> str:
return hashlib.sha256(data).hexdigest()
def emit(media_path: Path) -> None:
if media_path.suffix.lower() != ".mp4":
raise ValueError(f"Candidate must be an MP4: {media_path}")
data = media_path.read_bytes()
size = len(data)
if not 0 < size <= MAX_MEDIA_BYTES:
raise ValueError(f"Candidate size must be between 1 and {MAX_MEDIA_BYTES} bytes; got {size}")
digest = _sha256(data)
chunks = [data[offset:offset + CHUNK_BYTES] for offset in range(0, size, CHUNK_BYTES)]
print(
f"{MARKER}_BEGIN version=1 sha256={digest} size={size} "
f"chunks={len(chunks)} chunk_bytes={CHUNK_BYTES}",
flush=True,
)
for index, chunk in enumerate(chunks):
encoded = base64.b64encode(chunk).decode("ascii")
print(f"{MARKER}_CHUNK index={index:06d} data={encoded}", flush=True)
print(
f"{MARKER}_END version=1 sha256={digest} size={size} chunks={len(chunks)}",
flush=True,
)
def _buildkite_output(text: str) -> str:
"""Unwrap Buildkite's public JSON and reverse its HTML entity escaping."""
try:
payload = json.loads(text)
except json.JSONDecodeError:
return html.unescape(text)
if isinstance(payload, dict) and isinstance(payload.get("output"), str):
return html.unescape(payload["output"])
return html.unescape(text)
def decode(log_text: str) -> tuple[bytes, str]:
log_text = _buildkite_output(log_text)
begin_pattern = re.compile(
rf"{MARKER}_BEGIN version=1 sha256=([0-9a-f]{{64}}) size=([0-9]+) "
rf"chunks=([0-9]+) chunk_bytes=([0-9]+)"
)
end_pattern = re.compile(
rf"{MARKER}_END version=1 sha256=([0-9a-f]{{64}}) size=([0-9]+) chunks=([0-9]+)"
)
chunk_pattern = re.compile(rf"{MARKER}_CHUNK index=([0-9]{{6}}) data=([A-Za-z0-9+/]+={{0,2}})")
begin_matches = begin_pattern.findall(log_text)
end_matches = end_pattern.findall(log_text)
if len(begin_matches) != 1 or len(end_matches) != 1:
raise ValueError(
"Expected exactly one candidate envelope; "
f"found {len(begin_matches)} begin and {len(end_matches)} end markers"
)
digest, size_text, count_text, chunk_bytes_text = begin_matches[0]
end_digest, end_size_text, end_count_text = end_matches[0]
if (digest, size_text, count_text) != (end_digest, end_size_text, end_count_text):
raise ValueError("Candidate begin/end metadata does not match")
size = int(size_text)
count = int(count_text)
chunk_bytes = int(chunk_bytes_text)
if not 0 < size <= MAX_MEDIA_BYTES:
raise ValueError(f"Candidate size is outside the accepted range: {size}")
if chunk_bytes != CHUNK_BYTES:
raise ValueError(f"Unexpected chunk size: {chunk_bytes}")
matches = chunk_pattern.findall(log_text)
if len(matches) != count:
raise ValueError(f"Expected {count} candidate chunks; found {len(matches)}")
encoded_chunks: dict[int, str] = {}
for index_text, encoded in matches:
index = int(index_text)
if index in encoded_chunks:
raise ValueError(f"Duplicate candidate chunk: {index}")
encoded_chunks[index] = encoded
if sorted(encoded_chunks) != list(range(count)):
raise ValueError("Candidate chunk sequence is incomplete or out of range")
try:
data = b"".join(base64.b64decode(encoded_chunks[index], validate=True) for index in range(count))
except binascii.Error as error:
raise ValueError(f"Candidate chunk is not valid base64: {error}") from error
if len(data) != size:
raise ValueError(f"Decoded candidate size mismatch: expected {size}, got {len(data)}")
actual_digest = _sha256(data)
if actual_digest != digest:
raise ValueError(f"Decoded candidate sha256 mismatch: expected {digest}, got {actual_digest}")
return data, digest
def _parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
subparsers = parser.add_subparsers(dest="command", required=True)
emit_parser = subparsers.add_parser("emit", help="Encode one MP4 to stdout between strict log markers")
emit_parser.add_argument("media_path", type=Path)
decode_parser = subparsers.add_parser("decode", help="Recover and verify one MP4 from a raw or public JSON log")
decode_parser.add_argument("--input", type=Path, help="Log file (default: stdin)")
decode_parser.add_argument("--output", type=Path, required=True)
return parser.parse_args()
def main() -> int:
args = _parse_args()
if args.command == "emit":
emit(args.media_path)
return 0
log_text = args.input.read_text() if args.input else sys.stdin.read()
data, digest = decode(log_text)
if args.output.exists():
raise FileExistsError(f"Refusing to overwrite existing output: {args.output}")
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_bytes(data)
print(f"Recovered {len(data)} bytes with sha256={digest} to {args.output}")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the OpenAI-compatible API lane.
set -euo pipefail
exec pytest ./fastvideo/tests/entrypoints/test_openai_api_integration.py -vs
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the distillation-DMD lane.
set -euo pipefail
exec pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs
+87
View File
@@ -0,0 +1,87 @@
#!/usr/bin/env bash
# DreamVerse needs a GPU for import-time device resolution, but it does not
# build or exercise fastvideo-kernel. A checksummed Node archive is installed
# in the disposable Slurm container because the shared CI image is
# Python/CUDA focused.
set -euo pipefail
node_version=v22.23.2
case $(uname -m) in
aarch64 | arm64)
node_arch=arm64
node_archive_sha256=013b59cfd2819703a6f4a14ab891fc46fc2a4e3f5bcd92de3fb4929b43e35b30
;;
x86_64 | amd64)
node_arch=x64
node_archive_sha256=b294a556e639d64338823920e5866c21c02741742d2e1529ee1a225c1ec9252a
;;
*)
echo "Unsupported architecture for DreamVerse Node runtime: $(uname -m)" >&2
exit 2
;;
esac
node_archive="node-${node_version}-linux-${node_arch}.tar.gz"
node_runtime_root=$(mktemp -d -t fastvideo-node.XXXXXX)
node_archive_path="${node_runtime_root}/${node_archive}"
node_install_dir="${node_runtime_root}/${node_archive%.tar.gz}"
curl --proto '=https' --tlsv1.2 --retry 5 --retry-all-errors \
--location --fail --silent --show-error \
"https://nodejs.org/dist/${node_version}/${node_archive}" \
--output "$node_archive_path"
printf '%s %s\n' "$node_archive_sha256" "$node_archive_path" | sha256sum --check --status
tar -xzf "$node_archive_path" -C "$node_runtime_root"
export PATH="${node_install_dir}/bin:${PATH}"
node --version
npm --version
export PYTHONPATH="$(pwd)/apps/dreamverse${PYTHONPATH:+:$PYTHONPATH}"
pytest apps/dreamverse/dreamverse/tests -q
cd apps/dreamverse/web
npm ci
npm run typecheck
npm test
machine_arch=$(uname -m)
if [[ $machine_arch =~ ^(aarch64|arm64)$ ]]; then
npx playwright install --with-deps chromium firefox
else
npx playwright install --with-deps chromium webkit firefox
fi
master_port=${MASTER_PORT:-7959}
BACKEND_PORT=${BACKEND_PORT:-$((master_port + 50))}
python -m uvicorn dreamverse.mock_server:app --host 127.0.0.1 --port "$BACKEND_PORT" &
mock_server_pid=$!
cleanup() {
kill "$mock_server_pid" 2>/dev/null || true
wait "$mock_server_pid" 2>/dev/null || true
}
trap cleanup EXIT INT TERM
for _ in {1..30}; do
curl -fsS "http://127.0.0.1:$BACKEND_PORT/healthz" && break
sleep 1
done
curl -fsS "http://127.0.0.1:$BACKEND_PORT/healthz"
if [[ $machine_arch =~ ^(aarch64|arm64)$ ]]; then
# Playwright WebKit traps before opening a page on Linux ARM64, and its
# bundled Chromium lacks the H.264/AAC codecs used by the fMP4 assertions.
# Firefox covers every flow, including streaming. Chromium and its mobile
# profile still cover all codec-independent UI behavior on GB200.
BACKEND_HOST=127.0.0.1 BACKEND_PORT="$BACKEND_PORT" CI=1 \
npm run e2e -- --project=firefox
BACKEND_HOST=127.0.0.1 BACKEND_PORT="$BACKEND_PORT" CI=1 \
npm run e2e -- \
--project=chromium \
--project=mobile-chromium \
--grep-invert='streams, plays, and surfaces a downloadable clip|starts a new project and switches back to the prior session|saved projects persist across a page reload'
else
BACKEND_HOST=127.0.0.1 BACKEND_PORT="$BACKEND_PORT" CI=1 \
npm run e2e -- \
--project=chromium \
--project=webkit \
--project=firefox \
--project=mobile-safari \
--project=mobile-chromium
fi
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the encoder lane.
set -euo pipefail
exec pytest ./fastvideo/tests/encoders -vs
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the evaluation lane.
set -euo pipefail
exec pytest ./fastvideo/tests/eval -vs
+35
View File
@@ -0,0 +1,35 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the golden-gate lane. Environment (HF_HOME
# and authentication) is the runner's responsibility.
set -euo pipefail
golden_root=./fastvideo/tests/golden_gate
selected=${FASTVIDEO_GOLDEN_TEST_FILES-}
if [ -z "$selected" ]; then
if [ "${TEST_SCOPE:-}" = merge ]; then
echo "Missing FASTVIDEO_GOLDEN_TEST_FILES for merge scope" >&2
exit 2
fi
selected=all
fi
if [ "$selected" = all ]; then
exec pytest "$golden_root" -vs
fi
[[ $selected =~ ^test_[a-z0-9_]+\.py(,test_[a-z0-9_]+\.py)*$ ]] || {
echo "Invalid FASTVIDEO_GOLDEN_TEST_FILES selection" >&2
exit 2
}
IFS=, read -r -a golden_files <<< "$selected"
golden_paths=()
for golden_file in "${golden_files[@]}"; do
golden_path="$golden_root/$golden_file"
[ -f "$golden_path" ] || {
echo "Selected golden test does not exist: $golden_file" >&2
exit 2
}
golden_paths+=("$golden_path")
done
exec pytest "${golden_paths[@]}" -vs
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the LoRA-inference lane.
set -euo pipefail
exec pytest ./fastvideo/tests/inference/lora/test_lora_inference_similarity.py -vs
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the VMoBA-inference lane.
set -euo pipefail
exec python fastvideo/tests/inference/vmoba/test_vmoba_inference.py
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the custom-kernel lane.
set -euo pipefail
exec pytest fastvideo-kernel/tests/ -vs
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the LoRA-extraction lane.
set -euo pipefail
exec pytest ./fastvideo/tests/lora_extraction/test_lora_extraction.py -vs
+52
View File
@@ -0,0 +1,52 @@
#!/usr/bin/env bash
# Canonical Slurm performance lane. Reports are written outside the checkout
# so the trusted host driver can upload them after untrusted code exits.
set -uo pipefail
export PERFORMANCE_TRACKING_ROOT=/tmp/perf-tracking
export PERF_REPORTS_DIR=/workspace/artifacts/performance
mkdir -p "$PERF_REPORTS_DIR"
if [[ ${BUILDKITE_PULL_REQUEST:-false} =~ ^[1-9][0-9]*$ ]]; then
export PERF_RUN_SOURCE=pr
export PERF_UPLOAD_POLICY=pass
elif [ "${BUILDKITE_BRANCH:-}" = main ] \
&& { [ "${BUILDKITE_SOURCE:-}" = schedule ] || [ "${TEST_SCOPE:-}" = full ]; }; then
export PERF_RUN_SOURCE=scheduled_main
export PERF_UPLOAD_POLICY=always
elif [ "${TEST_SCOPE:-}" = direct ]; then
export PERF_RUN_SOURCE=unknown
export PERF_UPLOAD_POLICY=pass
else
export PERF_RUN_SOURCE=unknown
export PERF_UPLOAD_POLICY=never
fi
nvidia-smi \
--query-gpu=index,timestamp,clocks.sm,clocks.max.sm,power.draw,power.limit,temperature.gpu \
--format=csv -l 10 > "$PERF_REPORTS_DIR/gpu_telemetry.csv" 2>/dev/null &
telemetry_pid=$!
cleanup() {
kill "$telemetry_pid" 2>/dev/null || true
wait "$telemetry_pid" 2>/dev/null || true
}
trap cleanup EXIT INT TERM
pytest ./fastvideo/tests/performance -vs
pytest_rc=$?
compare_rc=0
if [ "$pytest_rc" -eq 0 ] || [ "$PERF_UPLOAD_POLICY" = always ]; then
PERF_PYTEST_RC=$pytest_rc python ./fastvideo/tests/performance/compare_baseline.py
compare_rc=$?
fi
python ./fastvideo/tests/performance/dashboard.py || true
cp -f fastvideo/tests/performance/results/*.json "$PERF_REPORTS_DIR/" 2>/dev/null || true
echo "--- GPU telemetry (clocks.sm vs clocks.max.sm reveals capped hosts) ---"
cat "$PERF_REPORTS_DIR/gpu_telemetry.csv" || true
final_rc=$pytest_rc
if [ "$final_rc" -eq 0 ]; then
final_rc=$compare_rc
fi
exit "$final_rc"
+6
View File
@@ -0,0 +1,6 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the self-forcing lane.
set -euo pipefail
export WANDB_MODE=offline
exec pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs
+205
View File
@@ -0,0 +1,205 @@
#!/usr/bin/env bash
# Canonical four-GPU SSIM lane for the Slinky Slurm worker.
set -euo pipefail
args=()
if [ "${FASTVIDEO_SSIM_BOOTSTRAP_MODE:-0}" = 1 ]; then
args+=(--bootstrap-mode)
fi
selected=${FASTVIDEO_SSIM_TEST_FILES-}
if [ -z "$selected" ]; then
if [ "${TEST_SCOPE:-}" = merge ]; then
echo "Missing FASTVIDEO_SSIM_TEST_FILES for merge scope" >&2
exit 2
fi
selected=all
fi
if [ "$selected" != all ]; then
[[ $selected =~ ^test_[a-z0-9_]+\.py(,test_[a-z0-9_]+\.py)*$ ]] || {
echo "Invalid FASTVIDEO_SSIM_TEST_FILES selection" >&2
exit 2
}
IFS=, read -r -a ssim_files <<< "$selected"
for ssim_file in "${ssim_files[@]}"; do
args+=(--test-file "$ssim_file")
done
fi
# This disposable branch defaults to a focused GameCraft FA4 candidate run.
# The canonical lane remains available for local inspection by setting this to
# zero. Always exit 2 in candidate mode so the run cannot be mistaken for a
# passing quality gate and cannot trigger the production exit-1 retry policy.
candidate_enabled=${FASTVIDEO_GAMECRAFT_FA4_CANDIDATE_LOG:-1}
force_candidate_exit() {
trap - EXIT
exit 2
}
if [ "$candidate_enabled" = 1 ]; then
trap force_candidate_exit EXIT
fi
# MoGe's utils3d dependency builds glcontext from source on ARM64. The current
# runner image predates the baked-in X11 headers below, so keep this guarded
# bootstrap until every deployed image digest contains libx11-dev.
if [ ! -f /usr/include/X11/Xlib.h ]; then
apt-get -o Acquire::Retries=5 update
apt-get -o Acquire::Retries=5 install -y --no-install-recommends libx11-dev
rm -rf /var/lib/apt/lists/*
fi
uv pip install git+https://github.com/microsoft/MoGe.git
uv pip install k_diffusion einops_exts alias_free_torch torchsde
if [ "$candidate_enabled" = 1 ]; then
checkout_sha=$(git rev-parse HEAD)
echo "+++ GameCraft T2V current-FA4 candidate on GB200"
echo "checkout_sha=$checkout_sha"
echo "buildkite_commit=${BUILDKITE_COMMIT:-<unset>}"
if [ -n "${BUILDKITE_COMMIT:-}" ] && [ "$checkout_sha" != "$BUILDKITE_COMMIT" ]; then
echo "Checkout does not match BUILDKITE_COMMIT" >&2
exit 2
fi
visible_gpus=${CUDA_VISIBLE_DEVICES:-0}
candidate_gpu=${visible_gpus%%,*}
candidate_log=/tmp/fastvideo-gamecraft-fa4-candidate-pytest.log
candidate_sentinel=$(mktemp /tmp/fastvideo-gamecraft-fa4-candidate.XXXXXX)
candidate_generated_root=fastvideo/tests/ssim/generated_videos/default
candidate_env=(
"CUDA_VISIBLE_DEVICES=$candidate_gpu"
"PYTORCH_CUDA_ALLOC_CONF=expandable_segments:False"
"FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN"
"FASTVIDEO_FA4=1"
"FASTVIDEO_SSIM_BOOTSTRAP_MODE=0"
"FASTVIDEO_SSIM_FULL_QUALITY=0"
)
if env "${candidate_env[@]}" python - <<'PY'
import os
import torch
print(f"requested_backend={os.environ['FASTVIDEO_ATTENTION_BACKEND']}")
print(f"requested_FASTVIDEO_FA4={os.environ['FASTVIDEO_FA4']}")
print(f"torch_version={torch.__version__}")
print(f"torch_cuda_version={torch.version.cuda}")
print(f"cuda_available={torch.cuda.is_available()}")
print(f"cuda_visible_devices={os.environ.get('CUDA_VISIBLE_DEVICES', '<unset>')}")
print(f"gpu_count={torch.cuda.device_count()}")
if torch.cuda.device_count() != 1:
raise RuntimeError(f"Expected exactly one visible candidate GPU, got {torch.cuda.device_count()}")
gpu_name = torch.cuda.get_device_name(0)
print(f"gpu[0]={gpu_name} capability={torch.cuda.get_device_capability(0)}")
if "GB200" not in gpu_name:
raise RuntimeError(f"Candidate must run on GB200, got {gpu_name}")
from fastvideo.attention.selector import get_attn_backend
from fastvideo.attention.utils.flash_attn_default import fa_version
from fastvideo.platforms import AttentionBackendEnum
requested_backend = AttentionBackendEnum[os.environ["FASTVIDEO_ATTENTION_BACKEND"]]
resolved_backend = get_attn_backend(
128,
torch.bfloat16,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,),
requested=requested_backend,
)
print(f"resolved_backend={resolved_backend.get_name()}")
print(f"resolved_flash_attention=FA{fa_version}")
if fa_version != "4":
raise RuntimeError(f"Candidate must use FA4, resolved FA{fa_version}")
PY
then
probe_rc=0
else
probe_rc=$?
fi
if [ "$probe_rc" -ne 0 ]; then
echo "GameCraft candidate environment probe failed with rc=$probe_rc" >&2
exit 2
fi
if env "${candidate_env[@]}" python -m pytest \
fastvideo/tests/ssim/test_gamecraft_similarity.py::test_gamecraft_t2v_similarity \
-vs >"$candidate_log" 2>&1
then
test_rc=0
else
test_rc=$?
fi
echo "+++ GameCraft candidate pytest tail"
python - "$candidate_log" <<'PY'
import sys
from pathlib import Path
log_path = Path(sys.argv[1])
data = log_path.read_bytes()
print(f"pytest_log_bytes={len(data)}")
tail = data[-131072:].decode("utf-8", errors="replace").replace("\r", "\n")
lines = tail.splitlines()[-160:]
for line in lines:
print(line[:2000])
PY
echo "pytest_rc=$test_rc"
candidate_videos=()
if [ -d "$candidate_generated_root" ]; then
mapfile -d '' -t candidate_videos < <(
find "$candidate_generated_root" -type f -name '*.mp4' -newer "$candidate_sentinel" -print0
)
fi
if [ "${#candidate_videos[@]}" -ne 1 ]; then
echo "Expected exactly one newly generated candidate MP4; found ${#candidate_videos[@]}" >&2
printf 'candidate_path=%s\n' "${candidate_videos[@]}" >&2
exit 2
fi
candidate_video=${candidate_videos[0]}
case "$candidate_video" in
*/default/GB200_reference_videos/HunyuanGameCraft-T2V/FLASH_ATTN/*.mp4) ;;
*)
echo "Candidate path is outside the expected GB200 T2V subtree: $candidate_video" >&2
exit 2
;;
esac
candidate_results=()
mapfile -d '' -t candidate_results < <(
find "$(dirname "$candidate_video")" -type f -name '*_ssim.json' -newer "$candidate_sentinel" -print0
)
if [ "${#candidate_results[@]}" -ne 1 ]; then
echo "Expected exactly one newly generated SSIM JSON; found ${#candidate_results[@]}" >&2
exit 2
fi
candidate_result=${candidate_results[0]}
python - "$candidate_video" "$candidate_result" <<'PY'
import json
import sys
from pathlib import Path
video_path = Path(sys.argv[1]).resolve()
result_path = Path(sys.argv[2])
result = json.loads(result_path.read_text())
if Path(result["generated_video"]).resolve() != video_path:
raise RuntimeError("SSIM JSON does not describe the candidate MP4")
if result["parameters"]["num_inference_steps"] != 20:
raise RuntimeError("Candidate did not use the expected 20 inference steps")
print(f"candidate_path={video_path}")
print(f"candidate_ssim_mean={result['mean_ssim']}")
print(f"candidate_ssim_min={result['min_ssim']}")
print(f"candidate_ssim_max={result['max_ssim']}")
print(f"candidate_prompt={result['parameters']['prompt']}")
PY
echo "FV_GAMECRAFT_FA4_SSIM_JSON_BEGIN"
cat "$candidate_result"
echo "FV_GAMECRAFT_FA4_SSIM_JSON_END"
echo "+++ GameCraft candidate MP4 log envelope"
python .buildkite/scripts/gamecraft_candidate_log.py emit "$candidate_video"
echo "--- GameCraft candidate complete: probe_rc=$probe_rc pytest_rc=$test_rc (forced lane rc=2)"
exit 2
fi
exec python fastvideo/tests/ssim/ci_runner.py "${args[@]}"
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the modular training-framework lane.
set -euo pipefail
exec pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs
+6
View File
@@ -0,0 +1,6 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the legacy vanilla-training lane.
set -euo pipefail
export WANDB_MODE=offline
exec pytest ./fastvideo/tests/training/Vanilla -srP
+6
View File
@@ -0,0 +1,6 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the legacy LoRA-training lane.
set -euo pipefail
export WANDB_MODE=offline
exec pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP
+6
View File
@@ -0,0 +1,6 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the legacy VSA-training lane.
set -euo pipefail
export WANDB_MODE=offline
exec pytest ./fastvideo/tests/training/VSA -srP
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the transformer lane.
set -euo pipefail
exec pytest ./fastvideo/tests/transformers -vs
+5
View File
@@ -0,0 +1,5 @@
#!/usr/bin/env bash
# Canonical Slurm CI selection for the VAE lane.
set -euo pipefail
exec pytest ./fastvideo/tests/vaes -vs
+13
View File
@@ -1,6 +1,19 @@
#!/bin/bash
set -uo pipefail
# DORMANT ROLLBACK ONLY. Active CI is Slurm-only and pipeline.yml never calls
# this launcher. Refuse every Buildkite invocation even if a stale step or
# operator typo reaches this file; local rollback experiments require an
# explicit opt-in.
if [ -n "${BUILDKITE:-}" ]; then
echo "Legacy Modal CI is disabled; use the Slinky Slurm runner." >&2
exit 2
fi
if [ "${FASTVIDEO_ENABLE_LEGACY_MODAL_CI:-0}" != 1 ]; then
echo "Legacy Modal CI is dormant. Set FASTVIDEO_ENABLE_LEGACY_MODAL_CI=1 only for a manual rollback test." >&2
exit 2
fi
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
+25
View File
@@ -0,0 +1,25 @@
#!/usr/bin/env bash
set -euo pipefail
exec pytest \
./fastvideo/tests/api/ \
./fastvideo/tests/contract/ \
./fastvideo/tests/dataset/ \
./fastvideo/tests/workflow/ \
./fastvideo/tests/entrypoints/ \
./fastvideo/tests/loader/ \
./fastvideo/tests/pipelines/ \
./fastvideo/tests/platforms/ \
./fastvideo/tests/train/ \
./fastvideo/tests/stages/ \
./fastvideo/tests/ops/ \
./fastvideo/tests/worker/ \
./fastvideo/tests/training/test_trackers.py \
./fastvideo/tests/attention/test_sdpa_metadata_mask_contract.py \
./fastvideo/tests/modal/test_kernel_build_cache.py \
./fastvideo/tests/modal/test_pr_test.py \
./fastvideo/tests/modal/test_ssim_test.py \
--ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py \
--ignore=./fastvideo/tests/train/models \
--ignore=./fastvideo/tests/train/methods \
-vs
+2 -2
View File
@@ -8,10 +8,10 @@ PR TITLE: Must start with a type tag, e.g.:
MERGE WORKFLOW:
1. Ensure pre-commit passes and you have at least 1 approval
2. Comment /merge (or add the "ready" label) to enter the Merge Queue
3. Full Test Suite runs automatically on a staging branch → auto-merge on success
3. A path-aware merge gate runs only relevant integration tests → auto-merge on success
ON-DEMAND TESTING (write access required):
/test full — Full Test Suite /test ssim — SSIM regression
/test full — Explicit all-lane run /test ssim — Full SSIM regression
/test training — Training pipeline /test encoder — Encoder tests
/test transformer — Transformer tests /test vae — VAE tests
/test kernel — CUDA kernel tests /test unit — Unit tests
+10 -10
View File
@@ -1,14 +1,14 @@
#!/usr/bin/env bash
# Gate the expensive Buildkite full suite on the cheap GitHub checks.
# Gate the path-aware Buildkite merge plan on the cheap GitHub checks.
#
# Polls the workflow runs for the PR head commit and only exits 0 once the
# watched cheap workflows (pre-commit, docs build) have succeeded, so the
# 'ready' label cannot burn ~20 GPU lanes on a head that a cheap check has
# already doomed.
# 'ready' label cannot burn path-selected GPU lanes on a head that a cheap
# check has already doomed.
#
# Semantics:
# - watched run completed with a bad conclusion -> exit 1 (fail CLOSED:
# no full suite; the next push re-arms via the 'synchronize' trigger)
# no merge gate; the next push re-arms via the 'synchronize' trigger)
# - watched run cancelled -> still pending: the docs
# workflow's repo-global 'pages' concurrency group cancels runs superseded
# by unrelated pushes, so 'cancelled' is not a verdict on this PR
@@ -29,7 +29,7 @@ set -euo pipefail
: "${PR_NUMBER:?PR_NUMBER (pull request number) is required}"
: "${GITHUB_REPOSITORY:?GITHUB_REPOSITORY is required}"
# Workflow-level `name:` values that must be green before the full suite
# Workflow-level `name:` values that must be green before the merge gate
# may start. "Deploy Documentation" is path-filtered on PRs, so its run may
# legitimately never exist; pre-commit always runs, so it must appear.
WATCHED_NAMES='["pre-commit", "Deploy Documentation"]'
@@ -56,7 +56,7 @@ recheck_ready_label() {
if pr_json=$(gh_api "repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" 2>/dev/null); then
if ! jq -e '[.labels[]?.name] | index("ready")' <<<"$pr_json" >/dev/null 2>&1; then
echo "::error::PR #${PR_NUMBER} no longer has the 'ready' label —" \
"NOT triggering the Buildkite full suite. Re-add the label to re-arm."
"NOT triggering the Buildkite merge gate. Re-add the label to re-arm."
exit 1
fi
else
@@ -84,7 +84,7 @@ while true; do
| map(.name) | join(", ")' <<<"$state")
if [ -n "$failed" ]; then
echo "::error::Cheap check(s) failed on ${PR_SHA}: ${failed}." \
"NOT triggering the Buildkite full suite. Push a fix (the 'ready'" \
"NOT triggering the Buildkite merge gate. Push a fix (the 'ready'" \
"label re-arms on every push), or re-run the failed check and then" \
"re-run this workflow."
exit 1
@@ -97,7 +97,7 @@ while true; do
if [ "$pending" -eq 0 ]; then
if [ -z "$missing" ]; then
recheck_ready_label
echo "All watched cheap checks are green — full suite may proceed."
echo "All watched cheap checks are green — merge gate may proceed."
exit 0
fi
case "$missing" in
@@ -119,14 +119,14 @@ while true; do
echo "::warning::GitHub API error querying workflow runs for ${PR_SHA} (attempt ${api_fails}/3)."
if [ "$api_fails" -ge 3 ]; then
recheck_ready_label
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the full suite WITHOUT the cheap-check gate."
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the merge gate WITHOUT the cheap-check gate."
exit 0
fi
fi
if [ "$elapsed" -ge "$MAX_WAIT_SECS" ]; then
recheck_ready_label
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the full suite anyway."
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the merge gate anyway."
exit 0
fi
sleep "$POLL_SECS"
+570
View File
@@ -0,0 +1,570 @@
#!/usr/bin/env python3
"""Select the additive GPU integration lanes needed by a PR diff.
Fastcheck is the universal six-lane baseline and is intentionally not repeated
here. This planner selects only the more expensive merge-gate lanes. Unknown
source/build paths fail closed to the complete integration set, while explicit
documentation and repository-metadata paths require no additional GPU work.
"""
from __future__ import annotations
import argparse
import fnmatch
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import TextIO
MERGE_LANES = (
"golden-gate",
"ssim",
"lora-inference",
"lora-extraction",
"training",
"distillation",
"self-forcing",
"lora-training",
"training-vsa",
"inference-vmoba",
"performance",
"api-server",
"train-framework",
"eval",
)
LANE_SCRIPT_TO_KEY = {
"api_server.sh": "api-server",
"distillation_dmd.sh": "distillation",
"eval.sh": "eval",
"golden_gate.sh": "golden-gate",
"inference_lora.sh": "lora-inference",
"inference_vmoba.sh": "inference-vmoba",
"lora_extraction.sh": "lora-extraction",
"performance.sh": "performance",
"self_forcing.sh": "self-forcing",
"ssim.sh": "ssim",
"train_framework.sh": "train-framework",
"training.sh": "training",
"training_lora.sh": "lora-training",
"training_vsa.sh": "training-vsa",
}
FASTCHECK_LANE_SCRIPTS = {
"dreamverse.sh",
"encoder.sh",
"kernel_tests.sh",
"transformer.sh",
"vae.sh",
}
LEGACY_TRAINING_LANES = (
"training",
"distillation",
"self-forcing",
"lora-training",
"training-vsa",
)
ALL_TRAINING_LANES = (*LEGACY_TRAINING_LANES, "train-framework")
SSIM_SMOKE_TESTS = (
"test_flux_t2i_similarity.py",
"test_wan_t2v_similarity.py",
)
SAFE_PATTERNS = (
"*.md",
"*.rst",
".agents/**",
".claude/**",
".codex/**",
".github/ISSUE_TEMPLATE/**",
".github/PULL_REQUEST_TEMPLATE.md",
".github/dependabot.yml",
".github/mergify.yml",
".github/scripts/**",
".github/workflows/**",
".buildkite/scripts/pre_commit.sh",
".git-blame-ignore-revs",
".gitattributes",
".gitignore",
".pre-commit-config.yaml",
"AGENTS.md",
"CITATION.cff",
"CODE_OF_CONDUCT.md",
"CONTRIBUTING.md",
"LICENSE",
"NOTICE",
"__init__.py",
"collect_env.py",
"SECURITY.md",
"assets/**",
"comfyui/**",
"docs/**",
"examples/**",
"mkdocs.yml",
"requirements-mkdocs.in",
"requirements-mkdocs.txt",
"scripts/**",
"tests/__init__.py",
"tests/local_tests/**",
)
ALL_IMPACT_PATTERNS = (
".buildkite/pipeline.yml",
"docker/**",
"pyproject.toml",
"requirements*.txt",
"setup.cfg",
"setup.py",
"uv.lock",
)
@dataclass(frozen=True)
class FamilyCoverage:
pattern: re.Pattern[str]
golden_tests: tuple[str, ...]
ssim_tests: tuple[str, ...]
FAMILY_COVERAGE = (
FamilyCoverage(
re.compile(r"(^|[/_.-])dreamx(_world)?([/_.-]|$)"),
("test_dreamx.py", ),
("test_dreamx_world_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])flux[_-]?2([/_.-]|$)"),
("test_flux2_klein.py", ),
("test_flux2_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])flux(?![_-]?2)([/_.-]|$)"),
("test_flux.py", ),
("test_flux_t2i_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])(hunyuan)?gamecraft([/_.-]|$)"),
("test_gamecraft.py", ),
("test_gamecraft_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])gen3c([/_.-]|$)"),
("test_gen3c.py", ),
("test_gen3c_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])glm[_-]?image([/_.-]|$)"),
("test_glm_image.py", ),
("test_glm_image_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])kandinsky[_-]?5([/_.-]|$)"),
("test_kandinsky5.py", ),
("test_kandinsky5_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])lingbot([a-z0-9_-]*)([/_.-]|$)"),
("test_lingbot.py", ),
("test_lingbot_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])longcat([/_.-]|$)"),
("test_longcat.py", ),
("test_longcat_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])ltx[_-]?2([/_.-]|$)"),
("test_ltx2.py", ),
("test_ltx2_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])matrixgame[_-]?2([/_.-]|$)"),
("test_matrixgame.py", ),
("test_matrixgame2_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])matrixgame[_-]?3([/_.-]|$)"),
("test_matrixgame.py", ),
("test_matrixgame3_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])minimax[_-]?h3([/_.-]|$)"),
("test_minimax_h3_t2v.py", ),
("test_minimax_h3_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])sd[_-]?3([._-]?5)?([/_.-]|$)"),
("test_sd35.py", ),
("test_sd35_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])stable[_-]?audio([/_.-]|$)"),
("test_stable_audio.py", ),
("test_stable_audio_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])turbo(diffusion)?([/_.-]|$)"),
(),
("test_turbodiffusion_similarity.py", ),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])wan(video)?([/_.-]|$)"),
("test_wan_t2v.py", ),
(
"test_causal_similarity.py",
"test_wan_i2v_similarity.py",
"test_wan_t2v_similarity.py",
),
),
FamilyCoverage(
re.compile(r"(^|[/_.-])z[_-]?image([/_.-]|$)"),
("test_zimage.py", ),
("test_zimage_similarity.py", ),
),
)
@dataclass
class MergePlan:
lanes: set[str] = field(default_factory=set)
golden_tests: set[str] = field(default_factory=set)
ssim_tests: set[str] = field(default_factory=set)
golden_all: bool = False
ssim_all: bool = False
reasons: list[str] = field(default_factory=list)
def add_lanes(self, *lanes: str, reason: str) -> None:
unknown = set(lanes) - set(MERGE_LANES)
if unknown:
raise ValueError(f"Unknown merge lanes: {sorted(unknown)}")
self.lanes.update(lanes)
self.reasons.append(reason)
def add_golden(self, tests: tuple[str, ...], reason: str) -> None:
self.add_lanes("golden-gate", reason=reason)
self.golden_tests.update(tests)
def add_ssim(self, tests: tuple[str, ...], reason: str) -> None:
self.add_lanes("ssim", reason=reason)
self.ssim_tests.update(tests)
def require_all(self, reason: str) -> None:
self.lanes.update(MERGE_LANES)
self.golden_all = True
self.ssim_all = True
self.reasons.append(reason)
def ordered_lanes(self) -> tuple[str, ...]:
return tuple(lane for lane in MERGE_LANES if lane in self.lanes)
def encoded_lanes(self) -> str:
lanes = self.ordered_lanes()
return "," + ",".join(lanes or ("none", )) + ","
def encoded_golden_tests(self) -> str:
if "golden-gate" not in self.lanes:
return "none"
if self.golden_all or not self.golden_tests:
return "all"
return ",".join(sorted(self.golden_tests))
def encoded_ssim_tests(self) -> str:
if "ssim" not in self.lanes:
return "none"
if self.ssim_all or not self.ssim_tests:
return "all"
return ",".join(sorted(self.ssim_tests))
def _matches_any(path: str, patterns: tuple[str, ...]) -> bool:
return any(fnmatch.fnmatchcase(path, pattern) for pattern in patterns)
def _family_coverage(path: str) -> tuple[set[str], set[str]]:
normalized = path.lower()
golden: set[str] = set()
ssim: set[str] = set()
for family in FAMILY_COVERAGE:
if family.pattern.search(normalized):
golden.update(family.golden_tests)
ssim.update(family.ssim_tests)
return golden, ssim
def _select_output_coverage(plan: MergePlan, path: str) -> None:
golden, ssim = _family_coverage(path)
if golden:
plan.add_golden(tuple(sorted(golden)), reason=f"model-family golden coverage: {path}")
else:
plan.golden_all = True
plan.add_lanes("golden-gate", reason=f"shared output golden coverage: {path}")
if ssim:
plan.add_ssim(tuple(sorted(ssim)), reason=f"model-family SSIM coverage: {path}")
else:
plan.add_ssim(SSIM_SMOKE_TESTS, reason=f"shared output SSIM smoke coverage: {path}")
def classify_paths(paths: list[str]) -> MergePlan:
plan = MergePlan()
normalized_paths: list[str] = []
for raw_path in paths:
path = raw_path.strip()
while path.startswith("./"):
path = path[2:]
if path:
normalized_paths.append(path)
normalized_paths = sorted(set(normalized_paths))
if not normalized_paths:
plan.require_all("changed-file list was empty; failing closed")
return plan
for path in normalized_paths:
if path == "__FASTVIDEO_CI_PLAN_ALL__":
plan.require_all("changed-file API failed; failing closed")
continue
if path in {"requirements-mkdocs.in", "requirements-mkdocs.txt"}:
plan.reasons.append(f"documentation dependencies need no GPU integration: {path}")
continue
if _matches_any(path, ALL_IMPACT_PATTERNS):
plan.require_all(f"cross-cutting build/runtime surface: {path}")
continue
lane_script_prefix = ".buildkite/scripts/lanes/"
if path.startswith(lane_script_prefix):
script_name = Path(path).name
lane = LANE_SCRIPT_TO_KEY.get(script_name)
if lane is None:
if script_name in FASTCHECK_LANE_SCRIPTS:
plan.reasons.append(f"covered by automatic Fastcheck lane: {path}")
else:
plan.require_all(f"unknown lane script: {path}")
elif lane == "golden-gate":
plan.golden_all = True
plan.add_lanes(lane, reason=f"golden lane implementation: {path}")
elif lane == "ssim":
plan.ssim_all = True
plan.add_lanes(lane, reason=f"SSIM lane implementation: {path}")
else:
plan.add_lanes(lane, reason=f"lane implementation: {path}")
continue
if path.startswith("fastvideo/tests/golden_gate/"):
name = Path(path).name
if name.startswith("test_") and name.endswith(".py"):
plan.add_golden((name, ), reason=f"changed golden test: {path}")
elif name in {"AGENTS.md", "README.md"}:
plan.reasons.append(f"golden documentation only: {path}")
else:
plan.golden_all = True
plan.add_lanes("golden-gate", reason=f"shared golden harness/reference: {path}")
continue
if path.startswith("fastvideo/tests/ssim/"):
name = Path(path).name
if name.startswith("test_") and name.endswith(".py"):
plan.add_ssim((name, ), reason=f"changed SSIM test: {path}")
elif path.endswith((".py", ".json", ".pt", ".png", ".mp4")):
plan.ssim_all = True
plan.add_lanes("ssim", reason=f"shared SSIM harness/reference: {path}")
continue
if path.startswith("fastvideo/tests/performance/") or path.startswith(".buildkite/performance-benchmarks/"):
plan.add_lanes("performance", reason=f"performance coverage: {path}")
continue
if path.startswith(("fastvideo/performance/", "fastvideo/performance_dashboard/",
"apps/performance_dashboard/")):
plan.add_lanes("performance", reason=f"performance implementation: {path}")
continue
if path.startswith("fastvideo/benchmarks/"):
if "/mlx_" in path or Path(path).name.startswith("mlx_"):
plan.reasons.append(f"covered by the path-filtered macOS MLX workflow: {path}")
else:
plan.add_lanes("performance", reason=f"benchmark implementation: {path}")
continue
if path.startswith("fastvideo/tests/eval/") or path.startswith("fastvideo/eval/"):
plan.add_lanes("eval", reason=f"evaluation coverage: {path}")
continue
if path.startswith("fastvideo/third_party/eval/"):
plan.add_lanes("eval", reason=f"vendored evaluation implementation: {path}")
continue
if path.startswith("fastvideo/tests/lora_extraction/") or path.startswith("scripts/lora_extraction/"):
plan.add_lanes("lora-extraction", reason=f"LoRA extraction coverage: {path}")
continue
if path.startswith("fastvideo/tests/inference/lora/"):
plan.add_lanes("lora-inference", reason=f"LoRA inference coverage: {path}")
continue
if path.startswith("fastvideo/tests/inference/vmoba/"):
plan.add_lanes("inference-vmoba", reason=f"VMoBA inference coverage: {path}")
continue
if path.startswith(("fastvideo/dataset/", "fastvideo/workflow/", "fastvideo/pipelines/preprocess/",
"fastvideo/pipelines/training/")):
plan.add_lanes(*ALL_TRAINING_LANES, reason=f"shared data/training input surface: {path}")
continue
if path.startswith("fastvideo/tests/train/") or path.startswith("fastvideo/train/"):
plan.add_lanes("train-framework", reason=f"modular training coverage: {path}")
continue
if path.startswith("fastvideo/tests/training/"):
lowered = path.lower()
if "/vanilla/" in lowered:
plan.add_lanes("training", reason=f"vanilla training coverage: {path}")
elif "/distill/" in lowered:
plan.add_lanes("distillation", reason=f"distillation coverage: {path}")
elif "/self-forcing/" in lowered:
plan.add_lanes("self-forcing", reason=f"self-forcing coverage: {path}")
elif "/lora/" in lowered:
plan.add_lanes("lora-training", reason=f"LoRA training coverage: {path}")
elif "/vsa/" in lowered:
plan.add_lanes("training-vsa", reason=f"VSA training coverage: {path}")
else:
plan.add_lanes(*LEGACY_TRAINING_LANES, reason=f"shared legacy training coverage: {path}")
continue
if path.startswith("fastvideo/training/"):
lowered = path.lower()
if "self_forcing" in lowered:
plan.add_lanes("self-forcing", reason=f"self-forcing implementation: {path}")
elif "distill" in lowered:
plan.add_lanes("distillation", reason=f"distillation implementation: {path}")
elif "lora" in lowered:
plan.add_lanes("lora-training", reason=f"LoRA training implementation: {path}")
else:
plan.add_lanes(*LEGACY_TRAINING_LANES, reason=f"shared legacy training implementation: {path}")
continue
lowered = path.lower()
if "vmoba" in lowered and path.startswith(("fastvideo/", ".buildkite/")):
plan.add_lanes("inference-vmoba", reason=f"VMoBA implementation: {path}")
plan.add_golden(("test_wan_t2v.py", ), reason=f"VMoBA end-to-end coverage: {path}")
continue
if "lora" in lowered and path.startswith("fastvideo/"):
plan.add_lanes(
"lora-inference",
"lora-extraction",
"lora-training",
reason=f"shared LoRA implementation: {path}",
)
_select_output_coverage(plan, path)
continue
if path.startswith("fastvideo/entrypoints/") or path.startswith("fastvideo/api/"):
plan.add_lanes("api-server", reason=f"API/entrypoint integration: {path}")
if "openai" not in lowered and "/cli/" not in lowered:
_select_output_coverage(plan, path)
continue
if path.startswith("fastvideo/worker/"):
plan.add_lanes("api-server", reason=f"worker/API integration: {path}")
_select_output_coverage(plan, path)
continue
if path.startswith("fastvideo/distributed/"):
plan.add_lanes(
"training",
"train-framework",
reason=f"distributed runtime integration: {path}",
)
_select_output_coverage(plan, path)
continue
if path.startswith(("fastvideo/hooks/", "fastvideo/platforms/", "fastvideo/third_party/")):
_select_output_coverage(plan, path)
continue
if path.startswith(("fastvideo/models/", "fastvideo/pipelines/", "fastvideo/configs/",
"fastvideo/layers/", "fastvideo/attention/")):
_select_output_coverage(plan, path)
continue
if path in {
"fastvideo/fastvideo_args.py",
"fastvideo/forward_context.py",
"fastvideo/image_processor.py",
"fastvideo/registry.py",
"fastvideo/utils.py",
}:
_select_output_coverage(plan, path)
continue
if path.startswith("fastvideo/mlx_runtime/"):
plan.reasons.append(f"covered by the path-filtered macOS MLX workflow: {path}")
continue
if path.startswith("fastvideo/logging_utils/") or path in {
"fastvideo/__init__.py",
"fastvideo/envs.py",
"fastvideo/logger.py",
"fastvideo/profiler.py",
"fastvideo/version.py",
}:
plan.reasons.append(f"covered by automatic Fastcheck: {path}")
continue
if path.startswith(("fastvideo-kernel/", "csrc/")):
plan.add_golden(("test_wan_t2v.py", ), reason=f"kernel integration smoke: {path}")
plan.add_ssim(("test_wan_t2v_similarity.py", ), reason=f"kernel numerical smoke: {path}")
continue
if path.startswith("apps/dreamverse/"):
# DreamVerse is already one of the six automatic Fastcheck lanes.
plan.reasons.append(f"covered by automatic DreamVerse Fastcheck: {path}")
continue
if path.startswith("fastvideo/tests/"):
# The automatic unit/component Fastcheck lanes own the remaining
# package tests. Domain-specific expensive test roots were handled
# above.
plan.reasons.append(f"covered by automatic Fastcheck: {path}")
continue
if path in {".buildkite/scripts/unit_test.sh", ".buildkite/scripts/pr_test.sh"}:
plan.reasons.append(f"covered by automatic unit Fastcheck: {path}")
continue
if _matches_any(path, SAFE_PATTERNS):
plan.reasons.append(f"no additional GPU integration needed: {path}")
continue
plan.require_all(f"unclassified path; failing closed: {path}")
return plan
def _write_github_output(output: TextIO, plan: MergePlan) -> None:
output.write(f"merge_test_plan={plan.encoded_lanes()}\n")
output.write(f"merge_golden_tests={plan.encoded_golden_tests()}\n")
output.write(f"merge_ssim_tests={plan.encoded_ssim_tests()}\n")
output.write(f"merge_plan_label={','.join(plan.ordered_lanes()) or 'none'}\n")
def _write_summary(output: TextIO, plan: MergePlan) -> None:
output.write("## Change-aware merge test plan\n\n")
output.write("| Selection | Value |\n|---|---|\n")
output.write(f"| Additional Slurm lanes | `{','.join(plan.ordered_lanes()) or 'none'}` |\n")
output.write(f"| Golden tests | `{plan.encoded_golden_tests()}` |\n")
output.write(f"| SSIM tests | `{plan.encoded_ssim_tests()}` |\n\n")
output.write("Fastcheck remains the universal six-lane baseline.\n")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--paths-file", type=Path, required=True)
parser.add_argument("--github-output", type=Path)
parser.add_argument("--summary-file", type=Path)
return parser.parse_args()
def main() -> int:
args = parse_args()
paths = args.paths_file.read_text(encoding="utf-8").splitlines()
plan = classify_paths(paths)
print(f"MERGE_TEST_PLAN={plan.encoded_lanes()}")
print(f"MERGE_GOLDEN_TESTS={plan.encoded_golden_tests()}")
print(f"MERGE_SSIM_TESTS={plan.encoded_ssim_tests()}")
for reason in plan.reasons:
print(f"- {reason}")
if args.github_output:
with args.github_output.open("a", encoding="utf-8") as output:
_write_github_output(output, plan)
if args.summary_file:
with args.summary_file.open("a", encoding="utf-8") as output:
_write_summary(output, plan)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+1 -1
View File
@@ -53,7 +53,7 @@ PC_PENDING='{"name": "pre-commit", "id": 1, "status": "in_progress", "conclusion
DOCS_OK='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "success"}'
DOCS_BAD='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "failure"}'
DOCS_CANCELLED='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "cancelled"}'
OTHER='{"name": "Trigger Full Suite", "id": 3, "status": "in_progress", "conclusion": null}'
OTHER='{"name": "Trigger Merge Gate", "id": 3, "status": "in_progress", "conclusion": null}'
NULL_NAME='{"name": null, "id": 4, "status": "completed", "conclusion": "failure"}'
PC_OK_RERUN='{"name": "pre-commit", "id": 5, "status": "completed", "conclusion": "success"}'
@@ -190,6 +190,7 @@ jobs:
if: ${{ !inputs.push_by_digest }}
run: |
echo "✅ Python ${{ inputs.python_version }} image successfully built and pushed to ${{ steps.image.outputs.name }}:${{ inputs.tag_suffix }}-sha-${GITHUB_SHA::7}"
echo "Digest: ${{ steps.build-push.outputs.digest }}"
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
- name: Digest success message
+39 -25
View File
@@ -26,29 +26,48 @@ jobs:
per_page: 100,
});
const bkStatuses = data.statuses.filter(
s => s.context.startsWith('buildkite/ci/')
);
const FASTCHECK_PREFIX = 'buildkite/ci/microscope-';
// Buildkite derives the GitHub context prefix from the label emoji.
// Keep hard Full Suite lanes in test-tube/bar-chart namespaces and
// Fastcheck lanes in microscope so targeted reruns cannot clear the
// wrong aggregate status. Automatic PR jobs use pr-fastcheck while
// slash-command and Full Suite jobs use ci; normalize the suffix
// and keep the newest status for each logical lane.
const FASTCHECK_PREFIXES = [
'buildkite/pr-fastcheck/microscope-',
'buildkite/ci/microscope-',
];
const FULL_SUITE_PREFIXES = [
'buildkite/ci/test-tube-',
'buildkite/ci/bar-chart-',
];
const fastcheck = bkStatuses.filter(
s => s.context.startsWith(FASTCHECK_PREFIX)
);
const fullSuite = bkStatuses.filter(
s => FULL_SUITE_PREFIXES.some(p => s.context.startsWith(p))
);
function newestByLane(prefixes) {
const statuses = new Map();
for (const status of data.statuses) {
const prefix = prefixes.find(p => status.context.startsWith(p));
if (!prefix) continue;
const lane = status.context.slice(prefix.length);
const previous = statuses.get(lane);
if (!previous || Date.parse(status.updated_at) > Date.parse(previous.updated_at)) {
statuses.set(lane, status);
}
}
return statuses;
}
if (
fastcheck.length > 0
&& fastcheck.every(s => s.state === 'success')
) {
const fastcheck = newestByLane(FASTCHECK_PREFIXES);
const fullSuiteOnly = newestByLane(FULL_SUITE_PREFIXES);
const fastcheckPassed =
fastcheck.size === 6
&& [...fastcheck.values()].every(s => s.state === 'success');
const fullSuitePassed =
fastcheckPassed
&& fullSuiteOnly.size === 14
&& [...fullSuiteOnly.values()].every(s => s.state === 'success');
if (fastcheckPassed) {
core.info(
`All ${fastcheck.length} fastcheck tests passed — updating fastcheck-passed`
`All ${fastcheck.size} fastcheck tests passed — updating fastcheck-passed`
);
await github.rest.repos.createCommitStatus({
owner: context.repo.owner,
@@ -56,17 +75,13 @@ jobs:
sha,
state: 'success',
context: 'fastcheck-passed',
description:
`All ${fastcheck.length} fastcheck tests passed`,
description: `All ${fastcheck.size} fastcheck tests passed`,
});
}
if (
fullSuite.length > 0
&& fullSuite.every(s => s.state === 'success')
) {
if (fullSuitePassed) {
core.info(
`All ${fullSuite.length} full suite tests passed — updating full-suite-passed`
'All 20 full suite tests passed — updating full-suite-passed'
);
await github.rest.repos.createCommitStatus({
owner: context.repo.owner,
@@ -74,7 +89,6 @@ jobs:
sha,
state: 'success',
context: 'full-suite-passed',
description:
`All ${fullSuite.length} full suite tests passed`,
description: 'All 20 full suite tests passed',
});
}
+2
View File
@@ -78,6 +78,7 @@ jobs:
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_checkpoint_compat.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
@@ -134,6 +135,7 @@ jobs:
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_checkpoint_compat.py \
fastvideo/tests/mlx/test_mlx_fastwan_benchmark.py \
fastvideo/tests/mlx/test_taehv_decode.py \
fastvideo/tests/mlx/test_frame_upsample.py \
+44
View File
@@ -0,0 +1,44 @@
name: Scheduled Full SSIM
on:
schedule:
- cron: "0 5 * * 0"
workflow_dispatch:
permissions:
contents: read
jobs:
trigger:
if: github.repository == 'hao-ai-lab/FastVideo'
runs-on: ubuntu-latest
steps:
- name: Trigger weekly full SSIM on Slinky Slurm
env:
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
SOURCE_SHA: ${{ github.sha }}
SOURCE_BRANCH: ${{ github.event.repository.default_branch }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
set -euo pipefail
curl -sS --fail-with-body -X POST \
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds" \
-H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
-H "Content-Type: application/json" \
--data-raw "$(jq -n \
--arg commit "$SOURCE_SHA" \
--arg branch "$SOURCE_BRANCH" \
'{
commit: $commit,
branch: $branch,
message: "Weekly full SSIM on Slinky Slurm",
ignore_pipeline_branch_filters: true,
env: {
TEST_SCOPE: "scheduled",
FULL_SUITE: "false",
TEST_TYPE: "ssim",
PR_NUMBER: "false",
PR_TITLE: "Scheduled full SSIM"
}
}')"
+12 -44
View File
@@ -33,7 +33,6 @@ jobs:
core.setOutput('has_write', String(hasWrite));
- name: Add ready label and react
id: label
if: steps.perm.outputs.has_write == 'true'
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
with:
@@ -48,47 +47,6 @@ jobs:
comment_id: context.payload.comment.id,
content: 'rocket',
});
const { data: pr } = await github.rest.pulls.get({ owner, repo, pull_number: prNumber });
core.setOutput('pr_sha', pr.head.sha);
core.setOutput('pr_branch', pr.head.ref);
core.setOutput('pr_number', String(prNumber));
core.setOutput('pr_title', pr.title);
- name: Trigger Full Suite
if: steps.perm.outputs.has_write == 'true'
env:
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
PR_SHA: ${{ steps.label.outputs.pr_sha }}
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
PR_TITLE: ${{ steps.label.outputs.pr_title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
curl -sS --fail-with-body -X POST \
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds" \
-H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
-H "Content-Type: application/json" \
--data-raw "$(jq -n \
--arg commit "$PR_SHA" \
--arg branch "$PR_BRANCH" \
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
--arg pr_title "$PR_TITLE" \
--argjson pr_id "$PR_NUMBER" \
'{
commit: $commit,
branch: $branch,
message: $message,
ignore_pipeline_branch_filters: true,
pull_request_id: $pr_id,
pull_request_base_branch: "main",
env: {
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring),
PR_TITLE: $pr_title
}
}')"
parse-command:
if: >-
@@ -129,7 +87,7 @@ jobs:
set -euo pipefail
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
VALID="encoder vae transformer kernel unit dreamverse ssim golden-gate training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
VALID="encoder vae transformer kernel unit dreamverse ssim golden-gate training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval unit-ci kernel-ci dreamverse-ci ssim-ci golden-gate-ci encoder-ci vae-ci transformer-ci lora-inference-ci lora-training-ci lora-extraction-ci training-ci distillation-ci self-forcing-ci vsa-ci vmoba-ci performance-ci api-ci train-framework-ci eval-ci full fastcheck pre-commit"
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
exit 1
@@ -137,7 +95,17 @@ jobs:
declare -A MAP=(
[encoder]=encoder [vae]=vae [transformer]=transformer
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
[kernel]=kernel_tests [unit]=unit_test [unit-ci]=unit_test_ci
[kernel-ci]=kernel_tests_ci [dreamverse-ci]=dreamverse_app_ci
[ssim-ci]=ssim_ci [vmoba-ci]=inference_vmoba_ci
[golden-gate-ci]=golden_gate_ci [training-ci]=training_ci
[encoder-ci]=encoder_ci [vae-ci]=vae_ci [transformer-ci]=transformer_ci
[lora-inference-ci]=inference_lora_ci [lora-training-ci]=training_lora_ci
[lora-extraction-ci]=lora_extraction_ci [distillation-ci]=distillation_dmd_ci
[self-forcing-ci]=self_forcing_ci [vsa-ci]=training_vsa_ci
[performance-ci]=performance_ci [api-ci]=api_server_ci
[train-framework-ci]=train_framework_ci [eval-ci]=eval_ci
[dreamverse]=dreamverse_app
[ssim]=ssim [golden-gate]=golden_gate [training]=training
[lora-inference]=inference_lora [lora-training]=training_lora
[lora-extraction]=lora_extraction
+66 -13
View File
@@ -1,4 +1,4 @@
name: Trigger Full Suite
name: Trigger Merge Gate
on:
pull_request_target:
@@ -10,7 +10,7 @@ permissions:
actions: read
concurrency:
group: full-suite-${{ github.event.pull_request.number }}
group: merge-gate-${{ github.event.pull_request.number }}
cancel-in-progress: false
jobs:
@@ -34,29 +34,72 @@ jobs:
});
const hasReady = pr.labels.some(l => l.name === 'ready');
core.setOutput('has_ready', String(hasReady));
if (!hasReady) core.info('No ready label — skipping Full Suite trigger.');
core.setOutput('changed_files', String(pr.changed_files));
if (!hasReady) core.info('No ready label — skipping merge-gate trigger.');
- name: Cancel previous Buildkite builds
if: steps.check.outputs.has_ready == 'true'
env:
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
PR_NUMBER: ${{ github.event.pull_request.number }}
run: |
# Find running builds for this branch with TEST_SCOPE=full and cancel them
builds=$(curl -sS -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds?branch=${PR_BRANCH}&state=running,scheduled" \
| jq -r '.[] | select(try (.env.TEST_SCOPE == "full") catch false) | .number')
# Match both branch and PR number: forks can reuse the same branch name.
builds=$(curl -sS --get -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
--data-urlencode "branch=$PR_BRANCH" \
--data-urlencode "state=running,scheduled" \
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds" \
| jq -r --arg pr_number "$PR_NUMBER" \
'.[] | select((.env.TEST_SCOPE? == "merge") and (.env.PR_NUMBER? == $pr_number)) | .number')
for build_num in $builds; do
echo "Cancelling Buildkite build #$build_num"
curl -sS -X PUT -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
done
# Checks out the BASE branch (default for pull_request_target), so PR
# authors cannot tamper with the gate script.
- name: Checkout gate script
# Check out the immutable BASE SHA: pull_request_target must never run a
# planner or gate script from the untrusted PR head.
- name: Checkout trusted merge planner
if: steps.check.outputs.has_ready == 'true'
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
with:
ref: ${{ github.event.pull_request.base.sha }}
persist-credentials: false
- name: Collect changed paths
if: steps.check.outputs.has_ready == 'true'
env:
GH_TOKEN: ${{ github.token }}
PR_NUMBER: ${{ github.event.pull_request.number }}
EXPECTED_CHANGED_FILES: ${{ steps.check.outputs.changed_files }}
run: |
set -euo pipefail
changed_json="$RUNNER_TEMP/merge-changed-files.json"
changed_paths="$RUNNER_TEMP/merge-changed-paths.txt"
if gh api --paginate --slurp \
"repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}/files?per_page=100" \
> "$changed_json"; then
observed=$(jq '[.[][] | .filename] | unique | length' "$changed_json")
if [ "$observed" = "$EXPECTED_CHANGED_FILES" ]; then
jq -r '.[][] | .filename, (.previous_filename // empty)' "$changed_json" \
| sort -u > "$changed_paths"
else
echo "::warning::Changed-file API returned $observed of $EXPECTED_CHANGED_FILES paths; selecting all merge lanes."
echo '__FASTVIDEO_CI_PLAN_ALL__' > "$changed_paths"
fi
else
echo "::warning::Changed-file API failed; selecting all merge lanes."
echo '__FASTVIDEO_CI_PLAN_ALL__' > "$changed_paths"
fi
- name: Select minimal merge tests
id: plan
if: steps.check.outputs.has_ready == 'true'
run: |
python3 .github/scripts/plan_merge_ci.py \
--paths-file "$RUNNER_TEMP/merge-changed-paths.txt" \
--github-output "$GITHUB_OUTPUT" \
--summary-file "$GITHUB_STEP_SUMMARY"
- name: Wait for pre-commit and docs build
if: steps.check.outputs.has_ready == 'true'
@@ -66,7 +109,7 @@ jobs:
PR_NUMBER: ${{ github.event.pull_request.number }}
run: bash .github/scripts/gate_full_suite.sh
- name: Trigger Buildkite Full Suite
- name: Trigger Buildkite merge gate
if: steps.check.outputs.has_ready == 'true'
env:
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
@@ -76,6 +119,10 @@ jobs:
PR_TITLE: ${{ github.event.pull_request.title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
MERGE_TEST_PLAN: ${{ steps.plan.outputs.merge_test_plan }}
MERGE_GOLDEN_TESTS: ${{ steps.plan.outputs.merge_golden_tests }}
MERGE_SSIM_TESTS: ${{ steps.plan.outputs.merge_ssim_tests }}
MERGE_PLAN_LABEL: ${{ steps.plan.outputs.merge_plan_label }}
run: |
curl -sS --fail-with-body -X POST \
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds" \
@@ -84,8 +131,11 @@ jobs:
--data-raw "$(jq -n \
--arg commit "$PR_SHA" \
--arg branch "$PR_BRANCH" \
--arg message "Full Suite for PR #${PR_NUMBER}" \
--arg message "Merge gate [${MERGE_PLAN_LABEL}] for PR #${PR_NUMBER}" \
--arg pr_title "$PR_TITLE" \
--arg merge_test_plan "$MERGE_TEST_PLAN" \
--arg merge_golden_tests "$MERGE_GOLDEN_TESTS" \
--arg merge_ssim_tests "$MERGE_SSIM_TESTS" \
--argjson pr_id "$PR_NUMBER" \
'{
commit: $commit,
@@ -95,8 +145,11 @@ jobs:
pull_request_id: $pr_id,
pull_request_base_branch: "main",
env: {
TEST_SCOPE: "full",
TEST_SCOPE: "merge",
FULL_SUITE: "true",
MERGE_TEST_PLAN: $merge_test_plan,
MERGE_GOLDEN_TESTS: $merge_golden_tests,
MERGE_SSIM_TESTS: $merge_ssim_tests,
PR_NUMBER: ($pr_id | tostring),
PR_TITLE: $pr_title
}
+4 -4
View File
@@ -38,17 +38,17 @@ jobs:
**How our CI works:**
PRs run a two-tier CI system:
PRs run a three-tier CI system:
1. **Pre-commit** — formatting (yapf), linting (ruff), type checking (mypy). Runs immediately on every PR.
2. **Fastcheck** — core GPU tests (encoders, VAEs, transformers, kernels, unit tests). Runs automatically via Buildkite on relevant file changes (~10-15 min).
3. **Full Suite** — integration tests, training pipelines, SSIM regression. Runs only when a reviewer adds the `ready` label.
2. **Fastcheck** — six core GPU lanes run automatically via Buildkite (~10-15 min).
3. **Merge gate** — a reviewer adds `ready`; changed paths select only the relevant integration, training, golden, or SSIM coverage.
**Before your PR is reviewed:**
- [ ] `pre-commit run --all-files` passes locally
- [ ] You've added or updated tests for your changes
- [ ] The PR description explains what and why
If pre-commit fails, a bot comment will explain how to fix it. Fastcheck and Full Suite results appear in the Checks section below.
If pre-commit fails, a bot comment will explain how to fix it. Fastcheck and merge-gate results appear in the Checks section below.
**Useful links:**
- [Contributing Guide](https://hao-ai-lab.github.io/FastVideo/contributing/overview/)
+27
View File
@@ -13,6 +13,11 @@ on:
required: false
default: false
type: boolean
build_ci_runner_image:
description: 'Build the ARM64 CUDA 13 CI runner image (sm_100)'
required: false
default: false
type: boolean
# Auto-rebuild the CUDA images when a repository-controlled image input
# changes on main. This includes the trusted SM89 kernel artifact's source,
# metadata/key helper, ABI dependency metadata, and build orchestration.
@@ -198,6 +203,28 @@ jobs:
docker buildx imagetools create "${TAG_ARGS[@]}" "${IMAGE_REFS[@]}"
docker buildx imagetools inspect "${TAGS[0]}"
# The CI runner is ARM64 like DGX Spark, but targets sm_100 rather than sm_121.
# Publish a single-architecture variant so the self-hosted CI runner can reuse
# the exact prebuilt kernel instead of compiling it in every job.
build-ci-runner-image:
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_ci_runner_image == 'true' }}
uses: ./.github/workflows/_template-build-image.yml
with:
python_version: '3.12'
dockerfile_path: docker/Dockerfile
tag_suffix: py3.12-cuda13.0.0-sm100
runner: ubuntu-24.04-arm
architecture: arm64
build_args: |
PYTHON_VERSION=3.12
CUDA_VERSION=13.0.0
UV_TORCH_BACKEND=cu130
TORCH_CUDA_ARCH_LIST=10.0
CMAKE_BUILD_PARALLEL_LEVEL=1
FLASH_ATTN_WHEEL_TAG=cu130torch2.12
FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.22
secrets: inherit
# Dreamverse matrix: {backend, UI} x {12.6.3, 13.0.0}, Python 3.12. Torch backend
# matches the base CUDA (cu126 / cu130). Keep these images amd64-only until the
# required FA4 dependency stack is available and validated on arm64.
+14 -7
View File
@@ -62,8 +62,9 @@ jobs:
cuda-version: '13.0.0'
torch-cuda-short: 'cu130'
platform:
# x86_64 builds the full cu126 + cu130 set (cu130 ships the consumer
# Blackwell sm_120a FP4 kernels).
# x86_64 builds the full cu126 + cu130 set. cu130 ships the
# data-center Blackwell sm_100a VSA and consumer sm_120a FP4
# kernels.
- os: ubuntu-22.04
arch: x86_64
wheel-plat: manylinux_2_35_x86_64
@@ -124,7 +125,7 @@ jobs:
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
run: |
sudo apt update
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
sudo apt install -y git gcc-11 g++-11 clang-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
# Allow Git to Access Safe Directory
@@ -168,7 +169,8 @@ jobs:
# covers sm_120a; turbodiffusion covers sm_100a+sm_120a. The sm_100 FP4
# forward is the FA4 CuTe DSL path in the fastvideo package (PR #1221),
# JIT-compiled at runtime — not built into this wheel.
# * x86_64 cu130 = Hopper TK + consumer Blackwell sm_120a FP4.
# * x86_64 cu130 = Hopper TK + data-center Blackwell sm_100a VSA
# + consumer Blackwell sm_120a FP4.
# * x86_64 cu126 = Hopper TK only (older drivers; CUDA < 12.8 has no FP4).
# The per-arch split in CMakeLists pins the FP4 targets to sm_120a and builds
# the main extension for the full arch list. CMAKE_BUILD_PARALLEL_LEVEL caps
@@ -178,7 +180,7 @@ jobs:
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=OFF -DFASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER=ON"
export CMAKE_BUILD_PARALLEL_LEVEL=1
elif [ "${{ matrix.torch-cuda.torch-cuda-short }}" = "cu130" ]; then
export TORCH_CUDA_ARCH_LIST="9.0a;12.0a"
export TORCH_CUDA_ARCH_LIST="9.0a;10.0a;12.0a"
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DFASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
# A single FP4 TU (attn_qat_infer) can use ~8-12 GB on its own, so serialize.
export CMAKE_BUILD_PARALLEL_LEVEL=1
@@ -194,7 +196,11 @@ jobs:
python -m build --wheel --outdir dist
# Fix the wheel to be manylinux compliant
uv pip install --system auditwheel
# Ubuntu 22.04 ships patchelf 0.14.3, while current auditwheel
# requires at least 0.14.5. Use the stable PyPI binary on both
# x86_64 and aarch64 release runners.
uv pip install --system auditwheel patchelf==0.17.2.4
patchelf --version
# Point auditwheel at torch libs, but do not vendor them into the wheel.
TORCH_LIB_DIR=$(python - <<'PY'
import os
@@ -211,7 +217,8 @@ jobs:
--exclude libtorch.so \
--exclude libc10.so \
--exclude libc10_cuda.so \
--exclude libtorch_python.so
--exclude libtorch_python.so \
--exclude libnccl.so.2
# Move fixed wheels back to dist for upload consistency
rm dist/*.whl
mv fixed_dist/*.whl dist/
+5 -3
View File
@@ -9,6 +9,7 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/08/23`: [FastH3 Preview v0.2](https://huggingface.co/FastVideo/FastVideo-Minimax-FastH3-Preview-v0.2) is a 4-step DMD2-distilled MiniMax-H3 checkpoint that generates synchronized video and audio. Run the verified [basic FastH3 example](examples/inference/basic/basic_fasth3.py), see the [inference guide](examples/inference/basic/README.md#fasth3-preview), or have a coding agent install FastVideo with the [agent setup prompt](#install-with-an-ai-coding-agent).
- `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/).
@@ -63,9 +64,10 @@ 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
> **On an Apple Silicon Mac?** FastVideo runs FastMetal-QAD through an MLX
> runtime. Install with `uv pip install -e '.[mlx]'`, download
> [`FastVideo/FastMetal-1.3B-QAD`](https://huggingface.co/FastVideo/FastMetal-1.3B-QAD),
> 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.
+26 -12
View File
@@ -89,11 +89,29 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
zsh \
vim \
curl \
ffmpeg \
libgl1 \
libglib2.0-0 \
libx11-dev \
gcc-11 \
g++-11 \
clang-11 \
cmake \
pkg-config \
build-essential \
libssl-dev \
&& rm -rf /var/lib/apt/lists/*
# Rust toolchain: some dependencies only ship sdists on aarch64 and need cargo
# to build. The dormant legacy Modal image layers the identical apt set +
# rustup on top of this image (fastvideo/tests/modal/pr_test.py); baking both
# here keeps its manual rollback path reproducible without changing the Slurm
# runner's package surface.
RUN set -o pipefail && \
curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable --profile minimal && \
/root/.cargo/bin/cargo --version && /root/.cargo/bin/rustc --version
ENV PATH=/root/.cargo/bin:${PATH}
# Set up C++20 compilers for ThunderKittens
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
@@ -138,6 +156,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
source /opt/venv/bin/activate && \
uv pip install --upgrade pip && \
uv pip install --excludes docker/uv-excludes ".[dev]" && \
python -c "import cv2; print('OpenCV', cv2.__version__)" && \
PYTAG=cp$(echo "${PYTHON_VERSION}" | tr -d .) && \
case "${TARGETARCH:-amd64}" in \
amd64) \
@@ -169,26 +188,21 @@ RUN --mount=type=cache,target=/opt/uv/cache \
# flash_attn/__init__.py), so FA2/varlen/bert_padding stay from the install above;
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
# Then verify both survive so a broken overlay fails the build instead of shipping
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
# arm64 and the image runs FA3/FA2 as usual.
# an FA2-less image. The pinned stack is validated on ARM64 GB200 (sm_100) as well
# as x86; FA4 remains opt-in through FASTVIDEO_FA4=1 so lanes with FA2 baselines
# keep their existing numerics.
RUN --mount=type=cache,target=/opt/uv/cache \
source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
if [ "${TARGETARCH}" = "arm64" ]; then \
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
else \
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
python -c "import flash_attn; assert hasattr(flash_attn, 'flash_attn_func'), 'FA2 was clobbered by the cute overlay'; import flash_attn.cute; print('FA2 + FA4 cute OK')"; \
fi
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
python -c "import flash_attn; assert hasattr(flash_attn, 'flash_attn_func'), 'FA2 was clobbered by the cute overlay'; import flash_attn.cute; print('FA2 + FA4 cute OK')"
COPY . .
# Build immutable FastVideo kernel wheels for the published image. The requested
# architecture remains installed for normal image users; amd64 images also carry
# an SM89 artifact so the predominant L40S Modal lanes can reuse it exactly.
# an SM89 artifact for L40S users and the dormant legacy rollback path.
ARG FASTVIDEO_KERNEL_PREBUILT_DIR=/opt/fastvideo-kernel-prebuilt
RUN --mount=type=cache,target=/opt/uv/cache \
source $HOME/.local/bin/env && \
+208 -76
View File
@@ -6,7 +6,8 @@ lives in [Testing](testing.md).
## Overview
FastVideo splits validation across GitHub Actions, Buildkite, and Modal:
FastVideo splits validation across GitHub Actions, Buildkite, Slinky Slurm,
and Mergify:
```text
PR opened or updated
@@ -16,17 +17,24 @@ PR opened or updated
| style, lint, type, spelling, Markdown, workflow syntax, filenames
|
|-- Tier 2: Fastcheck
| Buildkite orchestrates Modal GPU jobs
| path-filtered component and unit checks
| Buildkite schedules six lanes on the Slinky Slurm cluster
| encoder, VAE, transformer, kernel, unit, DreamVerse
|
|-- /merge, /test full, or ready label
|-- /merge or ready label
|
`-- Tier 3: Full Suite
Buildkite orchestrates Modal GPU jobs
path-filtered integration, SSIM, training, eval, and performance checks
`-- Tier 3: Change-aware merge gate
trusted base-branch planner classifies the complete PR diff
Buildkite adds only relevant integration, quality, training,
API, or performance lanes on Slinky Slurm
|
pass -> Mergify squash-merges when all merge conditions pass
fail -> fix, push, and re-run
|-- /test full
`-- Explicit all-20-lane diagnostic run
`-- weekly schedule on main
`-- Complete four-GPU SSIM matrix
```
CI is not one monolithic job:
@@ -34,10 +42,31 @@ CI is not one monolithic job:
- GitHub Actions owns pre-commit, slash-command handling, aggregate status
updates, docs deployment, image builds, package publishing, and community
automations.
- Buildkite owns the GPU test pipeline and path filtering.
- Modal owns the actual GPU execution environment for test jobs.
- Buildkite owns the GPU test graph, statuses, and trusted dispatch control
plane. Its agent runs on the Slurm login plane; it does not execute test
payloads.
- Slinky Slurm is the only active CI compute backend. A host-owned dispatcher
leases GPUs from a persistent four-GPU allocation and runs each lane in an
isolated Enroot container at the immutable PR SHA.
- Mergify owns merge protection, labeling, and the final squash merge.
The old files under `fastvideo/tests/modal/` are retained as dormant manual
rollback code. `.buildkite/scripts/pr_test.sh` rejects Buildkite invocations,
and no pipeline or slash-command route calls Modal.
Three Buildkite entry pipelines share the validated graph:
| Pipeline | Trigger | Scope |
|---|---|---|
| `pr-fastcheck` | Automatic pull-request webhook | Six Fastcheck lanes |
| `ci` | `/merge`, `ready`, schedules, and `/test` API builds | Change-aware merge gates, scheduled SSIM, explicit Full Suite, Fastcheck reruns, or one direct lane |
| `fastvideo-performance-lane` | Weekly scheduler | Direct performance lane |
Each entry pipeline starts with the same trusted `pipeline-upload` job on the
`ci-runner` queue. The `ci` pipeline's incoming GitHub webhook is disabled;
otherwise it would duplicate the automatic `pr-fastcheck` build. API and
scheduled builds continue to work with webhook processing disabled.
## CI Tiers
### Tier 1: Pre-commit
@@ -70,34 +99,37 @@ debugging a hook implementation.
| Attribute | Value |
|---|---|
| Triggered by | Buildkite PR builds with `TEST_SCOPE=fastcheck` or unset |
| Runner | Buildkite agent that launches Modal GPU jobs |
| Compute | Slinky Slurm (`ci-runner` queue) |
| Definition | `.buildkite/pipeline.yml` |
| Entrypoint | `.buildkite/scripts/pr_test.sh` -> `fastvideo/tests/modal/pr_test.py` |
| Entrypoint | Trusted host driver -> `.buildkite/scripts/unit_test.sh` or `.buildkite/scripts/lanes/*.sh` |
Fastcheck uses Buildkite's `monorepo-diff` plugin. Jobs whose watched paths did
not change are skipped and do not block the aggregate `fastcheck-passed`
status.
Fastcheck always schedules these six lanes: encoder, VAE, transformer, custom
kernels, unit tests, and DreamVerse. Static steps replace the former
host-side path-filter plugin: the login plane never checks out or executes PR
code.
| Buildkite label | `TEST_TYPE` | Main watched paths |
|---|---|---|
| Encoder Tests | `encoder` | `fastvideo/models/encoders/**`, `fastvideo/models/loader/**`, `fastvideo/tests/encoders/**`, `pyproject.toml`, `docker/Dockerfile` |
| VAE Tests | `vae` | `fastvideo/models/vaes/**`, `fastvideo/models/loader/**`, `fastvideo/tests/vaes/**`, `pyproject.toml`, `docker/Dockerfile` |
| Transformer Tests | `transformer` | `fastvideo/models/dits/**`, `fastvideo/models/loader/**`, `fastvideo/tests/transformers/**`, `fastvideo/layers/**`, `fastvideo/attention/**`, `pyproject.toml`, `docker/Dockerfile` |
| Kernel Tests | `kernel_tests` | `fastvideo-kernel/**`, `pyproject.toml`, `docker/Dockerfile` |
| Unit Tests | `unit_test` | `fastvideo/**`, `.buildkite/**`, `.github/**`, `pyproject.toml`, `docker/Dockerfile` |
| DreamVerse App Tests | `dreamverse_app` | `apps/dreamverse/**`, `pyproject.toml` |
### Tier 3: Full Suite
### Tier 3: Change-Aware Merge Gate
| Attribute | Value |
|---|---|
| Triggered by | `/merge`, adding `ready`, `/test full`, or a new push to a PR that already has `ready` |
| Runner | Buildkite agent that launches Modal GPU jobs |
| Triggered by | `/merge`, adding `ready`, or a new push to a PR that already has `ready` |
| Compute | Slinky Slurm only (`ci-runner` queue) |
| Definition | `.buildkite/pipeline.yml` |
| Entrypoint | `.buildkite/scripts/pr_test.sh` -> `fastvideo/tests/modal/pr_test.py` |
| Entrypoint | `/opt/fastvideo-ci-runner/run-ci` (`run-unit` is a compatibility wrapper) |
Full Suite is also path-filtered. It validates broader behavior before Mergify
can merge a PR.
Fastcheck is the universal six-lane baseline. The merge gate does not repeat
those jobs: it classifies every changed path and adds only the relevant lanes
from the fourteen-lane integration set below. Selected jobs are hard gates;
there are no soft-fail hardware lanes. A documentation-only PR can therefore
finish its merge build after the trusted uploader, while model-family changes
typically add focused golden-gate and SSIM files and a training-only change
adds only its owning training lane.
`.github/scripts/plan_merge_ci.py` is the canonical path policy. It runs from
the immutable base SHA under `pull_request_target`; PR code is never executed
on the GitHub runner. The changed-file list includes both sides of renames. An
API failure, truncated response, empty list, unknown build input, or unknown
source path fails closed to all fourteen integration lanes.
A `ready`-labeled PR does not hit Buildkite immediately:
`ci-trigger-full-suite.yml` first runs `.github/scripts/gate_full_suite.sh`,
@@ -106,25 +138,100 @@ head. A red cheap check blocks the suite (fail closed; the next push re-arms
it), while a GitHub outage or a >25 min wait lets it run anyway (fail open).
`/test full` bypasses the gate.
| Buildkite label | `TEST_TYPE` | Main watched paths |
|---|---|---|
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
| LoRA Training Tests | `training_lora` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
| Training Tests VSA | `training_vsa` | `fastvideo/**`, `fastvideo-kernel/**`, `pyproject.toml`, `docker/Dockerfile` |
| Inference Tests VMoBA | `inference_vmoba` | `fastvideo-kernel/**`, `fastvideo/attention/backends/vmoba.py` |
| Performance Tests | `performance` | DiTs, pipelines, attention, layers, worker, entrypoints, performance tests/configs |
| API Server Tests | `api_server` | OpenAI entrypoints, serve CLI, OpenAI API integration test |
| Train Framework Tests | `train_framework` | `fastvideo/train/**`, train model/method tests, model loader, DiTs |
| Eval Metrics Tests | `eval` | `fastvideo/eval/**`, `fastvideo/tests/eval/**`, `pyproject.toml`, `docker/Dockerfile` |
The complete static graph remains available through `/test full`; path
selection never deletes or dynamically invents a Buildkite step.
| Lane | Public `TEST_TYPE` | GPUs | Typical merge trigger |
|---|---|---:|---|
| Encoder | `encoder` | 1 | Universal Fastcheck |
| VAE | `vae` | 1 | Universal Fastcheck |
| Transformer | `transformer` | 1 | Universal Fastcheck |
| Kernel | `kernel_tests` | 1 | Universal Fastcheck |
| Unit | `unit_test` | 1 | Universal Fastcheck |
| DreamVerse | `dreamverse_app` | 1 | Universal Fastcheck |
| Golden gate | `golden_gate` | 1 | Model, pipeline, attention, layer, or output changes |
| SSIM | `ssim` | 4 | Matching model/SSIM paths; focused files when possible |
| LoRA inference | `inference_lora` | 1 | LoRA inference/shared LoRA paths |
| LoRA extraction | `lora_extraction` | 1 | LoRA extraction/shared LoRA paths |
| Vanilla training | `training` | 4 | Legacy vanilla/shared training paths |
| DMD distillation | `distillation_dmd` | 2 | DMD/shared training paths |
| Self-forcing | `self_forcing` | 2 | Self-forcing/shared training paths |
| LoRA training | `training_lora` | 2 | LoRA/shared training paths |
| VSA training | `training_vsa` | 2 | VSA/shared training paths |
| VMoBA inference | `inference_vmoba` | 1 | VMoBA backend/config paths |
| Performance | `performance` | 2 | Performance tests or benchmark policy |
| API server | `api_server` | 1 | API, worker, or server entrypoint paths |
| Modular train framework | `train_framework` | 1 | `fastvideo/train/` and its tests |
| Eval metrics | `eval` | 1 | `fastvideo/eval/` and its tests |
Golden-gate and SSIM selections are basenames, not arbitrary pytest arguments.
The private host checks the comma-separated allowlist before staging, and the
container checks it again before invoking pytest. Shared quality-harness
changes still run the complete owning lane. The full SSIM matrix also runs on
`main` every Sunday at 05:00 UTC through `ci-scheduled-ssim.yml`; `/test ssim`
and `/test full` remain available for deliberate complete runs.
The four Buildkite workers may accept multiple jobs concurrently. The
agent-owned lease broker packs their requested GPU counts onto one persistent
four-GPU Slurm allocation and waits when capacity is full. A four-GPU lane
such as SSIM or vanilla training owns the whole tray; two two-GPU lanes or up
to four one-GPU lanes can overlap without sharing devices. SSIM and vanilla
training also share the `fastvideo/slinky/whole-tray` Buildkite concurrency
group. That keeps the second whole-tray lane in Buildkite instead of consuming
an agent and its command timeout while the first lane waits for four free GPUs.
Because packed Enroot containers share the node network namespace, each GPU
lease receives its own 100-port rendezvous range. Tests preserve the
runner-assigned `MASTER_PORT`, and parallel SSIM tasks use distinct offsets
inside that range.
See [Performance Benchmarks](performance_benchmarks.md) for the performance
lane's thresholds, rolling baseline, artifacts, and reseeding process.
## `/merge` Request Flow
The Buildkite agent and Slurm have deliberately separate responsibilities:
```text
/merge PR comment
-> GitHub verifies write permission and refreshes the `ready` label
-> base-branch `ci-trigger-full-suite` workflow fetches the PR file list,
computes MERGE_TEST_PLAN plus focused golden/SSIM basenames, and gates on
cheap checks
-> sends the PR SHA to pipeline `ci` with TEST_SCOPE=merge and FULL_SUITE=true
-> trusted `pipeline-upload` job on queue `ci-runner`
fetch exact-SHA .buildkite/pipeline.yml
normalize + validate the complete static 20-lane policy and conditions
upload static Buildkite steps
-> each accepted step reaches the host policy hook
validate org/repo/SHA/ref/step key/command/scope/timeout
skip checkout on the login plane
stage a mode-0600 request on Lustre
lease 1-4 GPUs and attach an `srun` step to the Slinky tray
-> Enroot worker container
clone and verify the exact PR SHA
install the lane's project extras and cached kernel
run the repository-owned lane script
write numeric exit status and approved artifacts
-> trusted host returns that status to Buildkite
-> Buildkite publishes `full-suite-passed` only when every selected lane passes
```
The pipeline uploader and dispatcher run on the Slurm login plane, but those
are control-plane operations only. Python tests, model loading, CUDA kernels,
Node/Playwright checks, inference, training, SSIM generation, and performance
benchmarks all execute in Slurm allocations.
PR-controlled values never become host commands. The host policy accepts only
the pinned pipeline uploader or a known lane tuple. It rejects plugins,
artifact globs, shell injection variables, non-immutable commits, and unknown
commands before checkout. Hugging Face credentials are added only for lanes
that declare them, passed through a mode-0600 request file, and removed before
the PR payload starts. Active training lanes keep W&B offline and do not stage
a W&B credential. The ARM64 image includes the pinned FA4 CuTe overlay validated
on GB200. SSIM opts into FA4 to preserve its reference-video numerics; lanes
with FA2 baselines keep `FASTVIDEO_FA4=0`. Performance artifacts are relayed
afterward by the trusted host from an allowlisted directory and extension set.
## Slash Commands
Slash commands are handled by `.github/workflows/ci-slash-commands.yml`.
@@ -132,7 +239,7 @@ Repository write permission is required.
| Command | Effect |
|---|---|
| `/merge` | Adds `ready` and triggers Full Suite for the PR head branch. |
| `/merge` | Adds `ready` and triggers the path-aware merge gate for the PR head. |
| `/test full` | Runs the whole Full Suite with `TEST_SCOPE=full`. |
| `/test fastcheck` | Runs the whole Fastcheck suite with `TEST_SCOPE=fastcheck`. |
| `/test pre-commit` | Re-runs the pre-commit workflow on the PR merge ref. |
@@ -149,6 +256,7 @@ Valid direct test names:
| `/test unit` | `unit_test` |
| `/test dreamverse` | `dreamverse_app` |
| `/test ssim` | `ssim` |
| `/test golden-gate` | `golden_gate` |
| `/test training` | `training` |
| `/test lora-inference` | `inference_lora` |
| `/test lora-training` | `training_lora` |
@@ -162,12 +270,18 @@ Valid direct test names:
| `/test train-framework` | `train_framework` |
| `/test eval` | `eval` |
The temporary `<name>-ci` spellings remain accepted as compatibility aliases;
they select the same Slurm lane and do not identify a second backend.
When a direct test completes successfully, Buildkite posts
`direct-test-completed`. `.github/workflows/ci-aggregate-status.yml` then reads
the latest Buildkite statuses for the commit and updates `fastcheck-passed` or
`full-suite-passed` if all jobs in that group are green.
Skipped path-filtered jobs have no status entry and do not block the aggregate.
Buildkite label emojis define the status namespace used by that aggregation:
`:microscope:` is reserved for the six Fastcheck lanes, while Full-Suite-only
lanes use `:test_tube:` or `:bar_chart:`. Each active lane has exactly one
label and therefore one status context.
## Merge Protection
@@ -177,7 +291,7 @@ Mergify enforces these conditions before it squash-merges to `main`:
|---|---|
| `check-success~=pre-commit` | Tier 1 passed. |
| `check-success=fastcheck-passed` | All triggered Fastcheck jobs passed. |
| `check-success=full-suite-passed` | All triggered Full Suite jobs passed. |
| `check-success=full-suite-passed` | The selected merge gate or explicit Full Suite passed. |
| `#approved-reviews-by>=1` | At least one approving review. |
| Valid title regex | PR title starts with an accepted `[type]` tag. |
| `label=ready` | The PR has entered the merge flow. |
@@ -226,29 +340,40 @@ Process labels:
| Label | Who sets it | Meaning |
|---|---|---|
| `ready` | `/merge` or maintainer action | Triggers/keeps Full Suite active and enables auto-merge. |
| `ready` | `/merge` or maintainer action | Triggers/keeps the change-aware merge gate active and enables auto-merge. |
| `needs-rebase` | Mergify | PR has merge conflicts. |
| `do-not-merge` | Maintainer | Blocks merge regardless of CI status. |
## Modal Test Entrypoints
## Slurm Lane Entrypoints
All Buildkite test jobs go through `.buildkite/scripts/pr_test.sh`, which:
Every active test selection lives in `.buildkite/scripts/unit_test.sh` or a
focused `.buildkite/scripts/lanes/<lane>.sh`. The private, agent-owned lane
table binds each internal `*_ci` type to that script, its GPU count, wall-clock
limit, dependency extras, kernel-build policy, secrets, and artifacts. The
internal suffix is an implementation detail; there is only one active backend.
1. Reads Buildkite secrets for Modal, Hugging Face, and W&B when needed.
2. Selects a Modal function based on `TEST_TYPE`.
3. Passes Buildkite metadata into the Modal container.
4. Runs the selected test command from `fastvideo/tests/modal/pr_test.py` or
`fastvideo/tests/modal/ssim_test.py`.
5. Uploads performance artifacts for `TEST_TYPE=performance`.
SSIM uses `fastvideo/tests/ssim/ci_runner.py` inside a single four-GPU lease.
It discovers `REQUIRED_GPUS` and `*_MODEL_TO_PARAMS` with AST parsing, then
packs independent pytest subprocesses across the visible GPUs with fail-fast
termination. Performance writes reports to a host-mounted artifact directory;
the host uploads only `.md`, `.html`, `.json`, and `.csv` files after the
container exits.
The Modal launchers remain in the repository for manual rollback archaeology,
but they are not CI entrypoints. `pr_test.sh` rejects Buildkite calls and needs
`FASTVIDEO_ENABLE_LEGACY_MODAL_CI=1` even for a local manual invocation.
If you add a new CI test category:
1. Add the Modal function in `fastvideo/tests/modal/pr_test.py` or a focused
companion module.
2. Add the `TEST_TYPE` case in `.buildkite/scripts/pr_test.sh`.
3. Add the Buildkite direct-test step and any Fastcheck/Full Suite path filters
in `.buildkite/pipeline.yml`.
4. Add or update the `/test` mapping in `.github/workflows/ci-slash-commands.yml`.
1. Add an executable `.buildkite/scripts/lanes/<lane>.sh` containing the test
payload only.
2. Add the static Buildkite step in `.buildkite/pipeline.yml`, its source/test
ownership in `.github/scripts/plan_merge_ci.py`, and the `/test` mapping in
`.github/workflows/ci-slash-commands.yml`.
3. Extend `fastvideo/tests/contract/test_ci_test_collection.py`,
`test_merge_ci_plan.py`, and the trusted private runner's lane table plus
uploader policy in the same rollout.
4. Validate on the target GB200 hardware before making the lane a merge gate.
5. Document the lane here and link any domain-specific authoring guide.
## CD And Release Workflows
@@ -282,13 +407,18 @@ the explicit `py3.12-cuda13.0.0-latest` tag. This publication policy does not
change the unparameterized `docker/Dockerfile` build defaults, which remain CUDA
13 and `cu130`.
Published amd64 development images keep their configured Hopper kernel wheel
installed and also carry an immutable SM89 wheel under
`/opt/fastvideo-kernel-prebuilt`. Modal PR and SSIM jobs select the exact
source, ABI, and GPU-architecture match from that directory, so L40S jobs reuse
the trusted image artifact while kernel-changing PRs still build locally. Once
a kernel or artifact-key change reaches `main`, the image workflow republishes
the matching trusted artifact before later jobs consume the updated image tag.
Published development images carry architecture-specific kernel wheels under
`/opt/fastvideo-kernel-prebuilt`. The Slurm worker selects the exact source,
ABI, and GPU-architecture match, so normal lanes reuse the trusted artifact
while kernel-changing PRs still build locally. Once a kernel or artifact-key
change reaches `main`, the image workflow republishes the matching artifact
before later jobs consume the updated image pin.
The same workflow publishes a single-architecture ARM64, CUDA 13, SM100 image
for the self-hosted CI runner under the
`py3.12-cuda13.0.0-sm100-{latest,sha-*}` tags. It carries the matching prebuilt
kernel wheel so runner jobs can validate and install the exact source and ABI
match instead of recompiling it in every lane.
The optional Dreamverse matrix builds backend and UI images for CUDA 12.6 and
CUDA 13 on `amd64`. Dreamverse remains `amd64`-only because its FA4 dependency
@@ -320,15 +450,17 @@ The reusable implementation lives in
| `.github/mergify.yml` | Merge protection, PR title validation, PR labels, conflict labels, auto-merge |
| `.github/workflows/ci-precommit.yml` | Tier 1 pre-commit |
| `.github/workflows/ci-slash-commands.yml` | `/merge` and `/test` handling |
| `.github/workflows/ci-trigger-full-suite.yml` | Full Suite trigger for `ready` PRs and new pushes to ready PRs |
| `.github/workflows/ci-aggregate-status.yml` | Aggregate Fastcheck/Full Suite commit statuses |
| `.buildkite/pipeline.yml` | Buildkite test graph and path filters |
| `.buildkite/scripts/pr_test.sh` | Buildkite-to-Modal test dispatcher |
| `fastvideo/tests/modal/pr_test.py` | Modal functions for most GPU CI lanes |
| `fastvideo/tests/modal/ssim_test.py` | Modal functions and partitioning for SSIM |
| `.github/workflows/ci-trigger-full-suite.yml` | Change-aware merge-gate trigger for `ready` PRs and new pushes |
| `.github/workflows/ci-scheduled-ssim.yml` | Weekly complete SSIM trigger on `main` |
| `.github/scripts/plan_merge_ci.py` | Trusted changed-path to integration-lane planner |
| `.github/workflows/ci-aggregate-status.yml` | Aggregate Fastcheck and explicit Full Suite direct-rerun statuses |
| `.buildkite/pipeline.yml` | Static 20-lane Slurm Buildkite graph |
| `.buildkite/scripts/unit_test.sh`, `.buildkite/scripts/lanes/*.sh` | Active Slurm lane payloads |
| `fastvideo/tests/ssim/ci_runner.py` | Four-GPU Slurm SSIM scheduler |
| `.buildkite/scripts/pr_test.sh`, `fastvideo/tests/modal/*.py` | Dormant manual Modal rollback path (disabled in Buildkite) |
| `.buildkite/performance-benchmarks/tests/*.json` | Performance benchmark configs and thresholds |
| `.github/workflows/infra-docs.yml` | Docs build and GitHub Pages deploy |
| `.github/workflows/infra-build-image.yml` | Automatic CUDA matrix and manual Docker image builds |
| `.github/workflows/infra-build-image.yml` | CUDA matrix, CI runner image, and manual Docker image builds |
| `.github/workflows/publish-fastvideo.yml` | FastVideo PyPI publishing |
| `.github/workflows/publish-kernel.yml` | FastVideo kernel PyPI publishing |
| `.github/workflows/publish-comfyui.yml` | ComfyUI registry publishing |
+10 -8
View File
@@ -23,8 +23,8 @@ It serves three audiences:
pytest fastvideo/tests/performance/ -vs
# Optional: compare against the rolling HF baseline.
# PERF_REPORTS_DIR defaults to /root/data/perf_reports for Modal/CI, so
# override it when running outside the container.
# PERF_REPORTS_DIR defaults to /root/data/perf_reports in a CI container, so
# override it for a local run.
PERF_REPORTS_DIR=/tmp/fastvideo_perf_reports \
python fastvideo/tests/performance/compare_baseline.py
@@ -459,16 +459,18 @@ FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
| `FASTVIDEO_FA4` | `0` | `test_inference_performance.py` | FlashAttention-4 toggle included in `software_profile_id`. |
| `FASTVIDEO_PERFORMANCE_PROFILE_VERSION` | unset | `test_inference_performance.py` | Optional explicit software cohort/profile version included in `software_profile_id`. |
| `IMAGE_VERSION` | unset | `test_inference_performance.py` | CI container image/profile version included in `software_profile_id` when available. |
| `FASTVIDEO_CONTAINER_IMAGE_REF` | unset | `pr_test.py`, `launch_l40s_job.py`, `test_inference_performance.py` | Resolved CI container image ref or digest recorded in `environment_metadata` for audit without changing `software_profile_id`. |
| `FASTVIDEO_CONTAINER_IMAGE_REF` | unset | Slurm runner, `test_inference_performance.py` | Pinned CI container image digest recorded in `environment_metadata` and `software_profile_id`. |
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
## CI integration
The performance step can run on demand with `/test performance` and as part of
the Full Suite (see [CI/CD Architecture](ci_architecture.md)). The Modal entry
point is `fastvideo/tests/modal/pr_test.py:run_performance_tests` and the
Buildkite artifact upload is in
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
The performance step can run on demand with `/test performance`, through a
merge gate when performance tests or benchmark policy changed, and as part of
an explicit `/test full` run (see [CI/CD Architecture](ci_architecture.md)).
The weekly `fastvideo-performance-lane` schedule runs the same Slurm payload.
The active entry point is `.buildkite/scripts/lanes/performance.sh`; the
trusted host dispatcher relays its allowlisted reports to Buildkite after the
isolated container exits.
Each performance build runs pytest first. PR and direct runs only continue to
`compare_baseline.py` when that fixed-threshold phase passes; if pytest fails,
+10 -8
View File
@@ -42,19 +42,20 @@ The important process labels are:
| Label | Meaning |
|---|---|
| `ready` | The PR is ready for Full Suite and auto-merge consideration. |
| `ready` | The PR is ready for the change-aware merge gate and auto-merge consideration. |
| `needs-rebase` | The PR has merge conflicts with `main`. |
| `do-not-merge` | A maintainer has blocked merge. |
## CI Summary
FastVideo has three validation tiers:
FastVideo has three routine validation tiers plus an explicit full diagnostic:
| Tier | Runs when | What it does |
|---|---|---|
| Pre-commit | Pull requests and `/test pre-commit` | Formatting, linting, typing, spelling, Markdown, workflow syntax, filename checks |
| Fastcheck | PR Buildkite builds | Path-filtered component and unit checks on Modal GPU runners |
| Full Suite | `/merge`, `ready`, `/test full`, or new pushes to ready PRs | Path-filtered integration, SSIM, training, eval, API, and performance checks |
| Fastcheck | PR Buildkite builds | Six component, kernel, unit, and app lanes on Slinky Slurm |
| Merge gate | `/merge`, `ready`, or new pushes to ready PRs | Only path-relevant integration lanes on Slinky Slurm; Fastcheck remains the baseline |
| Full Suite | `/test full` | Explicit all-twenty-lane diagnostic run on Slinky Slurm |
See [CI/CD Architecture](ci_architecture.md#ci-tiers) for exact jobs, path
filters, and workflow files.
@@ -66,10 +67,11 @@ filters, and workflow files.
3. Fix pre-commit failures locally with `pre-commit run --all-files`.
4. Wait for at least one approving review.
5. When the PR is approved and ready, comment `/merge`.
6. `/merge` adds `ready` and triggers the Full Suite for the PR branch.
6. `/merge` adds `ready`, waits for cheap checks, and triggers the minimal
path-relevant integration lanes for the PR branch.
7. If all required checks pass, Mergify squash-merges the PR to `main`.
8. If Full Suite fails, fix the regression, push again, and re-run `/merge` or
the failed test.
8. If the merge gate fails, fix the regression, push again, and re-run
`/merge`. Use a targeted `/test` command for diagnosis.
Only contributors with repository write permission can use slash commands. If
you are an external contributor, ask a maintainer to run `/merge` or add
@@ -128,7 +130,7 @@ git push --force-with-lease
Mergify removes `needs-rebase` after conflicts are resolved.
### Full Suite Fails
### Merge Gate Or Full Suite Fails
The failing Buildkite step is the source of truth. Common causes are:
+34 -48
View File
@@ -138,47 +138,31 @@ pytest fastvideo/tests/ssim/ -vs
Use a machine whose GPU and backend match the reference folder you are testing.
## Modal Runs For SSIM
## Slurm CI Runs For SSIM
For CI-like SSIM execution, use `fastvideo/tests/modal/ssim_test.py`:
Comment `/test ssim` on a pull request to run the canonical four-GPU SSIM
lane on the Slinky Slurm cluster. `fastvideo/tests/ssim/ci_runner.py`
discovers the suite without importing test modules, packs independent pytest
processes across the four assigned GPUs, and stops the lane on the first
failure.
The change-aware `/merge` planner may run only the SSIM test basenames owned
by the changed model family. Shared SSIM harness changes still select the
complete lane. Independently, `main` runs the full SSIM matrix every Sunday at
05:00 UTC so infrequently touched model families retain periodic coverage.
For a focused developer run, invoke pytest directly and optionally select one
model from a parameterized test through `FASTVIDEO_SSIM_MODEL_ID`:
```bash
python -m modal run fastvideo/tests/modal/ssim_test.py::run_ssim_tests
pytest fastvideo/tests/ssim/test_wan_t2v_similarity.py -vs
FASTVIDEO_SSIM_MODEL_ID=Wan2.1-T2V-1.3B-Diffusers \
pytest fastvideo/tests/ssim/test_wan_t2v_similarity.py -vs
```
Target specific files or model ids:
```bash
python -m modal run fastvideo/tests/modal/ssim_test.py::run_ssim_tests \
--test-files test_wan_t2v_similarity.py \
--model-ids Wan2.1-T2V-1.3B-Diffusers
```
If `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` is not set, the local
entrypoint fails fast.
To export raw generated videos from Modal to the shared volume:
```bash
python -m modal run fastvideo/tests/modal/ssim_test.py::run_ssim_tests \
--sync-generated-to-volume
```
The raw export path is quality-tiered:
- default params: `ssim_generated_videos/default/<subdir>/generated_videos`
- full-quality params: `ssim_generated_videos/full_quality/<subdir>/generated_videos`
The printed `modal volume get` command downloads into
`./generated_videos_modal/<quality-tier>`. Convert those outputs into local
references with `copy-local`:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
--quality-tier full_quality \
--generated-dir ./generated_videos_modal/full_quality/L40S_reference_videos \
--device-folder L40S_reference_videos
```
The files under `fastvideo/tests/modal/` are retained only as a disabled
manual rollback implementation. No active CI trigger invokes them.
### SSIM Bootstrap Mode
@@ -206,28 +190,30 @@ python fastvideo/tests/ssim/reference_videos_cli.py promote-draft \
## CI Integration
FastVideo CI tests are orchestrated by Buildkite and run on Modal GPU
instances. The main files are:
FastVideo GPU CI is orchestrated by Buildkite and runs only on isolated Slinky
Slurm workers. The main files are:
| File | Purpose |
|---|---|
| `.buildkite/pipeline.yml` | Buildkite test graph and path filters. |
| `.buildkite/scripts/pr_test.sh` | Dispatches `TEST_TYPE` to a Modal function. |
| `fastvideo/tests/modal/pr_test.py` | Modal functions for most test lanes. |
| `fastvideo/tests/modal/ssim_test.py` | Modal functions and partitioning for SSIM. |
| `.buildkite/pipeline.yml` | Static, validated 20-lane Slurm test graph. |
| `.github/scripts/plan_merge_ci.py` | Trusted path-to-lane and focused quality-test policy for `/merge`. |
| `.buildkite/scripts/unit_test.sh`, `.buildkite/scripts/lanes/*.sh` | Repository-owned test payloads executed inside Slurm containers. |
| `fastvideo/tests/ssim/ci_runner.py` | Four-GPU SSIM task discovery and scheduling. |
| `.buildkite/scripts/pr_test.sh`, `fastvideo/tests/modal/*.py` | Dormant manual rollback path; rejected in Buildkite. |
For exact tier membership, path filters, slash commands, and aggregate statuses,
For exact tier membership, slash commands, runner isolation, and aggregate statuses,
see [CI/CD Architecture](ci_architecture.md).
### Adding A New CI Test Category
If a new test does not fit an existing lane:
1. Add a Modal function in `fastvideo/tests/modal/pr_test.py` or a focused
companion module.
2. Add a matching `TEST_TYPE` case in `.buildkite/scripts/pr_test.sh`.
3. Add Buildkite direct-test and path-filter entries in `.buildkite/pipeline.yml`.
4. Add the `/test` mapping in `.github/workflows/ci-slash-commands.yml`.
1. Put the test payload in an executable `.buildkite/scripts/lanes/<lane>.sh`.
2. Add its static step to `.buildkite/pipeline.yml`, its changed-path ownership
to `.github/scripts/plan_merge_ci.py`, and extend the CI contract tests.
3. Add the `/test` mapping in `.github/workflows/ci-slash-commands.yml`.
4. Coordinate the matching GPU, timeout, dependency, secret, and artifact
policy in the private Slurm runner allowlist.
5. Document the new category in [CI/CD Architecture](ci_architecture.md) and add
authoring notes here if contributors need them.
+7 -6
View File
@@ -8,7 +8,7 @@ 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
Measured on the 1.3B INT8 QAD model (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`.
@@ -26,10 +26,11 @@ uv pip install -e ".[mlx]" # RIFE ships vendored; this only needs MLX
```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" \
--model-root ./FastMetal-1.3B-QAD \
--mlx-checkpoint ./FastMetal-1.3B-QAD \
--prompt "A bird's-eye view of a misty forest valley at dawn." \
--num-frames 81 --fast \
--output-path video_samples/fox_fast.mp4
--output-path video_samples/forest_fast.mp4
```
`--num-frames` stays the *target* length; fast mode generates the smallest
@@ -57,9 +58,9 @@ of denoise. It composes with `--fast`; both together run the same clip in
```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" \
--prompt "A bird's-eye view of a misty forest valley at dawn." \
--height 480 --width 832 --num-frames 81 --fast-spatial \
--output-path video_samples/fox_fast_spatial.mp4
--output-path video_samples/forest_fast_spatial.mp4
```
| Flag | Default | Meaning |
@@ -76,6 +76,11 @@ 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."
inference_torch_compile: "Regional inference compile opt-in currently carried through PipelineSelection.experimental rather than CompileConfig."
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."
@@ -119,12 +124,14 @@ surfaces:
vae_precision: "Precision override pending dedicated typed component precision design."
vae_decode_precision: "Decode-only precision override pending dedicated typed component precision design."
image_encoder_precision: "Precision override pending dedicated typed component precision design."
image_encoder_precisions: "Precision overrides pending dedicated typed component precision design."
text_encoder_precisions: "Precision override pending dedicated typed component precision design."
internal_only:
dit_config: "Legacy internal component config object."
upsampler_config: "Legacy internal component config object."
vae_config: "Legacy internal component config object."
image_encoder_config: "Legacy internal component config object."
image_encoder_configs: "Legacy internal component config objects."
text_encoder_configs: "Legacy internal component config object."
preprocess_text_funcs: "Internal text preprocessing hooks."
postprocess_text_funcs: "Internal text postprocessing hooks."
@@ -361,6 +368,30 @@ surfaces:
sources: [fastvideo.configs.pipelines.matrixgame2.MatrixGame2I2V480PConfig]
num_frames_per_block:
sources: [fastvideo.configs.pipelines.matrixgame2.MatrixGame2I2V480PConfig]
duration_s:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
spectrogram_frame_rate:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
latent_downsample_rate:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
clip_frame_rate:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
sync_frame_rate:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
sync_segment_size:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
sync_segment_stride:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
sync_downsample_rate:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
clip_image_size:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
sync_image_size:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
clip_batch_size_multiplier:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
sync_batch_size_multiplier:
sources: [fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig]
audio_channels:
sources:
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
@@ -375,6 +406,7 @@ surfaces:
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
max_audio_duration_s:
sources:
- fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
sample_size:
@@ -383,6 +415,7 @@ surfaces:
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
sampling_rate:
sources:
- fastvideo.configs.pipelines.mmaudio.MMAudioV2AConfig
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
audio_txt_guidance_scale:
+50 -5
View File
@@ -1,6 +1,10 @@
# MPS (Apple Silicon)
Instructions to install FastVideo for Apple Silicon.
Install FastVideo on Apple Silicon and run FastMetal-QAD.
Apple Silicon uses the MLX runtime and the FastMetal-QAD INT8 checkpoints.
See the [FastMetal-QAD blog](https://haoailab.com/blogs/fastmetal/) and the
[FastMetal collection](https://huggingface.co/collections/FastVideo/fastmetal).
## Requirements
@@ -49,7 +53,7 @@ brew install ffmpeg
### Installation
FastWan's native Apple Silicon runtime requires the `mlx` extra.
FastMetal's native Apple Silicon runtime requires the `mlx` extra.
#### With uv (recommended)
@@ -87,6 +91,47 @@ Alternative with Conda environment:
uv pip install -e ".[mlx]"
```
## Run FastMetal-QAD
Each release is self-contained. Download one checkpoint and point both
`--model-root` and `--mlx-checkpoint` at it (the example also auto-detects
`mlx_dit.json` under `--model-root`).
| Checkpoint | Script | Mac tier |
| --- | --- | --- |
| [`FastVideo/FastMetal-1.3B-QAD`](https://huggingface.co/FastVideo/FastMetal-1.3B-QAD) | `mlx_wan_prompt_to_video.py` | 16 GB+ |
| [`FastVideo/FastMetal-5B-QAD`](https://huggingface.co/FastVideo/FastMetal-5B-QAD) | `mlx_wan22_generate.py` | 16 GB+ |
| [`FastVideo/FastMetal-14B-QAD`](https://huggingface.co/FastVideo/FastMetal-14B-QAD) | `mlx_wan_prompt_to_video.py` | 36 GB+ |
```bash
hf download FastVideo/FastMetal-1.3B-QAD --local-dir ./FastMetal-1.3B-QAD
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--model-root ./FastMetal-1.3B-QAD \
--mlx-checkpoint ./FastMetal-1.3B-QAD \
--height 480 --width 832 --num-frames 81 \
--prompt "A bird's-eye view of a misty forest valley at dawn."
```
14B uses the same script. Point both flags at `./FastMetal-14B-QAD`. That repo also ships an EMA variant: keep `--model-root` at the repo root and set `--mlx-checkpoint ./FastMetal-14B-QAD/ema`.
Wan2.2 5B uses a different latent layout, so it has its own entrypoint:
```bash
hf download FastVideo/FastMetal-5B-QAD --local-dir ./FastMetal-5B-QAD
python examples/inference/basic/mlx_wan22_generate.py \
--mlx-checkpoint ./FastMetal-5B-QAD \
--text-encoder-root ./FastMetal-5B-QAD \
--vae-root ./FastMetal-5B-QAD/vae \
--height 704 --width 1280 --num-frames 81 \
--prompt "A cinematic portrait with soft neon lighting and smooth camera motion."
```
CUDA FastWan-QAD (`FastVideo/FastWan-QAD-1.3B`, `FastVideo/FastWan-QAD-FP8-1.3B`) is a separate NVIDIA release. The MLX examples look for FastMetal packed weights (`mlx_dit.json`).
`basic_mps.py` is a generic PyTorch MPS demo. For local video on Mac, use the FastMetal commands above.
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
@@ -94,9 +139,9 @@ If you're planning to contribute to FastVideo please see the following page:
## Hardware Requirements
### For Basic Inference
- Mac M1, M2, M3, or M4 (at least 32 GB RAM is preferable for high quality video generation)
- **1.3B / 5B:** 16 GB unified memory and up (M1 and later)
- **14B:** 36 GB unified memory and up
- Fanless 13-inch MacBook Air can run 1.3B and 5B at the same resolutions
## Troubleshooting
@@ -156,8 +156,11 @@ is power-cycled. To avoid it:
- **Builds** (flash-attn, kernel): `nice -n 19`, `MAX_JOBS=2`, `nohup`. Never a
bare foreground high-parallelism build.
- Leave `*_cpu_offload` at the example defaults — "CPU" offload is the *same*
unified RAM on the GB10, so the win is tiling + sane resolution, not offloading.
- FastVideo automatically disables DiT layerwise/CPU offload and encoder/VAE CPU
offload after each worker binds its GB10 device. Do not force those modes back
on: "CPU" offload uses the same unified RAM. Multi-GPU FSDP sharding remains
available because it partitions weights without parking them in a separate
host pool.
## Gotchas specific to the GB10
+7
View File
@@ -14,6 +14,13 @@ vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
```
On unified-memory accelerators such as NVIDIA GB10 and Apple silicon, FastVideo
detects the selected device inside each worker and disables all five host-offload
modes before loading modules. Host and accelerator allocations share one physical
pool there, so offload adds transfers and duplicate residency instead of freeing
memory. CUDA FSDP sharding remains enabled when requested; MPS continues to
disable FSDP. `pin_cpu_memory` is not an offload mode and is left unchanged.
## Behavior Explanation
!!! note
+78 -4
View File
@@ -88,6 +88,7 @@ runtime on some GPU/shape combinations. To use FA4, install the pinned
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
```bash
UV_TORCH_BACKEND=cu130 uv pip install -e ".[fasth3]"
export FASTVIDEO_FA4=1
```
@@ -97,6 +98,21 @@ sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
FastVideo fails loudly instead of silently falling back.
MiniMax-H3 can additionally use FA4's packed-varlen entry point for its long,
single-sequence dense DiT self-attention:
```bash
export FASTVIDEO_FA4=1
export FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN=1
```
This route is inference-only and remains disabled by default. Runtime guards
keep masked attention, batch sizes above one, unequal query/key lengths,
grad-enabled calls, and NVFP4 on their established paths. It does not apply to
the Preview checkpoint's sparse VSA blocks. Packed-varlen changes floating-point
reduction order relative to fixed-length FA4, so treat it as a speed/quality
evaluation option rather than an exact-parity mode.
### FP4 Flash Attention 4 (Blackwell only)
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
@@ -337,7 +353,42 @@ Only DiT submodules that declare `_compile_conditions` are compiled
(most shipped models). The text encoder and VAE are not compiled by this
flag.
### What to expect
### Regional fullgraph compile (experimental)
`inference_torch_compile` is a stricter, kwargs-free variant that ports the
training-side regional compile of
[#1718](https://github.com/hao-ai-lab/FastVideo/pull/1718) to inference: the
loader wraps each `_compile_conditions` block in
`torch.compile(fullgraph=True)` with inductor
`options={"emulate_precision_casts": True}` right after the transformer
loads. The ordinary compile path keeps its historical compiler-disabled
attention boundary by default; regional compile opts in only the compatible
attention instances owned by this transformer. MiniMax-H3 VSA is supported
only by the inference-only sm_100a tile-64 route
(`FASTVIDEO_VSA_SM100A=1` and `VSA_tile_size=64`); the loader probes that
route before capture and keeps the transformer eager when the kernel or
device is unsupported. Legacy VSA, MiniMax-H3 tile-256 VSA, and the explicit
`FASTVIDEO_DISABLE_ATTENTION_COMPILE=1` escape hatch keep the transformer
eager with one warning instead of failing mid-denoise.
```python
generator = VideoGenerator.from_pretrained(
"MiniMaxAI/MiniMax-H3",
inference_torch_compile=True, # or FASTVIDEO_INFERENCE_TORCH_COMPILE=1
)
```
Do not combine it with `torch_compile_kwargs['mode']` (the loader injects
inductor options, and torch.compile forbids mode+options); it is
independent of `enable_torch_compile`, and when both are set the regional
compile wins for the DiT.
### What to expect from generic compile
The Wan result below measures the existing generic
`enable_torch_compile=True` path. It is useful evidence that compile can help,
but it is **not** a benchmark or numerical gate for the stricter regional
fullgraph path above.
| Config | Effect |
|---|---|
@@ -353,6 +404,24 @@ generations with the same input shapes. Always exclude the first
(warmup) generation when measuring steady-state latency — measuring the
warmup is the most common way to wrongly conclude "compile is slower".
**Regional MiniMax-H3 accuracy caveat (job 2660).** On one GB200 at
768×1344×124, the native 50-point schedule ran exactly 49 transformer
forwards. After one warmup, three fixed-prompt/fixed-seed repeats averaged
**185.08s → 157.01s end to end** and **174.90s → 147.12s denoising**. Each
leg was independently pixel-deterministic, but compiled output did **not**
match eager: mean absolute pixel error **20.247/255**, PSNR **16.67 dB**,
mean SSIM **0.7108**, and mean MS-SSIM **0.6370** across 124 frames. Treat
regional MiniMax-H3 compile as an opt-in performance experiment, not an
eager-parity-safe mode.
The same caveat applies to sparse MiniMax-H3 regional compile. Its mask
compaction, sm_100a launch, trained compression gates, and inference-only H3
fusions are fullgraph-compatible, but compilation can still change model
numerics. The FastH3 `all` profile enables regional compile by default because
it is the fastest measured route at 124, 243, and 345 frames. Use
`--no-inference-torch-compile` when comparing against the eager sparse-DiT
route.
**Numerics.** Inductor's lowering is designed to preserve eager
semantics within floating-point tolerance, but per-model equivalence is
not asserted by any standing SSIM regression here — the SSIM tests in
@@ -371,9 +440,14 @@ MS-SSIM gate on *your* config, especially when combining
hao-ai-lab/FastVideo#1365 — keep that fix to get a clean compiled
region under the default offload path.
- **`mode="reduce-overhead"` / CUDA graphs**: not yet supported
end-to-end. The attention dispatch is an untraceable custom op and
still breaks the graph, which CUDA-graph trees cannot span. Use the
default inductor mode (shown above) until that is resolved.
by regional compile because that path injects inductor `options`, and
PyTorch rejects `mode` together with `options`. The generic compile path
can accept `mode`, but CUDA-graph compatibility remains backend- and
shape-dependent. FA2/FA3 inference and FA4 expose traceable custom-op
boundaries. MiniMax-H3's sm_100a tile-64 inference route is the only VSA
path in the regional support envelope; other VSA paths and the FA3
grad-enabled path remain outside it. Use the default inductor mode shown
above unless your exact configuration has its own gate.
Extra `torch.compile` options are passed through `torch_compile_kwargs`
(a dict), accepted by `VideoGenerator.from_pretrained(...)` and by the
+11 -8
View File
@@ -182,12 +182,14 @@ optimizations: absence means **untested**, not incompatible.
| 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 |
| MLX FastMetal T2V 1.3B | [`FastVideo/FastMetal-1.3B-QAD`](https://huggingface.co/FastVideo/FastMetal-1.3B-QAD) | 480x832, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 16 GB+ unified memory | Released |
| MLX FastMetal TI2V 5B | [`FastVideo/FastMetal-5B-QAD`](https://huggingface.co/FastVideo/FastMetal-5B-QAD) | 480p / 720p, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 16 GB+ unified memory | Released |
| MLX FastMetal T2V 14B | [`FastVideo/FastMetal-14B-QAD`](https://huggingface.co/FastVideo/FastMetal-14B-QAD) | 480p / 720p, 81 frames, 3-step DMD, INT8 DiT + TAEHV decode | Apple M4 Max, 36 GB+ unified memory | Released |
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.
Apple Silicon uses FastMetal-QAD. CUDA FastWan-QAD (`FastVideo/FastWan-QAD-1.3B`,
`FastVideo/FastWan-QAD-FP8-1.3B`) is the NVIDIA release. See the
[Apple Silicon guide](../getting_started/installation/mps.md) and the
[FastMetal-QAD blog](https://haoailab.com/blogs/fastmetal/).
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
@@ -216,9 +218,10 @@ Per the installation guides:
[GPU install guide](../getting_started/installation/gpu.md).
- **NVIDIA DGX Spark (GB10, aarch64)** — CUDA 13, from-source kernel build; see
the [DGX Spark install guide](../getting_started/installation/spark.md).
- **Apple silicon (MPS)** — macOS 14 or newer; see the
[MPS install guide](../getting_started/installation/mps.md) and
[`basic_mps.py`](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mps.py).
- **Apple silicon** — macOS 14 or newer; FastMetal-QAD via the MLX runtime. See the
[Apple Silicon guide](../getting_started/installation/mps.md). The older
[`basic_mps.py`](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mps.py)
demo is PyTorch MPS only.
Optimization-specific hardware constraints (e.g. STA requiring Hopper) are
listed under [Special requirements](#special-requirements).
+83 -3
View File
@@ -18,11 +18,27 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
python examples/inference/basic/basic.py
```
For an example on Apple silicon:
```
python examples/inference/basic/basic_mps.py
### Apple Silicon (FastMetal-QAD)
Use the MLX runtime with FastMetal-QAD. See the
[Apple Silicon guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation/mps/).
```bash
hf download FastVideo/FastMetal-1.3B-QAD --local-dir ./FastMetal-1.3B-QAD
python examples/inference/basic/mlx_wan_prompt_to_video.py \
--model-root ./FastMetal-1.3B-QAD \
--mlx-checkpoint ./FastMetal-1.3B-QAD \
--prompt "A bird's-eye view of a misty forest valley at dawn."
```
5B uses
[`mlx_wan22_generate.py`](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/mlx_wan22_generate.py)
with
`FastVideo/FastMetal-5B-QAD`.
`examples/inference/basic/basic_mps.py` is the older PyTorch MPS demo.
For an example running DMD+VSA inference:
```
python examples/inference/basic/basic_dmd.py
@@ -33,6 +49,70 @@ For the typed config/request path added during the inference API refactor:
python examples/inference/basic/basic_dmd_new_api.py
```
### FastH3 Preview
The verified [basic FastH3 example](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_fasth3.py)
runs the few-step (4-forward, DMD2-distilled) MiniMax-H3 preview, generating
synchronized video and audio with its trained block-sparse VSA attention:
```bash
UV_TORCH_BACKEND=cu130 uv pip install -e ".[fasth3]"
```
This installs the pinned FA4 CuTe package and FastVideo kernel release used by
the measured GB200 profile. Then run:
```
python examples/inference/basic/basic_fasth3.py --prompt "your prompt"
```
The default checkpoint, [FastH3 Preview v0.2](https://huggingface.co/FastVideo/FastVideo-Minimax-FastH3-Preview-v0.2), is public on the Hub under the MiniMax H3 Community License. Review its model card and license before use or redistribution.
The default `all` profile is the fastest measured four-GPU Preview recipe on GB200. It selects VSA sparsity 0.9 with 64-token tiles and the sm_100a sparse kernel, enables FA4 for eligible non-VSA paths, regionally compiles and replicates the sparse DiT, compiles and temporally parallelizes the video VAE with the `gather` strategy, and pins CPU-offloaded component memory. It also pins the benchmark protocol: five sigma-grid points (exactly four DiT forwards), one excluded seed-999 warmup, then three timed seed-1000 requests with distinct output paths.
The equivalent explicit command is:
```bash
python examples/inference/basic/basic_fasth3.py \
--prompt "your prompt" \
--profile all \
--num-gpus 4 \
--steps 5 \
--vsa-sparsity 0.9 \
--vsa-tile-size 64 \
--vsa-kernel sm100a \
--compile-vae \
--parallel-vae \
--replicated-dit \
--pin-cpu-memory \
--fa4 \
--no-torch-compile \
--inference-torch-compile \
--ulysses-a2a off \
--warmup \
--repeats 3 \
--seed 1000 \
--warmup-seed 999
```
`all` enables the inference-only H3 fusions and regional compile. Both can change floating-point operation order, so this is a report-only performance profile rather than an exact-parity route. Use `--profile strict` to disable the H3 fusions while preserving regional compile, or `--profile strict --no-inference-torch-compile` for the eager strict route. Individual `--no-*` switches are available for portability and attribution; in particular, use `--vsa-kernel triton --no-fa4` if the Blackwell kernels are unavailable. The script preserves the warmup and each measured video under distinct paths, then prints per-request wall time plus a warmup-excluded median.
One script covers each validated duration; regional compile is the fastest
measured DiT route for all three:
```bash
# 5 s
python examples/inference/basic/basic_fasth3.py \
--prompt "your prompt" --output outputs/fasth3_5s
# 10 s
python examples/inference/basic/basic_fasth3.py \
--prompt "your prompt" --num-frames 243 --output outputs/fasth3_10s
# 15 s
python examples/inference/basic/basic_fasth3.py \
--prompt "your prompt" --num-frames 345 --output outputs/fasth3_15s
```
Pass `--no-inference-torch-compile` to recover the eager sparse-DiT route.
## Basic Walkthrough
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
+341
View File
@@ -0,0 +1,341 @@
# SPDX-License-Identifier: Apache-2.0
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
The default ``all`` profile reproduces the fastest measured FastH3 Preview
recipe on four GB200 GPUs. It runs the checkpoint's native five-point sigma
grid (exactly four DiT forwards), trained VSA policy, Blackwell sparse kernel,
regional fullgraph DiT compile, compiled/parallel video VAE, and inference-only
H3 fusions. One compile warmup is excluded before three measured requests.
Both regional compile and the default fusions can change floating-point
operation order, so ``all`` is a report-only performance profile.
``--profile strict`` disables the H3 fusions but preserves regional compile;
combine it with ``--no-inference-torch-compile`` for the eager strict route.
"""
from __future__ import annotations
import argparse
import importlib.util
import os
import statistics
import time
from collections.abc import Sequence
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
DEFAULT_MODEL = "FastVideo/FastVideo-Minimax-FastH3-Preview-v0.2"
def parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-path", default=DEFAULT_MODEL)
# The HF repo may require authentication while the MiniMax H3 Community
# License review completes. A local snapshot can be passed here instead.
parser.add_argument("--prompt", required=True)
parser.add_argument("--output", default="outputs/fasth3")
parser.add_argument("--profile",
choices=("all", "strict"),
default="all",
help="all enables the fastest measured, non-parity H3 fusions; strict disables only them")
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. The distilled schedule is
# t=1000,750,500,250 -> 0: five points and exactly four DiT forwards.
parser.add_argument("--steps",
type=int,
default=5,
help="sigma-grid points; N points run N-1 DiT forwards (the trained default is 5)")
parser.add_argument("--seed", type=int, default=1000, help="seed reused for every measured request")
parser.add_argument("--warmup-seed", type=int, default=999)
parser.add_argument("--repeats", type=int, default=3, help="number of measured requests after warmup")
parser.add_argument("--warmup",
action=argparse.BooleanOptionalAction,
default=True,
help="run one excluded request before timing")
parser.add_argument("--num-gpus", type=int, default=4)
parser.add_argument("--vsa-sparsity",
type=float,
default=0.9,
help="run-level VSA sparsity in [0, 1); 0.9 is the checkpoint's trained policy")
parser.add_argument("--vsa-tile-size",
type=int,
choices=(64, 256),
default=64,
help="VSA-H3 tile size; 64 is the checkpoint's trained and measured geometry")
parser.add_argument("--vsa-kernel",
choices=("triton", "sm100a"),
default="sm100a",
help="tile-64 sparse kernel; sm100a is the measured GB200 route and requires a compatible "
"fastvideo-kernel build")
parser.add_argument("--fa4",
action=argparse.BooleanOptionalAction,
default=True,
help="use FA4 for eligible non-VSA attention paths")
parser.add_argument("--h3-fusions",
action=argparse.BooleanOptionalAction,
default=None,
help="override the profile's H3 fusion policy (changes model numerics when enabled)")
parser.add_argument("--compile-vae",
action=argparse.BooleanOptionalAction,
default=True,
help="compile the video VAE decoder independently of the DiT")
parser.add_argument("--parallel-vae",
action=argparse.BooleanOptionalAction,
default=True,
help="round-robin VAE temporal chunks across sequence-parallel ranks")
parser.add_argument("--replicated-dit",
action=argparse.BooleanOptionalAction,
default=True,
help="replicate DiT weights instead of FSDP-sharding them")
parser.add_argument("--pin-cpu-memory",
action=argparse.BooleanOptionalAction,
default=True,
help="pin CPU-offloaded text-encoder and VAE weights")
parser.add_argument("--torch-compile",
action=argparse.BooleanOptionalAction,
default=False,
help="compile the whole DiT path (off in the fastest FastH3 profile)")
parser.add_argument("--inference-torch-compile",
action=argparse.BooleanOptionalAction,
default=True,
help="regionally compile DiT blocks (enabled in the fastest FastH3 profile)")
parser.add_argument("--ulysses-a2a",
choices=("off", "auto"),
default="off",
help="sequence-parallel all-to-all route; off reproduces the fastest FastH3 profile, while "
"auto opts into the fused NVLink kernel when the installed kernel package supports it")
parser.add_argument("--compile-mode",
default=None,
help='whole-DiT torch.compile mode, e.g. "reduce-overhead"; requires '
"--no-inference-torch-compile")
args = parser.parse_args(argv)
if args.repeats < 1:
parser.error("--repeats must be at least 1")
if args.num_gpus < 1:
parser.error("--num-gpus must be at least 1")
if not 0.0 <= args.vsa_sparsity < 1.0:
parser.error("--vsa-sparsity must be in [0, 1)")
if args.compile_mode is not None and args.inference_torch_compile:
parser.error("--compile-mode cannot be combined with regional compile; pass --no-inference-torch-compile")
return args
def _h3_fusions_enabled(args: argparse.Namespace) -> bool:
if args.h3_fusions is not None:
return bool(args.h3_fusions)
return args.profile == "all"
def profile_environment(args: argparse.Namespace) -> dict[str, str | None]:
"""Return the complete boot-time environment for this profile.
``None`` means the variable must be removed. Values are explicit even for
disabled features so a shell's inherited experiment settings cannot
silently change the advertised profile.
"""
return {
"FASTVIDEO_ATTENTION_BACKEND": "VIDEO_SPARSE_ATTN_H3",
"FASTVIDEO_VSA_SM100A": "1" if args.vsa_kernel == "sm100a" else "0",
"FASTVIDEO_VSA_CUTEDSL": "0",
# A non-empty output path enables the diagnostic probe.
"FASTVIDEO_H3_VSA_PROBE": None,
"FASTVIDEO_DISABLE_ATTENTION_COMPILE": "0",
"FASTVIDEO_FA4": "1" if args.fa4 else "0",
"FASTVIDEO_NVFP4_FA4": "0",
"FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN": "0",
"FASTVIDEO_MINIMAX_H3_FUSIONS": "all" if _h3_fusions_enabled(args) else "0",
"FASTVIDEO_INFERENCE_TORCH_COMPILE": "1" if args.inference_torch_compile else "0",
"FASTVIDEO_VAE_PARALLEL_DECODE": "1" if args.parallel_vae else "0",
"FASTVIDEO_VAE_PARALLEL_ENCODE": "0",
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY": "gather",
"FASTVIDEO_ULYSSES_A2A": args.ulysses_a2a,
"FASTVIDEO_STAGE_LOGGING": "1",
}
def configure_environment(args: argparse.Namespace) -> dict[str, str | None]:
environment = profile_environment(args)
for name, value in environment.items():
if value is None:
os.environ.pop(name, None)
else:
os.environ[name] = value
return environment
def _fa4_is_installed() -> bool:
try:
return importlib.util.find_spec("flash_attn.cute") is not None
except (ImportError, ModuleNotFoundError):
return False
def _sm100a_kernel_is_installed() -> bool:
try:
from fastvideo_kernel import block_sparse_attn_sm100a
except ImportError:
return False
return bool(getattr(block_sparse_attn_sm100a, "_HAS_VSA_SM100A", False))
def validate_profile_dependencies(args: argparse.Namespace) -> None:
"""Fail before model loading when the selected measured route is absent."""
if args.fa4 and not _fa4_is_installed():
raise RuntimeError(
"FastH3's FA4 profile requires the pinned flash-attn-4 package. Install it with "
"`UV_TORCH_BACKEND=cu130 uv pip install -e \".[fasth3]\"`, or pass --no-fa4.")
if args.vsa_kernel == "sm100a" and not _sm100a_kernel_is_installed():
raise RuntimeError(
"FastH3's sm100a profile requires fastvideo-kernel 0.3.4 built with the Blackwell VSA extension. "
"Install this checkout with `UV_TORCH_BACKEND=cu130 uv pip install -e \".[fasth3]\"` (or run "
"`cd fastvideo-kernel && ./build.sh`), or pass --vsa-kernel triton.")
def build_generator_config(args: argparse.Namespace) -> GeneratorConfig:
experimental: dict[str, object] = {
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
"VSA_sparsity": args.vsa_sparsity,
"VSA_tile_size": args.vsa_tile_size,
"inference_torch_compile": args.inference_torch_compile,
"vae_parallel_decode": args.parallel_vae,
"vae_parallel_decode_strategy": "gather",
}
return GeneratorConfig(
model_path=args.model_path,
pipeline=PipelineSelection(experimental=experimental),
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=args.num_gpus > 1 and not args.replicated_dit,
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=args.pin_cpu_memory,
),
compile=CompileConfig(
enabled=args.torch_compile,
mode=args.compile_mode,
vae_enabled=args.compile_vae,
),
),
)
def build_request(args: argparse.Namespace, output_path: Path, seed: int) -> GenerationRequest:
return 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,
# MiniMax-H3 is guidance-distilled; FastH3 inherits that contract.
guidance_scale=1.0,
batch_cfg=False,
seed=seed,
),
output=OutputConfig(
output_path=str(output_path),
save_video=True,
return_frames=False,
),
)
def _actual_output_path(result: object, requested: Path) -> Path:
video_path = getattr(result, "video_path", None)
return Path(video_path) if video_path else requested
def _denoise_seconds(result: object) -> float | None:
stages = getattr(getattr(result, "logging_info", None), "stages", None)
if not stages:
return None
for stage_name, metrics in stages.items():
if "denois" not in stage_name.lower():
continue
execution_time = metrics.get("execution_time")
return float(execution_time) if execution_time is not None else None
return None
def run(args: argparse.Namespace) -> list[float]:
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
environment = configure_environment(args)
validate_profile_dependencies(args)
print(f"Profile: {args.profile} ({'non-parity fusions' if _h3_fusions_enabled(args) else 'fusions off'})")
print(f"Output directory: {output_dir.resolve()}")
print("Denoising contract: 5 sigma points = 4 DiT forwards" if args.steps == 5 else
f"Denoising contract override: {args.steps} sigma points = {args.steps - 1} DiT forwards")
print("Profile environment: " + " ".join(f"{key}={value if value is not None else '<unset>'}"
for key, value in environment.items()))
generator = VideoGenerator.from_config(build_generator_config(args))
measured_wall_times: list[float] = []
measured_denoise_times: list[float] = []
try:
if args.warmup:
warmup_path = output_dir / "_fasth3_warmup.mp4"
print(f"[warmup] generating (excluded from timing summary): {warmup_path}")
started = time.perf_counter()
warmup_result = generator.generate(build_request(args, warmup_path, args.warmup_seed))
warmup_wall = time.perf_counter() - started
actual_warmup_path = _actual_output_path(warmup_result, warmup_path)
print(f"[warmup] wall={warmup_wall:.3f}s (excluded)")
print(f"Warmup output written to: {actual_warmup_path}")
for index in range(1, args.repeats + 1):
requested_path = output_dir / f"fasth3_{args.profile}_run_{index:02d}.mp4"
print(f"[measured {index}/{args.repeats}] generating: {requested_path}")
started = time.perf_counter()
result = generator.generate(build_request(args, requested_path, args.seed))
wall = time.perf_counter() - started
measured_wall_times.append(wall)
actual_path = _actual_output_path(result, requested_path)
print(f"Output written to: {actual_path}")
print(f"E2E wall time: {wall:.3f}s")
generation_time = getattr(result, "generation_time", None)
if generation_time is not None:
print(f"Generation time: {float(generation_time):.3f}s")
denoise_time = _denoise_seconds(result)
if denoise_time is not None:
measured_denoise_times.append(denoise_time)
print(f"Denoising time: {denoise_time:.3f}s")
median = statistics.median(measured_wall_times)
print(f"Measured E2E wall times (n={len(measured_wall_times)}, warmup excluded): "
f"{[round(value, 3) for value in measured_wall_times]}")
print(f"Median E2E wall time: {median:.3f}s")
if measured_denoise_times:
print(f"Median denoising time: {statistics.median(measured_denoise_times):.3f}s")
return measured_wall_times
finally:
generator.shutdown()
def main() -> None:
run(parse_args())
if __name__ == "__main__":
main()
@@ -15,6 +15,7 @@ from fastvideo.api import (
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
@@ -41,6 +42,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load (the #1718 "
"training-port semantics: no kwargs; fullgraph + emulate_precision_casts injected). "
"First generation pays the inductor JIT (~1-2 min); use --repeats >= 2 and time "
"the last repeat. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent")
parser.add_argument("--repeats",
type=int,
default=1,
@@ -54,9 +61,16 @@ def main() -> None:
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
# Boot-time run configuration folded into FastVideoArgs (the same
# experimental-dict route basic_fasth3.py uses for the VSA knobs).
experimental: dict[str, object] = {}
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
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,
+29 -6
View File
@@ -1,14 +1,18 @@
# SPDX-License-Identifier: Apache-2.0
"""End-to-end Wan2.2-TI2V-5B generation on Apple Silicon (MLX DiT + MLX TAEHV).
"""End-to-end FastMetal-5B-QAD generation on Apple Silicon (MLX DiT + MLX TAEHV).
This is the Wan2.2 TI2V entrypoint. Use FastVideo/FastMetal-5B-QAD:
hf download FastVideo/FastMetal-5B-QAD --local-dir ./FastMetal-5B-QAD
python examples/inference/basic/mlx_wan22_generate.py \\
--mlx-checkpoint ./FastMetal-5B-QAD \\
--text-encoder-root ./FastMetal-5B-QAD \\
--vae-root ./FastMetal-5B-QAD/vae
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).
"""
@@ -31,6 +35,11 @@ from fastvideo.mlx_runtime.prompt_cache import (
save_prompt_cache,
text_encoder_fingerprint,
)
from fastvideo.mlx_runtime.checkpoint_compat import (
UnsupportedMLXCheckpointError,
raise_if_unsupported_mlx_checkpoint,
resolve_mlx_checkpoint,
)
from fastvideo.mlx_runtime.rife_interp import aligned_keyframe_count
FASTWAN21_MODEL_ID = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
@@ -165,7 +174,9 @@ def main() -> None:
"--mlx-checkpoint",
type=Path,
default=None,
help="Pre-quantized MLX DiT checkpoint directory. Rewrapped with Wan2.2 per-token conditioning.",
help="Packed FastMetal-5B-QAD MLX DiT directory (mlx_dit.json + mlx_dit.safetensors). "
"If omitted, a FastMetal directory passed as --text-encoder-root is used when it "
"already contains those files.",
)
parser.add_argument("--vae-root", type=Path, default=None)
parser.add_argument("--height", type=int, default=DEFAULT_HEIGHT)
@@ -241,6 +252,18 @@ def main() -> None:
# 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.mlx_checkpoint = resolve_mlx_checkpoint(args.mlx_checkpoint, args.text_encoder_root)
if args.mlx_checkpoint is not None:
if args.text_encoder_root is None and (args.mlx_checkpoint / "text_encoder").is_dir():
args.text_encoder_root = args.mlx_checkpoint
if args.vae_root is None and (args.mlx_checkpoint / "vae").is_dir():
args.vae_root = args.mlx_checkpoint / "vae"
try:
raise_if_unsupported_mlx_checkpoint(args.mlx_checkpoint, args.dit_checkpoint)
except UnsupportedMLXCheckpointError as exc:
raise SystemExit(str(exc)) from exc
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,
@@ -1,11 +1,25 @@
"""Generate a FastWan text-to-video clip with the Apple Silicon MLX runtime.
"""Generate a FastMetal text-to-video clip with the Apple Silicon MLX runtime.
This is the supported source-tree entrypoint for the FastWan-QAD-INT8-1.3B
Apple release:
This is the supported source-tree entrypoint for FastMetal-QAD (Wan2.1 1.3B
and 14B). Use ``mlx_wan22_generate.py`` for FastMetal-5B-QAD.
Download FastMetal-QAD and point ``--model-root`` / ``--mlx-checkpoint`` at it:
hf download FastVideo/FastMetal-1.3B-QAD --local-dir ./FastMetal-1.3B-QAD
python examples/inference/basic/mlx_wan_prompt_to_video.py \\
--model-root ./FastMetal-1.3B-QAD --mlx-checkpoint ./FastMetal-1.3B-QAD
CUDA FastWan-QAD (``FastVideo/FastWan-QAD-1.3B``, ``FastVideo/FastWan-QAD-FP8-1.3B``)
is a separate NVIDIA release.
FastMetal-QAD Hugging Face repos ship ``mlx_dit.json`` + ``mlx_dit.safetensors``,
not a Diffusers ``transformer/`` tree. Do not copy ``transformer/config.json``
from Wan2.1 or other checkpoints; point ``--mlx-checkpoint`` at the FastMetal
directory and the example reads the DiT config from ``mlx_dit.json``.
- Hugging Face/torch encodes the prompt with UMT5 (bf16 by default: fp32
exponent range without fp16 overflow risk, at fp16 memory cost).
- MLX runs the FastWan DiT denoising loop (INT8 by default, compiled with
- MLX runs the FastMetal DiT denoising loop (INT8 by default, compiled with
``mx.compile`` unless ``--no-mlx-compile``).
- TAEHV (default, fast/low-memory) or the full Wan VAE (``--decode-backend
wan-vae``, higher fidelity, bf16) decodes the final latents.
@@ -56,13 +70,18 @@ from fastvideo.mlx_runtime.prompt_cache import (
save_prompt_cache,
text_encoder_fingerprint,
)
from fastvideo.mlx_runtime.checkpoint_compat import (
UnsupportedMLXCheckpointError,
raise_if_unsupported_mlx_checkpoint,
resolve_mlx_checkpoint,
)
from fastvideo.mlx_runtime.rife_interp import aligned_keyframe_count
if TYPE_CHECKING: # pragma: no cover - typing only
from fastvideo.mlx_runtime.fast_spatial import FastSpatialPlan
DEFAULT_MODEL_ID = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
DEFAULT_MODEL_ID = "FastVideo/FastMetal-1.3B-QAD"
# Legacy pinned-snapshot location, kept for callers that import it (the MLX
# benchmark harness). New code should prefer resolve_model_root(None), which
@@ -99,6 +118,10 @@ def resolve_model_root(
"tokenizer/*",
"text_encoder/*",
"vae/*",
"mlx_dit.json",
"mlx_dit.safetensors",
"ema/mlx_dit.json",
"ema/mlx_dit.safetensors",
"transformer/*" if include_transformer else "transformer/config.json",
]
return Path(snapshot_download(
@@ -138,6 +161,7 @@ def encode_prompt(
text_encoder = UMT5EncoderModel.from_pretrained(
model_root / "text_encoder",
torch_dtype=dtype,
low_cpu_mem_usage=True,
local_files_only=True,
).to(device)
text_encoder.eval()
@@ -379,6 +403,7 @@ def decode_latents_to_video(
vae = AutoencoderKLWan.from_pretrained(
model_root / "vae",
torch_dtype=dtype,
low_cpu_mem_usage=True,
local_files_only=True,
).to(device)
vae.eval()
@@ -471,11 +496,16 @@ def _rife_interpolate_video(*, video_path: Path, target_frames: int, factor: int
def main() -> None:
parser = argparse.ArgumentParser(description="Prompt-to-video FastWan generation using MLX for the DiT")
parser = argparse.ArgumentParser(
description="Prompt-to-video FastMetal-QAD generation using the Apple Silicon MLX runtime")
parser.add_argument("--model-root", type=Path, default=None,
help=f"Model directory. Defaults to the local HF cache for {DEFAULT_MODEL_ID} "
help="FastMetal-QAD directory (tokenizer, UMT5, VAE, packed MLX DiT). "
f"Defaults to the local HF cache for {DEFAULT_MODEL_ID} "
"(downloading it if missing).")
parser.add_argument("--prompt", default="A paper boat sails through a shallow stream in a mossy forest.")
parser.add_argument(
"--prompt",
default="A bird's-eye view of a misty forest valley at dawn.",
)
parser.add_argument("--output-path", type=Path, default=Path("video_samples/mlx_fastwan_prompt_to_video.mp4"))
parser.add_argument("--height", type=int, default=480)
parser.add_argument("--width", type=int, default=832)
@@ -639,9 +669,8 @@ def main() -> None:
"keyed by (model, prompt, length, dtype), so repeat runs skip "
"the text encoder entirely. Default: on.")
parser.add_argument("--mlx-checkpoint", type=Path, default=None,
help="Load the DiT from a pre-quantized MLX checkpoint directory "
"(created with --save-mlx-checkpoint) instead of casting/quantizing "
"the Diffusers weights on every run.")
help="Packed FastMetal MLX DiT directory (mlx_dit.json + mlx_dit.safetensors). "
"Defaults to --model-root when that directory already contains those files.")
parser.add_argument("--save-mlx-checkpoint", type=Path, default=None,
help="After loading the DiT, save it (cast + quantized) as an MLX "
"checkpoint directory for fast reloads via --mlx-checkpoint.")
@@ -690,6 +719,12 @@ def main() -> None:
np.save(args.encode_prompt_only, prompt_embeds.cpu().numpy())
return
mlx_checkpoint = resolve_mlx_checkpoint(args.mlx_checkpoint, model_root)
try:
raise_if_unsupported_mlx_checkpoint(mlx_checkpoint or model_root)
except UnsupportedMLXCheckpointError as exc:
raise SystemExit(str(exc)) from exc
import mlx.core as mx
import torch
from diffusers import UniPCMultistepScheduler
@@ -719,17 +754,21 @@ def main() -> None:
config_path = model_root / "transformer/config.json"
checkpoint_path = model_root / "transformer/diffusion_pytorch_model.safetensors"
config = json.loads(config_path.read_text())
# A pre-quantized MLX DiT can be paired with a lightweight asset root for
# UMT5/TAEHV. In that case the model root is *not* the architecture
# authority: use the checkpoint's embedded transformer config for the
# sampler guard and latent geometry.
dit_config = config
if args.mlx_checkpoint is not None:
mlx_config_path = Path(args.mlx_checkpoint) / "mlx_dit.json"
if mlx_config_path.is_file():
mlx_checkpoint_config = json.loads(mlx_config_path.read_text())
dit_config = mlx_checkpoint_config.get("config", mlx_checkpoint_config)
# Packed FastMetal checkpoints are the architecture authority. Do not
# require transformer/config.json when mlx_dit.json is already present.
if mlx_checkpoint is not None:
mlx_checkpoint_config = json.loads((mlx_checkpoint / "mlx_dit.json").read_text())
dit_config = mlx_checkpoint_config.get("config", mlx_checkpoint_config)
config = dit_config
else:
if not config_path.is_file():
raise SystemExit(
f"No packed MLX DiT (mlx_dit.json) and no Diffusers transformer config at {config_path}. "
"FastMetal-QAD checkpoints intentionally omit transformer/; download "
"FastVideo/FastMetal-1.3B-QAD and pass --model-root / --mlx-checkpoint at that directory."
)
config = json.loads(config_path.read_text())
dit_config = config
if int(dit_config.get("in_channels", 0)) == 48 and int(dit_config.get("out_channels", 0)) == 48:
raise SystemExit(
"Wan2.2-TI2V-5B uses 48-channel, per-token timestep conditioning. "
@@ -835,10 +874,10 @@ def main() -> None:
load_start = time.perf_counter()
mx.clear_cache()
mx.reset_peak_memory()
if args.mlx_checkpoint is not None:
if mlx_checkpoint is not None:
from fastvideo.mlx_runtime.checkpoint import load_mlx_dit_checkpoint
dit = load_mlx_dit_checkpoint(args.mlx_checkpoint, compile=args.mlx_compile)
dit = load_mlx_dit_checkpoint(mlx_checkpoint, compile=args.mlx_compile)
config = dit.config
else:
dit = mlx_dit_from_diffusers_safetensors(
+71 -1
View File
@@ -114,7 +114,63 @@ list(APPEND TORCH_INCLUDE_DIRS ${TORCH_INCLUDE_PATHS})
# Find Torch package (still useful for libraries)
find_package(Torch REQUIRED)
# Include directories
# The Ulysses kernel needs NCCL's 2.29 device API. Keep it optional so ROCm,
# older NCCL installs, and minimal CUDA builders still produce a usable wheel.
set(FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A "AUTO" CACHE STRING
"Build the NCCL-device-API Ulysses all-to-all: AUTO/ON/OFF")
set_property(CACHE FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A PROPERTY STRINGS AUTO ON OFF)
set(ENABLE_ULYSSES_A2A OFF)
if(NOT GPU_BACKEND STREQUAL "ROCM" AND NOT FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A STREQUAL "OFF")
# nvidia.nccl is a namespace package, so find_spec rather than __file__.
execute_process(COMMAND "${Python_EXECUTABLE}" -c
"import importlib.util as u;s=u.find_spec('nvidia.nccl');print(list(s.submodule_search_locations)[0] if s else '')"
OUTPUT_VARIABLE NCCL_PIP_ROOT OUTPUT_STRIP_TRAILING_WHITESPACE ERROR_QUIET)
find_path(NCCL_INCLUDE_DIR nccl_device/coop.h
HINTS $ENV{NCCL_HOME}/include ${NCCL_PIP_ROOT}/include)
find_library(NCCL_LIBRARY NAMES nccl
HINTS $ENV{NCCL_HOME}/lib $ENV{NCCL_HOME}/lib64 ${NCCL_PIP_ROOT}/lib)
# PyPI's nvidia-nccl-cu* wheels ship the SONAME but not the development
# symlink (libnccl.so.2, no libnccl.so). Accept that exact versioned name;
# target_link_libraries can link an absolute SONAME path directly.
if(NOT NCCL_LIBRARY)
find_file(NCCL_VERSIONED_LIBRARY NAMES libnccl.so.2
HINTS $ENV{NCCL_HOME}/lib $ENV{NCCL_HOME}/lib64 ${NCCL_PIP_ROOT}/lib)
if(NCCL_VERSIONED_LIBRARY)
set(NCCL_LIBRARY "${NCCL_VERSIONED_LIBRARY}")
endif()
endif()
if(NCCL_INCLUDE_DIR AND NCCL_LIBRARY)
include(CheckCXXSourceCompiles)
set(_FASTVIDEO_REQUIRED_INCLUDES "${CMAKE_REQUIRED_INCLUDES}")
# nccl.h includes cuda_runtime.h, so a host-compiler feature probe
# needs the toolkit includes explicitly even though CUDA is enabled.
set(CMAKE_REQUIRED_INCLUDES "${NCCL_INCLUDE_DIR};${CUDAToolkit_INCLUDE_DIRS}")
unset(NCCL_HAS_REQUIRED_DEVICE_API CACHE)
check_cxx_source_compiles("\
#define NCCL_HOSTLIB_ONLY
#include <cstddef>
#include <nccl_device.h>
int main() {
ncclDevCommRequirements reqs = NCCL_DEV_COMM_REQUIREMENTS_INITIALIZER;
ncclCommProperties props = NCCL_COMM_PROPERTIES_INITIALIZER;
return (reqs.size == 0 || props.size == 0);
}"
NCCL_HAS_REQUIRED_DEVICE_API)
set(CMAKE_REQUIRED_INCLUDES "${_FASTVIDEO_REQUIRED_INCLUDES}")
endif()
if(NCCL_HAS_REQUIRED_DEVICE_API)
set(ENABLE_ULYSSES_A2A ON)
elseif(FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A STREQUAL "ON")
message(FATAL_ERROR
"FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A=ON requires NCCL device headers/library "
"with NCCL_DEV_COMM_REQUIREMENTS_INITIALIZER (NCCL 2.29+). "
"Resolved include='${NCCL_INCLUDE_DIR}', library='${NCCL_LIBRARY}'.")
endif()
endif()
include_directories(
${CMAKE_SOURCE_DIR}/include
${CMAKE_SOURCE_DIR}/include/cutlass/include
@@ -124,6 +180,9 @@ include_directories(
${CMAKE_SOURCE_DIR}/csrc/turbodiffusion
${TORCH_INCLUDE_DIRS}
)
if(ENABLE_ULYSSES_A2A)
include_directories(${NCCL_INCLUDE_DIR})
endif()
# ---------------------------
# ThunderKittens (TK) toggles
@@ -157,6 +216,7 @@ endif()
message(STATUS "TORCH_CUDA_ARCH_LIST (cmake/env): ${TORCH_CUDA_ARCH_LIST}")
message(STATUS "FASTVIDEO_KERNEL_BUILD_TK: ${FASTVIDEO_KERNEL_BUILD_TK}")
message(STATUS "FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER: ${FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER}")
message(STATUS "FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A: ${FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A}")
set(ENABLE_TK_KERNELS OFF)
if(FASTVIDEO_KERNEL_BUILD_TK STREQUAL "ON")
@@ -307,6 +367,9 @@ if(BUILD_CXX_KERNELS)
csrc/turbodiffusion/norm/layernorm.cu
csrc/turbodiffusion/quant/quant.cu
)
if(ENABLE_ULYSSES_A2A)
list(APPEND EXTENSION_SOURCES csrc/comm/ulysses_all_to_all.cu)
endif()
# Conditionally add TK kernels
if(ENABLE_TK_KERNELS)
@@ -371,6 +434,9 @@ if(BUILD_CXX_KERNELS)
if(ENABLE_TK_KERNELS)
list(APPEND COMPILE_DEFS TK_COMPILE_ST_ATTN TK_COMPILE_BLOCK_SPARSE)
endif()
if(ENABLE_ULYSSES_A2A)
list(APPEND COMPILE_DEFS FASTVIDEO_KERNEL_COMPILE_ULYSSES_A2A)
endif()
target_compile_definitions(fastvideo_kernel_ops PRIVATE ${COMPILE_DEFS})
@@ -382,6 +448,9 @@ if(BUILD_CXX_KERNELS)
# Link against Torch libraries to avoid undefined symbols at import time
# (e.g., torch::autograd vtables) when loading the extension module.
target_link_libraries(fastvideo_kernel_ops PRIVATE ${TORCH_LIBRARIES})
if(ENABLE_ULYSSES_A2A)
target_link_libraries(fastvideo_kernel_ops PRIVATE "${NCCL_LIBRARY}")
endif()
# Also link against libtorch_python to satisfy Python-binding symbols
# (e.g., torch::PyWarningHandler) required by torch/extension.h.
@@ -485,6 +554,7 @@ message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
message(STATUS " + Ulysses NCCL-device all-to-all: ${ENABLE_ULYSSES_A2A}")
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
message(STATUS "============================================================")
+14 -9
View File
@@ -10,6 +10,8 @@ Compiled CUDA extensions (CMake, see the build summary printed at the end of eve
|---|---|---|---|---|
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
| same extension, optional part | MiniMax-H3 block-sparse VSA forward (64- and 128-token blocks) | `csrc/attention/block_sparse*_sm100a.cu` | Blackwell `sm_100a` only | ON iff `10.0a` is in `TORCH_CUDA_ARCH_LIST` |
| same extension, optional part | fused NVLink Ulysses all-to-all | `csrc/comm/ulysses_all_to_all.cu` | CUDA | `FASTVIDEO_KERNEL_BUILD_ULYSSES_A2A` (AUTO = ON with NCCL 2.29+ device headers and library; always OFF on ROCm) |
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
Runtime-JIT kernels (no build step, ship in every wheel/image):
@@ -22,20 +24,23 @@ Runtime-JIT kernels (no build step, ship in every wheel/image):
## What gets built where, and when
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|---|---|---|---|---|---|
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | Ulysses | FP4 |
|---|---|---|---|---|---|---|
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | AUTO | — (CUDA < 12.8) |
| | | x86_64 cu130 | `9.0a;10.0a;12.0a` | ON | AUTO | ON |
| | | aarch64 cu130 | `10.0a;12.0a` | — | AUTO | ON |
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | AUTO | — |
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | AUTO | — |
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | AUTO | — |
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | AUTO | ON iff sm_120 |
Notes:
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
- Both cu130 PyPI wheels ship the MiniMax-H3 sm_100a VSA forward.
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
- Ulysses AUTO builds only when CMake finds a NCCL library and device-API headers with the 2.29 initializers. Use
`-DFASTVIDEO_KERNEL_BUILD_ULYSSES_A2A=ON` to require it or `OFF` to test the portable build.
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
## Installation
@@ -0,0 +1,279 @@
/*
* Copyright (c) 2025 by FlashInfer team.
*
* Adapted from flashinfer-ai/flashinfer @ 8a94642d83cba0939035868fb6c309b4474a13d6
* (PR #3820), csrc/ulysses_all_to_all.cu.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
// Torch host bindings for the fused Ulysses all-to-all. Kernel in
// include/comm/ulysses_all_to_all.cuh.
//
// The per-group context is an ncclDevComm plus a registered symmetric window,
// both created here from the caller's ncclComm_t.
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAStream.h>
#include <torch/extension.h>
#include <algorithm>
#include <cstdint>
#include <memory>
#include <vector>
#include <nccl.h>
#include <nccl_device.h>
#include "comm/ulysses_all_to_all.cuh"
namespace fi = fastvideo::comm::ulysses;
namespace {
// A symmetric window every rank can store into, plus the ncclDevComm the
// kernel opens barrier sessions on.
struct UlyssesContext {
ncclComm_t comm = nullptr;
ncclWindow_t win = nullptr;
void* buf = nullptr;
size_t nbytes = 0;
ncclDevComm devComm{};
bool dev_comm_created = false;
int device = -1;
int rank = 0;
int world = 0;
};
#define NCCL_TRY(expr, what) \
do { \
ncclResult_t _r = (expr); \
TORCH_CHECK(_r == ncclSuccess, what " failed: rc=", static_cast<int>(_r)); \
} while (0)
} // namespace
// Allocate the local half of a context. This is intentionally separate from
// registration so Python can vote after local allocation: if one rank is OOM,
// no peer enters a collective window registration alone.
int64_t allocate_ulysses_a2a(int64_t nbytes, int64_t rank, int64_t world_size,
int64_t device_index) {
TORCH_CHECK(world_size == 2 || world_size == 4 || world_size == 6 || world_size == 8,
"ulysses a2a only supports world size in (2, 4, 6, 8), got ", world_size);
TORCH_CHECK(rank >= 0 && rank < world_size, "invalid rank");
TORCH_CHECK(nbytes > 0, "nbytes must be positive");
TORCH_CHECK(device_index >= 0, "device index must be non-negative");
const at::cuda::CUDAGuard device_guard(
c10::Device(c10::DeviceType::CUDA, static_cast<c10::DeviceIndex>(device_index)));
auto ctx = std::make_unique<UlyssesContext>();
ctx->nbytes = static_cast<size_t>(nbytes);
ctx->device = static_cast<int>(device_index);
ctx->rank = static_cast<int>(rank);
ctx->world = static_cast<int>(world_size);
// The window must come from NCCL's allocator (4096B aligned per
// NCCL_WIN_REQUIRED_ALIGNMENT), which is why this is not a torch tensor.
NCCL_TRY(ncclMemAlloc(&ctx->buf, ctx->nbytes), "ncclMemAlloc");
return reinterpret_cast<int64_t>(ctx.release());
}
// Register the user window. Collective: every rank must call together.
void register_ulysses_a2a_window(int64_t handle, int64_t comm_ptr) {
auto* ctx = reinterpret_cast<UlyssesContext*>(handle);
TORCH_CHECK(ctx != nullptr, "handle must come from allocate_ulysses_a2a");
TORCH_CHECK(ctx->buf != nullptr && ctx->win == nullptr && !ctx->dev_comm_created,
"ulysses a2a context is not in the allocated state");
const at::cuda::CUDAGuard device_guard(
c10::Device(c10::DeviceType::CUDA, static_cast<c10::DeviceIndex>(ctx->device)));
ctx->comm = reinterpret_cast<ncclComm_t>(comm_ptr);
NCCL_TRY(ncclCommWindowRegister(ctx->comm, ctx->buf, ctx->nbytes, &ctx->win,
NCCL_WIN_COLL_SYMMETRIC),
"ncclCommWindowRegister");
}
// Create the device communicator only after Python has voted that every rank
// registered its window. This operation is collective as well.
void create_ulysses_a2a_dev_comm(int64_t handle) {
auto* ctx = reinterpret_cast<UlyssesContext*>(handle);
TORCH_CHECK(ctx != nullptr, "handle must come from allocate_ulysses_a2a");
TORCH_CHECK(ctx->comm != nullptr && ctx->win != nullptr && !ctx->dev_comm_created,
"ulysses a2a context is not in the window-registered state");
const at::cuda::CUDAGuard device_guard(
c10::Device(c10::DeviceType::CUDA, static_cast<c10::DeviceIndex>(ctx->device)));
ncclDevCommRequirements reqs = NCCL_DEV_COMM_REQUIREMENTS_INITIALIZER;
reqs.lsaBarrierCount = fi::kMaxBlocks;
NCCL_TRY(ncclDevCommCreate(ctx->comm, &reqs, &ctx->devComm), "ncclDevCommCreate");
ctx->dev_comm_created = true;
}
static ncclResult_t first_error(ncclResult_t current, ncclResult_t next) {
return current == ncclSuccess ? next : current;
}
void dispose_ulysses_a2a(int64_t handle) {
auto* ctx = reinterpret_cast<UlyssesContext*>(handle);
if (ctx == nullptr) return;
const at::cuda::CUDAGuard device_guard(
c10::Device(c10::DeviceType::CUDA, static_cast<c10::DeviceIndex>(ctx->device)));
ncclResult_t result = ncclSuccess;
if (ctx->comm != nullptr && ctx->dev_comm_created) {
result = first_error(result, ncclDevCommDestroy(ctx->comm, &ctx->devComm));
ctx->dev_comm_created = false;
}
if (ctx->comm != nullptr && ctx->win != nullptr) {
result = first_error(result, ncclCommWindowDeregister(ctx->comm, ctx->win));
ctx->win = nullptr;
}
if (ctx->buf != nullptr) {
result = first_error(result, ncclMemFree(ctx->buf));
ctx->buf = nullptr;
}
delete ctx;
TORCH_CHECK(result == ncclSuccess, "ulysses a2a cleanup failed: rc=", static_cast<int>(result));
}
// Whether the whole group is load-store accessible. NCCL determined this at
// ncclCommInitRank.
bool ulysses_lsa_covers_group(int64_t comm_ptr, int64_t world_size) {
auto comm = reinterpret_cast<ncclComm_t>(comm_ptr);
ncclCommProperties properties = NCCL_COMM_PROPERTIES_INITIALIZER;
NCCL_TRY(ncclCommQueryProperties(comm, &properties), "ncclCommQueryProperties");
ncclTeam_t lsa = ncclTeamLsa(comm);
return properties.deviceApiSupport && lsa.nRanks == static_cast<int>(world_size);
}
// Fused-transpose Ulysses all-to-all.
// mode == 0: inp [B, S_local, H, D] -> out [B, S_global, H_local, D]
// mode == 1: inp [B, S_global, H_local, D] -> out [B, S_local, H, D]
// where H is the *global* head count and H_local = H / world_size.
void ulysses_a2a(int64_t handle, torch::Tensor inp, torch::Tensor out, int64_t B, int64_t S_local,
int64_t H, int64_t D, int64_t mode) {
auto* ctx = reinterpret_cast<UlyssesContext*>(handle);
TORCH_CHECK(ctx != nullptr, "handle must come from allocate_ulysses_a2a");
const at::cuda::CUDAGuard device_guard(inp.device());
auto stream = at::cuda::getCurrentCUDAStream();
TORCH_CHECK(inp.is_cuda() && out.is_cuda(), "inp and out must be CUDA tensors");
TORCH_CHECK(inp.is_contiguous() && out.is_contiguous(), "inp and out must be contiguous");
TORCH_CHECK(inp.device() == out.device(), "inp and out must be on the same device");
TORCH_CHECK(inp.scalar_type() == out.scalar_type(), "inp and out must share a dtype");
TORCH_CHECK(inp.numel() == out.numel(), "inp and out must have equal element counts");
TORCH_CHECK(inp.get_device() == ctx->device, "input is on CUDA device ", inp.get_device(),
" but the Ulysses context belongs to device ", ctx->device);
TORCH_CHECK(mode == 0 || mode == 1, "mode must be 0 or 1");
TORCH_CHECK(inp.dim() == 4 && out.dim() == 4, "inp and out must be 4-D");
const int W = ctx->world;
TORCH_CHECK(H % W == 0, "global head count must be divisible by world size");
const int H_local = static_cast<int>(H / W);
const torch::Tensor& local_op = (mode == 0) ? inp : out; // [B, S_local, H, D]
const torch::Tensor& global_op = (mode == 0) ? out : inp; // [B, S_global, H_local, D]
TORCH_CHECK(local_op.size(0) == B && local_op.size(1) == S_local && local_op.size(2) == H &&
local_op.size(3) == D,
"the [B, S_local, H, D] operand of mode ", mode, " has shape (", local_op.size(0),
", ", local_op.size(1), ", ", local_op.size(2), ", ", local_op.size(3),
"), expected (", B, ", ", S_local, ", ", H, ", ", D, ")");
TORCH_CHECK(global_op.size(0) == B && global_op.size(1) == W * S_local &&
global_op.size(2) == H_local && global_op.size(3) == D,
"the [B, S_global, H_local, D] operand of mode ", mode, " has shape (",
global_op.size(0), ", ", global_op.size(1), ", ", global_op.size(2), ", ",
global_op.size(3), "), expected (", B, ", ", W * S_local, ", ", H_local, ", ", D,
")");
const size_t out_bytes = out.numel() * out.element_size();
TORCH_CHECK(out_bytes <= ctx->nbytes, "operand of ", out_bytes,
" bytes exceeds the window capacity ", ctx->nbytes);
const int64_t num_rows = B * static_cast<int64_t>(W) * S_local;
const int blocks =
static_cast<int>(std::max<int64_t>(1, std::min<int64_t>(fi::kMaxBlocks, num_rows)));
const int threads = fi::kUlyssesThreads;
#define LAUNCH_ULYSSES_A2A(T, NG, MODE) \
fi::ulysses_a2a_kernel<T, NG, MODE><<<blocks, threads, 0, stream>>>( \
reinterpret_cast<const T*>(inp.data_ptr()), ctx->devComm, ctx->win, /*off=*/0, \
ctx->rank, static_cast<int>(B), static_cast<int>(S_local), H_local, static_cast<int>(D))
#define DISPATCH_NGPUS(T, MODE) \
switch (W) { \
case 2: \
LAUNCH_ULYSSES_A2A(T, 2, MODE); \
break; \
case 4: \
LAUNCH_ULYSSES_A2A(T, 4, MODE); \
break; \
case 6: \
LAUNCH_ULYSSES_A2A(T, 6, MODE); \
break; \
case 8: \
LAUNCH_ULYSSES_A2A(T, 8, MODE); \
break; \
default: \
TORCH_CHECK(false, "ulysses_a2a only supports world size in (2,4,6,8)"); \
}
#define DISPATCH_DTYPE(MODE) \
switch (out.scalar_type()) { \
case at::ScalarType::Float: { \
DISPATCH_NGPUS(float, MODE); \
break; \
} \
case at::ScalarType::Half: { \
DISPATCH_NGPUS(half, MODE); \
break; \
} \
case at::ScalarType::BFloat16: { \
DISPATCH_NGPUS(nv_bfloat16, MODE); \
break; \
} \
default: \
TORCH_CHECK(false, "ulysses_a2a only supports float32, float16 and bfloat16, got ", \
out.scalar_type()); \
}
if (mode == 0) {
DISPATCH_DTYPE(0);
} else {
DISPATCH_DTYPE(1);
}
#undef DISPATCH_DTYPE
#undef DISPATCH_NGPUS
#undef LAUNCH_ULYSSES_A2A
TORCH_CHECK(cudaGetLastError() == cudaSuccess, "ulysses_a2a kernel launch failed");
// Copy this rank's completed result out of the window.
auto status = cudaMemcpyAsync(out.data_ptr(), ctx->buf, out_bytes, cudaMemcpyDeviceToDevice,
stream);
TORCH_CHECK(status == cudaSuccess, "ulysses_a2a copy-out failed: ", cudaGetErrorString(status));
}
void register_ulysses_a2a(pybind11::module_& m) {
m.def("allocate_ulysses_a2a", &allocate_ulysses_a2a, "allocate a local ulysses a2a window");
m.def("register_ulysses_a2a_window", &register_ulysses_a2a_window,
"register the ulysses a2a window collectively");
m.def("create_ulysses_a2a_dev_comm", &create_ulysses_a2a_dev_comm,
"create the ulysses a2a device communicator collectively");
m.def("dispose_ulysses_a2a", &dispose_ulysses_a2a, "release a ulysses a2a context");
m.def("ulysses_lsa_covers_group", &ulysses_lsa_covers_group,
"whether the whole group is load-store accessible");
m.def("ulysses_a2a", &ulysses_a2a, "fused-transpose Ulysses all-to-all over NVLink");
}
@@ -22,6 +22,11 @@ extern std::vector<torch::Tensor> block_sparse_attention_backward(
);
#endif
#ifdef FASTVIDEO_KERNEL_COMPILE_ULYSSES_A2A
// Ulysses sequence-parallel all-to-all (csrc/comm/)
void register_ulysses_a2a(pybind11::module_ &);
#endif
// TurboDiffusion kernels
void register_quant(pybind11::module_ &);
void register_rms_norm(pybind11::module_ &);
@@ -61,6 +66,11 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward (Hopper)");
#endif
#ifdef FASTVIDEO_KERNEL_COMPILE_ULYSSES_A2A
// Ulysses sequence-parallel all-to-all
register_ulysses_a2a(m);
#endif
// TurboDiffusion
register_quant(m);
register_rms_norm(m);
@@ -0,0 +1,165 @@
/*
* Copyright (c) 2025 by FlashInfer team.
*
* Adapted from flashinfer-ai/flashinfer @ 8a94642d83cba0939035868fb6c309b4474a13d6
* (PR #3820), which in turn adapted ThunderKittens' NVLink all-to-all:
* https://github.com/HazyResearch/ThunderKittens/blob/main/kernels/parallel/all_to_all/all_to_all.cu
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
// Fused-transpose Ulysses all-to-all over NVLink. Peer addresses come from
// ncclGetLsaPointer and synchronization from ncclLsaBarrierSession; the index
// math and slab decomposition below are upstream's.
//
// head_dim == 2 layout, uniform sequence splits. With
// W = ulysses world size
// H_local = H / W
// S_global = S_local * W
//
// mode == 0 (input a2a): [B, S_local, H, D] -> [B, S_global, H_local, D]
// y_r[b, j*S_local + s, hl, d] = x_j[b, s, r*H_local + hl, d]
// mode == 1 (output a2a): [B, S_global, H_local, D] -> [B, S_local, H, D]
// out_j[b, s, r*H_local + hl, d] = u_r[b, j*S_local + s, hl, d]
//
// In both modes the unit of transfer is a contiguous (H_local * D) block, so
// every cross-GPU store is fully coalesced.
#ifndef FASTVIDEO_COMM_ULYSSES_ALL_TO_ALL_CUH_
#define FASTVIDEO_COMM_ULYSSES_ALL_TO_ALL_CUH_
#include <cstdint>
#include <nccl.h>
#include <nccl_device.h>
namespace fastvideo {
namespace comm {
namespace ulysses {
constexpr int kUlyssesThreads = 512;
// Deliberately modest: this is link-bandwidth bound, so a small grid leaves the
// rest of the GPU free without costing throughput.
constexpr int kMaxBlocks = 36;
// Shared movement body for the fused-transpose all-to-all (no barriers).
//
// Rows are ordered ((b * W + peer) * S_local + s), so consecutive rows share a
// (batch, peer) and are contiguous on the gather side of the transpose. Each
// block takes a contiguous slab of rows and flattens its threads over the 16B
// units in it, so consecutive lanes address one peer buffer back to back and the
// remote writes coalesce into large bursts rather than (H_local * D)-sized
// scattered ones.
template <typename T, int NGPUS, int MODE>
__device__ __forceinline__ void ulysses_a2a_move(const T* __restrict__ local_in,
void* const* peer_ptrs, int rank, int B,
int S_local, int H_local, int D) {
static_assert(MODE == 0 || MODE == 1, "MODE must be 0 or 1");
const int W = NGPUS;
const int64_t H = static_cast<int64_t>(H_local) * W;
const int64_t S_global = static_cast<int64_t>(S_local) * W;
const int64_t block_len = static_cast<int64_t>(H_local) * D; // elements/row
const int64_t num_rows = static_cast<int64_t>(B) * W * S_local;
// 16B-vectorized fast path when every row is 16B aligned (the common case:
// contiguous bf16/fp16/fp32 tensors with block_len * sizeof(T) % 16 == 0).
using Vec = int4;
constexpr int kVecBytes = sizeof(Vec);
const int64_t row_bytes = block_len * static_cast<int64_t>(sizeof(T));
const bool vec_ok =
(row_bytes % kVecBytes) == 0 && (reinterpret_cast<uintptr_t>(local_in) % kVecBytes) == 0;
// Contiguous slab of rows for this block.
const int64_t rows_per_block = (num_rows + gridDim.x - 1) / gridDim.x;
const int64_t row_lo = static_cast<int64_t>(blockIdx.x) * rows_per_block;
int64_t row_hi = row_lo + rows_per_block;
if (row_hi > num_rows) row_hi = num_rows;
if (row_lo >= row_hi) return;
const int tid = threadIdx.x;
const int nthr = blockDim.x;
// Decode (b, peer, s) and compute src/dst element offsets for a given row.
auto offsets = [&](int64_t row, int64_t& src_off, int64_t& dst_off) {
const int64_t s = row % S_local;
const int64_t tmp = row / S_local;
const int64_t peer = tmp % W;
const int64_t b = tmp / W;
if constexpr (MODE == 0) {
src_off = ((b * S_local + s) * H + peer * H_local) * D;
dst_off = (b * S_global + static_cast<int64_t>(rank) * S_local + s) * block_len;
} else {
src_off = (b * S_global + peer * S_local + s) * block_len;
dst_off = ((b * S_local + s) * H + static_cast<int64_t>(rank) * H_local) * D;
}
return peer;
};
if (vec_ok) {
const int64_t units_per_row = row_bytes / kVecBytes;
const int64_t total_units = (row_hi - row_lo) * units_per_row;
for (int64_t u = tid; u < total_units; u += nthr) {
const int64_t local_row = u / units_per_row;
const int64_t unit = u - local_row * units_per_row;
const int64_t row = row_lo + local_row;
int64_t src_off, dst_off;
const int64_t peer = offsets(row, src_off, dst_off);
const Vec* s4 = reinterpret_cast<const Vec*>(local_in + src_off);
Vec* d4 = reinterpret_cast<Vec*>((T*)peer_ptrs[peer] + dst_off);
d4[unit] = s4[unit];
}
} else {
// Scalar fallback (unaligned / odd shapes).
for (int64_t row = row_lo; row < row_hi; ++row) {
int64_t src_off, dst_off;
const int64_t peer = offsets(row, src_off, dst_off);
const T* s_ptr = local_in + src_off;
T* d_ptr = (T*)peer_ptrs[peer] + dst_off;
for (int64_t i = tid; i < block_len; i += nthr) {
d_ptr[i] = s_ptr[i];
}
}
}
}
// The transfer mode is a compile-time template parameter so the address math
// specializes and the coalesced slab decomposition is fully unrolled per mode.
template <typename T, int NGPUS, int MODE>
__global__ void __launch_bounds__(kUlyssesThreads, 1)
ulysses_a2a_kernel(const T* __restrict__ local_in, ncclDevComm devComm, ncclWindow_t win,
size_t win_offset, int rank, int B, int S_local, int H_local, int D) {
// Resolved once; the movement loop would otherwise call this per 16B store.
void* peer_ptrs[NGPUS];
#pragma unroll
for (int p = 0; p < NGPUS; ++p) {
peer_ptrs[p] = ncclGetLsaPointer(win, win_offset, p);
}
// Each CTA owns one barrier generation. Sharing index 0 across independently
// scheduled CTAs races the generation counter and is unsupported by NCCL's
// device API. The host reserves kMaxBlocks slots when it creates devComm.
ncclLsaBarrierSession<ncclCoopCta> bar(
ncclCoopCta(), devComm, ncclTeamTagLsa(), /*index=*/blockIdx.x);
// Every rank must have entered before anyone writes into peer buffers.
bar.sync(ncclCoopCta(), cuda::memory_order_relaxed);
ulysses_a2a_move<T, NGPUS, MODE>(local_in, peer_ptrs, rank, B, S_local, H_local, D);
// Release-acquire: all peer writes visible before a rank reads its window.
bar.sync(ncclCoopCta(), cuda::memory_order_acq_rel);
}
} // namespace ulysses
} // namespace comm
} // namespace fastvideo
#endif // FASTVIDEO_COMM_ULYSSES_ALL_TO_ALL_CUH_
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.3.2"
version = "0.3.4"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -49,6 +49,36 @@ def _force_tk() -> bool:
return os.environ.get("FASTVIDEO_VSA_TK", "0") == "1"
def _force_sm100a() -> bool:
"""True iff the sm_100a (Blackwell) forward is explicitly opted into.
Opt-in only (same env the H3 backend honors): the sm_100a extension is
forward-only, so this routing pairs it with the Triton backward -- its lse
is already in Triton's M format. Honored only when
``block_sparse_attn_sm100a.is_supported`` passes. Unsupported 64-token
metadata falls through to the default selection; unsupported 128-token
metadata raises because Triton has no compatible fallback.
``FASTVIDEO_VSA_TRITON`` still wins.
"""
return os.environ.get("FASTVIDEO_VSA_SM100A", "0") == "1"
def _sm100a_is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
try:
from fastvideo_kernel import block_sparse_attn_sm100a as vsa_sm100a
except Exception:
return False
return vsa_sm100a.is_supported(q, variable_block_sizes)
def _infer_block_size(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> int:
num_blocks = variable_block_sizes.numel()
seq_len = q.shape[2]
if num_blocks == 0 or seq_len % num_blocks != 0:
return 0
return seq_len // num_blocks
# ---------------------------------------------------------------------------
# Index helpers
# ---------------------------------------------------------------------------
@@ -339,6 +369,76 @@ def _backward_sm90(ctx, grad_o, grad_lse):
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
# ---------------------------------------------------------------------------
# SM100A backend custom op (index-native)
#
# Forward runs the sm_100a CUDA extension; backward reuses the Triton kernels.
# The sm_100a forward emits lse in exactly Triton's M format (max*log2e +
# log2(l)), so the pairing needs no conversion. The Triton backward is
# hardcoded to 64-token blocks, hence the block-size assert below.
# ---------------------------------------------------------------------------
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_sm100a",
mutates_args=(),
device_types="cuda",
)
def block_sparse_attn_sm100a_op(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
from fastvideo_kernel.block_sparse_attn_sm100a import block_sparse_attn_sm100a
o, M = block_sparse_attn_sm100a(q, k, v, q2k_idx, q2k_num, variable_block_sizes,
need_lse=True)
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_sm100a")
def _block_sparse_attn_sm100a_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q)
M = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
device=q.device,
dtype=torch.float32,
)
return o, M
def _setup_context_sm100a(ctx, inputs, output):
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
def _backward_sm100a(ctx, grad_o, grad_M):
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
block = q.shape[2] // variable_block_sizes.numel()
if block != 64:
raise RuntimeError(
"block_sparse_attn_sm100a backward pairs the sm_100a forward with the "
f"Triton backward, which is hardcoded to 64-token blocks; got {block}. "
"Run 128-token-block metadata without grad, or use the Triton forward.")
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, q2k_idx,
q2k_num, variable_block_sizes)
return dq, dk, dv, None, None, None
block_sparse_attn_sm100a_op.register_autograd(_backward_sm100a,
setup_context=_setup_context_sm100a)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
@@ -364,11 +464,24 @@ def block_sparse_attn_from_indices(
# Backend resolution:
# - FASTVIDEO_VSA_TRITON forces Triton everywhere.
# - FASTVIDEO_VSA_SM100A opts into the sm_100a forward (Triton backward).
# Unsupported 64-token metadata falls through; unsupported 128-token
# metadata raises because Triton cannot consume it.
# - FASTVIDEO_VSA_TK requests sm_90 TK; honored only when it's actually
# available (else falls through to the default below).
# - Otherwise: TK on sm_90 if available, else Triton.
if _force_triton():
use_sm90 = False
elif _force_sm100a():
if _sm100a_is_supported(q, variable_block_sizes):
return block_sparse_attn_sm100a_op(q, k, v, q2k_idx, q2k_num,
variable_block_sizes)
if _infer_block_size(q, variable_block_sizes) == 128:
raise NotImplementedError(
"128-token block-sparse attention requires the sm_100a forward; "
"the Triton fallback only supports 64-token blocks, and the "
"sm_100a route is unavailable for this input.")
use_sm90 = sm90_available
elif _force_tk():
use_sm90 = sm90_available
else:
@@ -84,6 +84,107 @@ def is_supported(q: torch.Tensor, variable_block_sizes: torch.Tensor) -> bool:
return True
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_sm100a_inference",
mutates_args=(),
device_types="cuda",
)
def _block_sparse_attn_sm100a_inference(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
"""Opaque no-LSE launch used by the inference-only sm_100a route.
The extension is exposed as a raw pybind function rather than a dispatcher
op. Calling it directly makes Dynamo descend through a Python/C++ boundary
that has no fake implementation, so ``torch.compile(fullgraph=True)``
cannot capture a sparse H3 block. Keep that boundary inside this custom op;
its inputs have already been normalized by the public wrapper below.
"""
fwd = _FWD_BY_BLOCK[_block_size(q, variable_block_sizes)]
sm_scale = 1.0 / (q.shape[-1]**0.5)
res = fwd(q, k, v, None, q2k_idx, q2k_num, variable_block_sizes, sm_scale, False)
return res[0]
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_sm100a_inference")
def _block_sparse_attn_sm100a_inference_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
# The C++ binding allocates its output with torch::empty_like(q).
return torch.empty_like(q)
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_sm100a_from_mask_inference",
mutates_args=(),
device_types="cuda",
)
def _block_sparse_attn_sm100a_from_mask_inference(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
"""Opaque mask compaction plus no-LSE sm_100a launch for inference.
H3 naturally produces a bool block map. Its Triton ``map_to_index`` call
must live behind the same opaque boundary as the raw pybind launch;
otherwise Dynamo sees that kernel before reaching the index-native custom
op and full-graph capture still fails.
"""
from fastvideo_kernel.triton_kernels.index import map_to_index
q2k_idx, q2k_num = map_to_index(block_map)
return _block_sparse_attn_sm100a_inference(
q,
k,
v,
q2k_idx.to(torch.int32).contiguous(),
q2k_num.to(torch.int32).contiguous(),
variable_block_sizes,
)
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_sm100a_from_mask_inference")
def _block_sparse_attn_sm100a_from_mask_inference_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
return torch.empty_like(q)
def block_sparse_attn_sm100a_from_mask(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, None]:
"""Inference forward from a bool block map, with compaction kept opaque."""
out = _block_sparse_attn_sm100a_from_mask_inference(
q.contiguous(),
k.contiguous(),
v.contiguous(),
block_map.to(torch.bool).contiguous(),
variable_block_sizes.to(torch.int32).contiguous(),
)
return out, None
def block_sparse_attn_sm100a(
q: torch.Tensor,
k: torch.Tensor,
@@ -92,13 +193,24 @@ def block_sparse_attn_sm100a(
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
need_lse: bool = True,
) -> Tuple[torch.Tensor, torch.Tensor]:
) -> Tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass. Returns ``(out, lse)``; ``out`` has q's layout."""
fwd = _FWD_BY_BLOCK[_block_size(q, variable_block_sizes)]
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
idx = q2k_idx.to(torch.int32).contiguous()
num = q2k_num.to(torch.int32).contiguous()
vbs = variable_block_sizes.to(torch.int32).contiguous()
if not need_lse:
# This is the production inference path. The custom op keeps the raw
# pybind launch opaque to Dynamo while its fake kernel carries output
# metadata through full-graph capture.
return _block_sparse_attn_sm100a_inference(q, k, v, idx, num, vbs), None
# Preserve the established LSE-producing path for correctness tests and
# any future forward/backward pairing; only inference needs the opaque op.
fwd = _FWD_BY_BLOCK[_block_size(q, vbs)]
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)
res = fwd(q, k, v, None, idx, num, vbs, sm_scale, True)
return res[0], res[1]
@@ -0,0 +1,86 @@
# SPDX-License-Identifier: Apache-2.0
"""Ulysses sequence-parallel all-to-all ops.
Thin wrappers over csrc/comm/ulysses_all_to_all.cu. The kernel stores directly
into peers' memory through NCCL's device API, so the caller supplies an
ncclComm_t for the group.
"""
import torch
try:
from fastvideo_kernel._C import fastvideo_kernel_ops as _ops
except ImportError: # pragma: no cover - no compiled extension in this install
_ops = None
_SUPPORTED_WORLD_SIZES = (2, 4, 6, 8)
_REQUIRED_OPS = (
"allocate_ulysses_a2a",
"register_ulysses_a2a_window",
"create_ulysses_a2a_dev_comm",
"dispose_ulysses_a2a",
"ulysses_lsa_covers_group",
"ulysses_a2a",
)
def is_available() -> bool:
"""Whether this wheel was built with the Ulysses all-to-all kernel."""
return _ops is not None and all(hasattr(_ops, name) for name in _REQUIRED_OPS)
def _require() -> None:
if not is_available():
raise RuntimeError(
"the Ulysses all-to-all kernel is not present in this fastvideo-kernel build; "
"rebuild with ./build.sh or install a wheel that includes csrc/comm/")
def lsa_covers_group(comm_ptr: int, world_size: int) -> bool:
"""Whether every rank in the group is load-store accessible to every other."""
_require()
return bool(_ops.ulysses_lsa_covers_group(int(comm_ptr), int(world_size)))
def allocate(nbytes: int, rank: int, world_size: int, device_index: int) -> int:
"""Allocate one rank's local symmetric window without a collective."""
_require()
if world_size not in _SUPPORTED_WORLD_SIZES:
raise ValueError(f"ulysses a2a supports world sizes {_SUPPORTED_WORLD_SIZES}, "
f"got {world_size}")
return int(_ops.allocate_ulysses_a2a(int(nbytes), int(rank), int(world_size), int(device_index)))
def register_window(handle: int, comm_ptr: int) -> None:
"""Register an allocated window with the supplied communicator.
Every rank in ``comm_ptr`` must call this together.
"""
_require()
_ops.register_ulysses_a2a_window(int(handle), int(comm_ptr))
def create_dev_comm(handle: int) -> None:
"""Create the device communicator for a registered window collectively."""
_require()
_ops.create_ulysses_a2a_dev_comm(int(handle))
def dispose(handle: int) -> None:
"""Release a handle from :func:`allocate`. It is dangling afterwards."""
_require()
_ops.dispose_ulysses_a2a(int(handle))
def all_to_all(handle: int, inp: torch.Tensor, out: torch.Tensor, B: int, S_local: int, H: int,
D: int, mode: int) -> None:
"""Run one fused all-to-all on the current stream, writing into ``out``.
``mode == 0``: ``[B, S_local, H, D] -> [B, S_global, H_local, D]``
``mode == 1``: ``[B, S_global, H_local, D] -> [B, S_local, H, D]``
``H`` is the global head count. Every rank must call with consistent
geometry in the same order.
"""
_require()
_ops.ulysses_a2a(int(handle), inp, out, int(B), int(S_local), int(H), int(D), int(mode))
@@ -16,14 +16,21 @@ import triton.language as tl
import math # small utility needed by the sparse wrapper
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
# the code below and commenting out the equivalent parameters is convenient for
# re-tuning.
# BLOCK_M / BLOCK_N are fixed at 64 because they are structural, not tunable:
# the kernel indexes the top-k list per BLOCK_M q-tile and addresses keys as
# kv_idx * BLOCK_N, so both must match the granularity q2k_index and
# variable_block_sizes were built at.
#
# num_stages / num_warps ARE free, and the previous {3, 4, 7} was inherited from
# the upstream tutorial rather than tuned here. It skips 5 and 6; on Blackwell
# (sm_121) the optimum is num_stages=5, so the search could not reach it. Both
# block paths independently select 5 once it is available. Autotune still picks
# per architecture, so other GPUs re-tune rather than inheriting this choice.
configs = [
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
for BM in [64]\
for BN in [64]\
for s in [3, 4, 7]\
for s in [2, 3, 4, 5, 6, 7]\
for w in [4, 8]\
]
@@ -1 +1 @@
__version__ = "0.3.2"
__version__ = "0.3.4"
+67 -2
View File
@@ -10,6 +10,14 @@ from fastvideo.attention.utils.flash_attn_default import (
flash_attn_func_compilable,
)
if fa_version == "4":
# The FA4 varlen wrapper is already a compile-safe custom op. Keep the
# import conditional so FA2/FA3 environments do not need flash_attn.cute.
from fastvideo.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_compilable, )
else:
flash_attn_varlen_func_compilable = None
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
@@ -19,7 +27,12 @@ from fastvideo.attention.backends.abstract import (
from fastvideo.logger import init_logger
logger = init_logger(__name__)
logger.info("Using FlashAttention-%s backend", fa_version)
# Every worker records the loaded FlashAttention implementation so a
# distributed profiling log contains one backend receipt per rank.
logger.info("Worker %s Using FlashAttention-%s backend",
os.environ.get("RANK", "0"),
fa_version,
local_main_process_only=False)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
@@ -33,6 +46,21 @@ except ImportError:
_FA4_FP4_AVAILABLE = False
_FA4_QUANT_OPS: tuple | None = None
_FA4_PACKED_VARLEN_CONFIG_LOGGED = False
def _log_fa4_packed_varlen_config() -> None:
"""Emit one configuration receipt per worker process.
This intentionally does not claim that a particular forward used the
kernel: masks, gradients, batch size, unequal Q/K lengths, and NVFP4 are
runtime guards evaluated later in ``_forward_impl``.
"""
global _FA4_PACKED_VARLEN_CONFIG_LOGGED
if not _FA4_PACKED_VARLEN_CONFIG_LOGGED:
logger.info("MiniMax-H3 dense attention: FA4 packed-varlen route configured (runtime guards apply)",
local_main_process_only=False)
_FA4_PACKED_VARLEN_CONFIG_LOGGED = True
def _import_fa4_quant_ops() -> tuple:
@@ -223,7 +251,21 @@ class FlashAttentionImpl(AttentionImpl):
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.nvfp4_fa4 = extra_impl_args.get("nvfp4_fa4", False) or os.environ.get("FASTVIDEO_NVFP4_FA4", "0") == "1"
# MiniMax-H3's dense DiT explicitly enables this faster FA4 entry
# point. It remains off for every other model and for the H3 text
# refiner; grad-enabled calls stay on the established fixed path.
self.fa4_packed_varlen = bool(extra_impl_args.get("fa4_packed_varlen", False))
if self.fa4_packed_varlen and fa_version == "4":
_log_fa4_packed_varlen_config()
# An explicit ``nvfp4_fa4`` impl arg wins over the process-wide
# FASTVIDEO_NVFP4_FA4 env opt-in, so precision-sensitive layers (e.g.
# the FP32-pinned H3 VAE attention) can force-disable FP4 Q/K
# quantization while the DiT keeps it. When the arg is absent the env
# keeps its previous semantics.
nvfp4_fa4 = extra_impl_args.get("nvfp4_fa4")
if nvfp4_fa4 is None:
nvfp4_fa4 = os.environ.get("FASTVIDEO_NVFP4_FA4", "0") == "1"
self.nvfp4_fa4 = bool(nvfp4_fa4)
if self.nvfp4_fa4:
cap = torch.cuda.get_device_capability()
assert cap in [(10, 0), (10, 3)], (f"NVFP4 FA4 requires Blackwell (sm100a/sm103a), got sm{cap[0]}{cap[1]}")
@@ -318,6 +360,29 @@ class FlashAttentionImpl(AttentionImpl):
elif self.nvfp4_fa4:
output = self._forward_nvfp4(query, key, value)
elif (self.fa4_packed_varlen and fa_version == "4" and not torch.is_grad_enabled() and query.shape[0] == 1
and query.shape[1] == key.shape[1] == value.shape[1]):
# FA4's packed-varlen entry point is materially faster for H3's
# long, single-document self-attention. Flatten only the batch
# dimension and describe that one sequence with CUDA int32
# cumulative lengths; the existing custom-op wrapper keeps this
# route traceable under torch.compile(fullgraph=True).
assert flash_attn_varlen_func_compilable is not None
sequence_length = query.shape[1]
cu_seqlens = torch.arange(2, dtype=torch.int32, device=query.device) * sequence_length
output = flash_attn_varlen_func_compilable(
query.squeeze(0),
key.squeeze(0),
value.squeeze(0),
cu_seqlens,
cu_seqlens,
sequence_length,
sequence_length,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
causal=self.causal,
).unsqueeze(0)
else:
# Route through the compilable wrapper so dynamo sees a
# registered op (no graph break) for FA2/FA3; identical
@@ -5,13 +5,17 @@ H3 runs one joint bidirectional attention over
``[text | condition keyframes | audio | generated video]``, so this
backend differs from the Wan-tuned ``video_sparse_attn``:
- Tiles are ``[segment-pure prefix chunks] + [3D (4,8,8) video tiles]``;
prefix tiles never straddle segment boundaries.
- Tiles are ``[segment-pure prefix chunks] + [3D video tiles]``; prefix
tiles never straddle segment boundaries. The tile size is selectable at
metadata build time: 256 tokens ``(4,8,8)`` (default) or 64 tokens
``(4,4,4)`` (see ``VSA_H3_TILE_SHAPES``).
- Selection is pure Python on pooled tile scores; the block-sparse kernel
consumes an explicit bool mask, so no kernel changes are needed.
- The compression branch is gated by ``to_gate_compress``, which the H3
checkpoint does not carry: the loader zero-initializes it, so untrained
inference is exactly pure sparse and finetuning can learn the gate.
- The compression branch is gated by ``to_gate_compress``, which the base
H3 checkpoint does not carry: the loader zero-initializes it, so
untrained inference is exactly pure sparse and finetuning can learn the
gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship
trained gates, which load and activate the branch.
- Non-video *queries* are always dense. Non-video *keys* are either
always-selected for every query ("exempt", default) or compete in
top-k under a FLOP-matched budget ("compete") — the ablation axis,
@@ -20,22 +24,50 @@ backend differs from the Wan-tuned ``video_sparse_attn``:
(``vsa_dense_first_n_steps``, ``vsa_dense_layers``) let mixed schedules
run the diffuse steps/layers dense while pushing the rest harder.
Targets sm10.x through the FA4 CuTe 256-tile path
At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path
(``FASTVIDEO_VSA_CUTEDSL=1``); the Triton 256→64 expansion is the
fallback and keeps identical mask semantics.
fallback and keeps identical mask semantics. At tile 64 the block map is
already at the kernels' native 64-token granularity, so both forward and
backward run the Triton block-sparse kernels directly (no expansion,
``FASTVIDEO_VSA_CUTEDSL`` does not apply). A third, opt-in route exists
for the tile-64 FORWARD only: ``FASTVIDEO_VSA_SM100A=1`` sends no-grad
forwards through the sm_100a CUDA block-sparse kernel
(``fastvideo_kernel.block_sparse_attn_sm100a``, upstream PR #1719 plus
our per-q-tile ``q2k_num`` fix) when the extension is built, the device
is sm_100, and the geometry qualifies. The CUDA kernel assigns adjacent
pairs of query tiles to CTAs, so an odd logical tile count receives one
internal, zero-valid partner tile for the no-grad call only. Score search,
the trained mask, gate-compress, and the returned packed sequence remain on
the original logical tiles. Grad-tracking forwards and every backward stay
on Triton unchanged. If the env is set but a precondition fails, the route
logs one warning and falls back.
"""
import functools
import math
import os
from dataclasses import dataclass
from typing import Any
import torch
try:
from fastvideo_kernel.block_sparse_attn import block_sparse_attn as block_sparse_attn_64_bhsd
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
from fastvideo_kernel.triton_kernels.index import map_to_index
except ImportError:
block_sparse_attn_64_bhsd = None
block_sparse_attn_256_bshd = None
map_to_index = None
try:
# Optional: only present in fastvideo_kernel builds that carry the sm_100a
# CUDA block-sparse forward (upstream PR #1719). The module itself imports
# fine without the compiled symbols (`_HAS_VSA_SM100A` is then False and
# `is_supported` says no), so this only guards *module* availability.
from fastvideo_kernel import block_sparse_attn_sm100a as _sm100a
except ImportError:
_sm100a = None
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder, layer_idx_from_prefix)
@@ -43,51 +75,174 @@ from fastvideo.attention.backends.video_sparse_attn import (compute_topk, constr
get_non_pad_index, get_tile_partition_indices,
scatter_into_tile_buf)
from fastvideo.attention.backends.video_sparse_attn_h3_probe import probe_enabled, record_probe
from fastvideo.logger import init_logger
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
logger = init_logger(__name__)
# Opt-in switch for the sm_100a CUDA forward on the tile-64 no-grad path.
VSA_SM100A_ENV = "FASTVIDEO_VSA_SM100A"
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x (default)
_TILE_ELEMS = math.prod(VSA_H3_TILE_SIZE)
# Selectable tile geometries, keyed by element count (= the build-time
# ``tile_size``). 64 runs the native 64-token Triton block-sparse kernels for
# forward AND backward — the block map is already at kernel granularity, so no
# 256->64 mask expansion is involved and FASTVIDEO_VSA_CUTEDSL does not apply.
VSA_H3_TILE_SHAPES: dict[int, tuple[int, int, int]] = {
_TILE_ELEMS: VSA_H3_TILE_SIZE,
64: (4, 4, 4),
}
def token_tile_and_valid(variable_block_sizes: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
@torch.library.custom_op(
"fastvideo::h3_vsa_sm100a_from_mask_compat",
mutates_args=(),
device_types="cuda",
)
def _h3_vsa_sm100a_from_mask_compat(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
"""Compile-safe mask adapter for kernel wheels predating the mask API."""
if _sm100a is None or map_to_index is None:
raise RuntimeError("The sm100a compatibility route requires the raw kernel and map_to_index")
q2k_idx, q2k_num = map_to_index(block_map)
out, _ = _sm100a.block_sparse_attn_sm100a(
q,
k,
v,
q2k_idx.to(torch.int32).contiguous(),
q2k_num.to(torch.int32).contiguous(),
variable_block_sizes.to(torch.int32).contiguous(),
need_lse=False,
)
return out
@torch.library.register_fake("fastvideo::h3_vsa_sm100a_from_mask_compat")
def _h3_vsa_sm100a_from_mask_compat_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> torch.Tensor:
del k, v, block_map, variable_block_sizes
return torch.empty_like(q)
def _sm100a_has_compile_safe_mask_route(sm100a_mod: Any) -> bool:
return (callable(getattr(sm100a_mod, "block_sparse_attn_sm100a_from_mask", None))
or (callable(getattr(sm100a_mod, "block_sparse_attn_sm100a", None)) and map_to_index is not None))
def _sm100a_from_mask(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> tuple[torch.Tensor, None]:
"""Use the native mask entry when installed, otherwise the local adapter."""
native = getattr(_sm100a, "block_sparse_attn_sm100a_from_mask", None)
if callable(native):
return native(q, k, v, block_map, variable_block_sizes)
return _h3_vsa_sm100a_from_mask_compat(q, k, v, block_map, variable_block_sizes), None
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
"""Per padded-token tile id and pad-validity mask.
The single encoding of the padding contract, shared by the probe and the
test oracle so they cannot drift from the backend's tile geometry.
``tile_elems`` must match the metadata the sizes came from
(``MiniMaxH3VSAMetadata.tile_elems``).
"""
device = variable_block_sizes.device
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(_TILE_ELEMS)
token_valid = (torch.arange(_TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
return token_tile, token_valid
def _validate_h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
variable_block_sizes: torch.Tensor,
untile_combined_index: torch.Tensor,
tile_elems: int = _TILE_ELEMS,
) -> None:
"""Fail synchronously on out-of-bounds tile geometry.
Invariants the block-sparse kernel trusts without checking:
every tile's valid size is in (0, tile_elems]; the sizes sum to the
packed sequence length; and ``untile_combined_index`` maps each packed
row to exactly one non-pad slot of the padded tile buffer. A violation
would surface only as an async device fault at some later kernel or
collective (e.g. an FSDP all-gather), which is unattributable — so raise
here, once per cached geometry, with the numbers in hand.
"""
total = sum(prefix_segments) + math.prod(dit_seq_shape)
n_pad = variable_block_sizes.numel() * tile_elems
sizes_min = int(variable_block_sizes.min())
sizes_max = int(variable_block_sizes.max())
sizes_sum = int(variable_block_sizes.sum())
if sizes_min < 1 or sizes_max > tile_elems or sizes_sum != total:
raise ValueError(f"VSA-H3 tile sizes out of bounds for prefix={prefix_segments}, video={dit_seq_shape}, "
f"tile_elems={tile_elems}: min={sizes_min}, max={sizes_max}, sum={sizes_sum}, "
f"expected sum={total}.")
if untile_combined_index.numel() != total:
raise ValueError(f"VSA-H3 untile index has {untile_combined_index.numel()} entries for a packed "
f"sequence of {total} rows (prefix={prefix_segments}, video={dit_seq_shape}).")
idx_min = int(untile_combined_index.min())
idx_max = int(untile_combined_index.max())
if idx_min < 0 or idx_max >= n_pad:
# Range first: the pad-slot gather below would itself index out of
# bounds (the very async fault this guard exists to preempt).
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: range "
f"[{idx_min}, {idx_max}] vs padded length {n_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
in_tile_offset = untile_combined_index % tile_elems
maps_into_pad = bool((in_tile_offset >= variable_block_sizes[untile_combined_index // tile_elems]).any())
if maps_into_pad or int(torch.unique(untile_combined_index).numel()) != total:
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: "
f"pad-slot hit={maps_into_pad} "
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
@functools.lru_cache(maxsize=10)
def _h3_tile_geometry(
prefix_segments: tuple[int, ...],
dit_seq_shape: tuple[int, int, int],
device: torch.device,
tile_shape: tuple[int, int, int] = VSA_H3_TILE_SIZE,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
"""Tile the packed sequence: segment-pure prefix chunks, then video tiles.
Returns (tile_partition_indices, variable_block_sizes,
untile_combined_index, num_prefix_tiles, num_video_tiles).
"""
tile_elems = math.prod(tile_shape)
prefix_len = sum(prefix_segments)
prefix_sizes: list[int] = []
for segment in prefix_segments:
full, rem = divmod(segment, _TILE_ELEMS)
prefix_sizes.extend([_TILE_ELEMS] * full)
full, rem = divmod(segment, tile_elems)
prefix_sizes.extend([tile_elems] * full)
if rem:
prefix_sizes.append(rem)
num_prefix_tiles = len(prefix_sizes)
ts_t, ts_h, ts_w = VSA_H3_TILE_SIZE
ts_t, ts_h, ts_w = tile_shape
t, h, w = dit_seq_shape
num_tiles = (math.ceil(t / ts_t), math.ceil(h / ts_h), math.ceil(w / ts_w))
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, VSA_H3_TILE_SIZE)
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
num_video_tiles = int(video_sizes.numel())
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
tile_partition_indices = torch.cat([
torch.arange(prefix_len, device=device, dtype=torch.long),
video_indices,
@@ -100,9 +255,11 @@ def _h3_tile_geometry(
# get_non_pad_index is lru-cached on tensor identity; variable_block_sizes
# is itself cached by this function, so the identity stays stable.
non_pad_index = get_non_pad_index(variable_block_sizes, _TILE_ELEMS)
non_pad_index = get_non_pad_index(variable_block_sizes, tile_elems)
untile_combined_index = non_pad_index[torch.argsort(tile_partition_indices)]
# One-time (lru-cached) synchronous bounds check; see _validate_h3_tile_geometry.
_validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
return (tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles)
@@ -131,6 +288,14 @@ class MiniMaxH3VSABackend(AttentionBackend):
return MiniMaxH3VSAMetadataBuilder
class _MiniMaxH3VSATileBufferHolder:
"""Builder-owned no-grad tile scratch and its active geometry."""
def __init__(self) -> None:
self.buffer: torch.Tensor | None = None
self.untile_geometry: torch.Tensor | None = None
@dataclass
class MiniMaxH3VSAMetadata(AttentionMetadata):
total_seq_length: int
@@ -139,44 +304,54 @@ class MiniMaxH3VSAMetadata(AttentionMetadata):
exempt: bool
variable_block_sizes: torch.Tensor
untile_combined_index: torch.Tensor
# Device-side copy of ``dense_layers``. Regional fullgraph capture uses
# this tensor with each implementation's tensor-valued layer index so the
# shared block code does not specialize once per Python ``layer_idx``.
dense_layers_tensor: 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
# one buffer serves the whole denoising loop (pad slots stay zero and
# every non-pad slot is fully overwritten per tile(), so cross-step reuse
# is valid; saves a ~1.4 GB alloc+memset per step at 720p). VSA-H3 runs
# eager today — revisit the reuse if it ever goes under cudagraphs.
tile_buf_holder: list = None # type: ignore[assignment]
# Builder-owned padded tile buffer. It records the geometry that last
# populated the allocation so a same-shaped geometry change can clear
# stale pad rows once while steady-state denoising reuses the buffer.
tile_buf_holder: _MiniMaxH3VSATileBufferHolder | None = None
class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
self._tile_buf_holder: list = [None]
self._tile_buf_holder = _MiniMaxH3VSATileBufferHolder()
def prepare(self) -> None:
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
prefix_segments: tuple[int, ...],
device: torch.device,
exempt: bool = True,
dense_layers: tuple[int, ...] = (),
**kwargs: dict[str, Any],
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
prefix_segments: tuple[int, ...],
device: torch.device,
exempt: bool = True,
dense_layers: tuple[int, ...] = (),
tile_size: int = _TILE_ELEMS,
**kwargs: dict[str, Any],
) -> MiniMaxH3VSAMetadata:
tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
if tile_shape is None:
raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
raw_latent_shape[2] // patch_size[2])
prefix_segments = tuple(int(s) for s in prefix_segments if s > 0)
total_seq_length = sum(prefix_segments) + math.prod(dit_seq_shape)
(_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles,
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
dense_layers = tuple(int(layer) for layer in dense_layers)
return MiniMaxH3VSAMetadata(
current_timestep=current_timestep,
VSA_sparsity=VSA_sparsity,
@@ -186,13 +361,15 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
exempt=exempt,
variable_block_sizes=variable_block_sizes,
untile_combined_index=untile_combined_index,
dense_layers=tuple(int(layer) for layer in dense_layers),
tile_elems=int(tile_size),
dense_layers=dense_layers,
dense_layers_tensor=torch.tensor(dense_layers, device=device, dtype=torch.int64),
tile_buf_holder=self._tile_buf_holder,
)
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems: int = _TILE_ELEMS) -> torch.Tensor:
"""fp32 mean over each tile_elems-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
written), so a plain sum with fp32 accumulation needs no validity mask
@@ -200,8 +377,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Te
the masked mean exactly.
"""
batch, seq_len, heads, dim = x.shape
n_tiles = seq_len // _TILE_ELEMS
pooled = x.view(batch, n_tiles, _TILE_ELEMS, heads, dim).sum(dim=2, dtype=torch.float32)
n_tiles = seq_len // tile_elems
pooled = x.view(batch, n_tiles, tile_elems, heads, dim).sum(dim=2, dtype=torch.float32)
pooled = pooled / variable_block_sizes.view(1, -1, 1, 1)
return pooled.permute(0, 2, 1, 3)
@@ -232,6 +409,24 @@ def _build_block_mask(
return mask
def _sm100a_unavailable_reason(sm100a_mod: Any, query_bhsd: torch.Tensor, variable_block_sizes: torch.Tensor,
grad_mode: bool) -> str | None:
"""Why the opt-in sm_100a forward route cannot run here, or None if it can.
Pure decision logic, split out so the routing is unit-testable without a
GPU or the compiled extension (tests substitute ``sm100a_mod``). Order
matters only for the message: the cheapest, most actionable reason first.
"""
if sm100a_mod is None:
return "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
if grad_mode:
return "inputs require grad and the sm_100a kernel is forward-only; grad paths keep Triton"
if not sm100a_mod.is_supported(query_bhsd, variable_block_sizes):
return ("block_sparse_attn_sm100a.is_supported returned False (needs an sm_100 device, a built "
"extension, bf16, head_dim 128, an even tile count, and integer tile sizes)")
return None
class MiniMaxH3VSAImpl(AttentionImpl):
def __init__(
@@ -246,26 +441,109 @@ class MiniMaxH3VSAImpl(AttentionImpl):
) -> None:
self.prefix = prefix
self.layer_idx = layer_idx_from_prefix(prefix, default=-1)
self.head_size = head_size
# None means the regional-compile preparation hook has not run. The
# eager path deliberately ignores this cache and preserves its
# request-time env/probe/fallback behavior; only Dynamo capture reads
# the prepared, static route.
self._regional_compile_sm100a_enabled: bool | None = None
self._regional_compile_layer_idx: torch.Tensor | None = None
def prepare_for_regional_compile(self, device: torch.device) -> str | None:
"""Resolve the inference-only sm_100a route before fullgraph capture.
The ordinary eager route probes the environment, extension, device,
and tensor contract at every call so it can warn and fall back. Those
Python/device-capability checks are not safe inside a regional
``fullgraph=True`` block. Probe one representative tile-64 input on
the loaded model's device now, then let ``forward`` specialize on the
resulting plain bool while Dynamo is compiling.
"""
requested = os.environ.get(VSA_SM100A_ENV, "0") == "1"
enabled = False
reason = None if requested else f"{VSA_SM100A_ENV}=1 is required for compile-safe VSA-H3 attention"
if requested:
if _sm100a is None:
reason = "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
elif not _sm100a_has_compile_safe_mask_route(_sm100a):
reason = ("neither a native block_sparse_attn_sm100a_from_mask entry nor the raw sm100a "
"kernel plus map_to_index compatibility route is installed")
else:
# Two 64-token blocks exercise the exact sm_100a inference
# specialization while keeping the one-time probe tiny. The
# kernel predicate checks extension presence, CUDA capability,
# dtype/layout, head size, block size, and even block count
# without reading metadata tensor contents.
probe_query = torch.empty((1, 1, 128, self.head_size), device=device, dtype=torch.bfloat16)
probe_block_sizes = torch.full((2, ), 64, device=device, dtype=torch.int32)
reason = _sm100a_unavailable_reason(
_sm100a,
probe_query,
probe_block_sizes,
grad_mode=False,
)
enabled = reason is None
self._regional_compile_sm100a_enabled = enabled
# Keep this marker unset when preparation fails. Generic/training
# torch.compile must retain the established Triton attention route.
self._regional_compile_layer_idx = (torch.tensor(self.layer_idx, device=device, dtype=torch.int64)
if enabled else None)
if enabled:
route = ("native fastvideo-kernel mask entry" if callable(
getattr(_sm100a, "block_sparse_attn_sm100a_from_mask", None)) else
"FastVideo compatibility mask adapter")
logger.info_once(f"VSA-H3 regional compile mask route: {route}")
if requested and reason is not None:
logger.warning_once(f"VSA-H3 regional compile is unavailable and will stay eager: {reason}")
return reason
def tile(self, x: torch.Tensor, attn_metadata: MiniMaxH3VSAMetadata) -> torch.Tensor:
"""Scatter rows into the padded tile buffer (pad positions stay zero).
The returned tensor aliases the builder-owned buffer; callers must
consume it before the next ``tile()`` (both call sites in
``forward()`` read it immediately).
``forward()`` read it immediately). Odd tile-64 no-grad sm100a
requests carry one additional all-zero tile internally; metadata and
all observable outputs retain the logical geometry.
"""
if x.shape[1] != attn_metadata.total_seq_length:
raise ValueError(f"VSA-H3 metadata was built for sequence length {attn_metadata.total_seq_length}, "
f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
"routed to the VSA-H3 backend; exclude it from the supported backends.")
n_tiles = attn_metadata.variable_block_sizes.numel()
target_shape = (x.shape[0], n_tiles * _TILE_ELEMS, x.shape[-2], x.shape[-1])
grad_mode = torch.is_grad_enabled() and x.requires_grad
compiling = torch.compiler.is_compiling()
regional_compiling = compiling and self._regional_compile_layer_idx is not None
if regional_compiling:
sm100a_requested = bool(self._regional_compile_sm100a_enabled)
elif compiling:
# Training/generic compile keeps the long-standing Triton route.
sm100a_requested = False
else:
sm100a_requested = os.environ.get(VSA_SM100A_ENV, "0") == "1"
needs_sm100a_pair = (attn_metadata.tile_elems == 64 and n_tiles % 2 != 0 and not grad_mode and sm100a_requested)
kernel_tiles = n_tiles + int(needs_sm100a_pair)
target_shape = (x.shape[0], kernel_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])
# single scatter: untile_combined_index maps original row i to its
# padded slot, so this is exactly the inverse of postprocess_output
# ``untile_combined_index`` maps each packed row to a logical tile
# slot. Different geometries can share one transport shape; clear a
# reused allocation once when the mapping identity changes so no old
# valid row can survive as padding.
holder = attn_metadata.tile_buf_holder
holder[0] = scatter_into_tile_buf(x, target_shape, attn_metadata.untile_combined_index, holder[0])
return holder[0]
if holder is None:
raise RuntimeError("VSA-H3 metadata has no builder-owned tile buffer holder")
buffer_matches = (holder.buffer is not None and holder.buffer.shape == target_shape
and holder.buffer.dtype == x.dtype and holder.buffer.device == x.device)
if buffer_matches and holder.untile_geometry is not attn_metadata.untile_combined_index:
holder.buffer.zero_()
holder.buffer = scatter_into_tile_buf(x, target_shape, attn_metadata.untile_combined_index, holder.buffer)
holder.untile_geometry = attn_metadata.untile_combined_index
if needs_sm100a_pair:
# A prior even geometry can reuse this allocation and may have
# written the last tile as logical data.
holder.buffer[:, n_tiles * attn_metadata.tile_elems:].zero_()
return holder.buffer
def preprocess_qkv(self, qkv: torch.Tensor, attn_metadata: MiniMaxH3VSAMetadata) -> torch.Tensor:
return self.tile(qkv, attn_metadata)
@@ -281,24 +559,71 @@ class MiniMaxH3VSAImpl(AttentionImpl):
gate_compress: torch.Tensor | None,
attn_metadata: MiniMaxH3VSAMetadata,
) -> torch.Tensor:
if block_sparse_attn_256_bshd is None:
compiling = torch.compiler.is_compiling()
regional_compiling = compiling and self._regional_compile_layer_idx is not None
tile_elems = attn_metadata.tile_elems
if regional_compiling and tile_elems != 64:
raise RuntimeError("VSA-H3 regional fullgraph compile requires 64-token tiles; disable "
"inference_torch_compile for tile-256/CuTe runs.")
if tile_elems == 64:
if block_sparse_attn_64_bhsd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn is not installed")
elif block_sparse_attn_256_bshd is None:
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
# mask) while the rest keep the configured sparsity
layer_sparsity = 0.0 if self.layer_idx in attn_metadata.dense_layers else attn_metadata.VSA_sparsity
probe_dir = probe_enabled()
# Probe recording performs filesystem writes and host synchronizations,
# so the loader keeps probe-enabled runs eager. Avoid even reading that
# environment switch while Dynamo captures a regional full graph.
# The metadata always describes the trained logical geometry.
# ``tile()`` may append exactly one transport-only partner for an odd
# tile-64 sm100a call. Keep score selection and the gate branch on the
# logical prefix, and reject every other shape before a kernel sees it.
n_tiles = attn_metadata.variable_block_sizes.numel()
logical_seq_len = n_tiles * tile_elems
pair_pad_seq_len = logical_seq_len + tile_elems
pair_pad_is_valid = tile_elems == 64 and n_tiles % 2 != 0
allowed_seq_lengths = (logical_seq_len, pair_pad_seq_len) if pair_pad_is_valid else (logical_seq_len, )
if query.shape[1] not in allowed_seq_lengths:
expected = (f"the logical length {logical_seq_len} or one sm100a partner tile "
f"({pair_pad_seq_len})" if pair_pad_is_valid else f"the logical length {logical_seq_len}")
raise ValueError(f"VSA-H3 tiled query has length {query.shape[1]}, expected {expected}.")
has_sm100a_pair = query.shape[1] == pair_pad_seq_len
for name, tensor in (("key", key), ("value", value)):
if tensor.shape[1] != query.shape[1]:
raise ValueError(f"VSA-H3 tiled {name} length {tensor.shape[1]} does not match query "
f"length {query.shape[1]}.")
if gate_compress is not None and gate_compress.shape[1] != query.shape[1]:
raise ValueError(f"VSA-H3 tiled gate length {gate_compress.shape[1]} does not match query "
f"length {query.shape[1]}.")
logical_query = query[:, :logical_seq_len]
logical_key = key[:, :logical_seq_len]
logical_value = value[:, :logical_seq_len]
logical_gate = gate_compress[:, :logical_seq_len] if gate_compress is not None else None
# Probe-guided per-layer opt-out: diffuse layers run dense (all-True
# mask) while the rest keep the configured sparsity. During regional
# capture, keep the layer decision tensor-valued so the 50 block
# instances reuse one graph instead of specializing on layer_idx.
force_dense = None
if regional_compiling:
assert self._regional_compile_layer_idx is not None
force_dense = (attn_metadata.dense_layers_tensor == self._regional_compile_layer_idx).any()
layer_sparsity = attn_metadata.VSA_sparsity
else:
layer_sparsity = 0.0 if self.layer_idx in attn_metadata.dense_layers else attn_metadata.VSA_sparsity
probe_dir = None if compiling else probe_enabled()
scores = None
if layer_sparsity > 0.0 or gate_compress is not None or probe_dir is not None:
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes)
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
q_pooled = _pool_tiles(logical_query, attn_metadata.variable_block_sizes, tile_elems)
k_pooled = _pool_tiles(logical_key, attn_metadata.variable_block_sizes, tile_elems)
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (query.shape[-1]**0.5)
if probe_dir is not None:
record_probe(probe_dir, self.layer_idx, query, key, scores, attn_metadata)
record_probe(probe_dir, self.layer_idx, logical_query, logical_key, scores, attn_metadata)
if scores is None:
n_tiles = attn_metadata.variable_block_sizes.numel()
mask = torch.ones(query.shape[0], query.shape[2], n_tiles, n_tiles, dtype=torch.bool, device=query.device)
else:
mask = _build_block_mask(
@@ -308,24 +633,134 @@ class MiniMaxH3VSAImpl(AttentionImpl):
layer_sparsity,
attn_metadata.exempt,
)
if force_dense is not None:
# A scalar bool tensor broadcasts over the block map. This exactly
# preserves the eager dense-layer contract without a Python branch.
mask = mask | force_dense
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
if tile_elems == 64:
# Native 64-token path: the block map is already at the kernels'
# granularity. Both 64-token entries take BHSD ([B, H, S_pad, D]);
# mirror block_sparse_attn_256_bshd's Triton branch and transpose
# around the call.
q_bhsd = query.transpose(1, 2).contiguous()
k_bhsd = key.transpose(1, 2).contiguous()
v_bhsd = value.transpose(1, 2).contiguous()
if gate_compress is not None:
sm100a_mask = mask
sm100a_variable_block_sizes = attn_metadata.variable_block_sizes
if has_sm100a_pair:
# The synthetic tile is neither a logical query nor key. Its
# all-False row yields q2k_num=0, the all-False column keeps it
# out of real rows, and vbs=0 masks all of its key slots.
sm100a_mask = torch.nn.functional.pad(mask, (0, 1, 0, 1), value=False)
sm100a_variable_block_sizes = torch.nn.functional.pad(
attn_metadata.variable_block_sizes,
(0, 1),
value=0,
)
# 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.
grad_mode = torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad)
use_sm100a = False
if regional_compiling:
if self._regional_compile_sm100a_enabled is None:
raise RuntimeError(
"VSA-H3 sm_100a routing was not resolved before torch.compile; "
"call prepare_for_regional_compile(device) on every MiniMaxH3VSAImpl after loading weights.")
# The preparation probe established module/device/kernel
# support. Keep only static tensor/geometry facts here; no
# env access, device-capability query, or is_supported call may
# enter the Dynamo graph.
if not (self._regional_compile_sm100a_enabled and not grad_mode and q_bhsd.dtype == torch.bfloat16
and q_bhsd.shape[-1] == 128 and sm100a_variable_block_sizes.numel() % 2 == 0):
raise RuntimeError(
"VSA-H3 regional fullgraph compile requires the prepared sm_100a BF16/head-128 route "
"on a supported device; disable inference_torch_compile for this request.")
use_sm100a = True
elif not compiling and os.environ.get(VSA_SM100A_ENV, "0") == "1":
reason = _sm100a_unavailable_reason(_sm100a, q_bhsd, sm100a_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
else:
logger.warning_once(f"{VSA_SM100A_ENV}=1 but falling back to the Triton-64 kernels: {reason}")
if use_sm100a:
# Regional preparation emits the compile-route receipt before
# capture. Logging from this branch would itself break a
# ``fullgraph=True`` forward.
if not compiling:
logger.info_once("MiniMax-H3 VSA tile-64 forward: using the sm100a CUDA block-sparse kernel")
if regional_compiling:
# The compile-safe wrapper keeps both Triton mask
# compaction and the raw pybind launch behind one
# fake-backed custom-op boundary.
out_bhsd, _ = _sm100a_from_mask(
q_bhsd,
k_bhsd,
v_bhsd,
sm100a_mask,
sm100a_variable_block_sizes,
)
else:
# Preserve the established eager/index-native route and
# compatibility with older kernel wheels. Per-row counts
# are non-uniform (prefix queries are dense; video queries
# run prefix+top-k), which the fixed kernel supports.
q2k_idx, q2k_num = map_to_index(sm100a_mask)
out_bhsd, _ = _sm100a.block_sparse_attn_sm100a(
q_bhsd,
k_bhsd,
v_bhsd,
q2k_idx,
q2k_num,
sm100a_variable_block_sizes.to(torch.int32),
need_lse=False,
)
else:
if has_sm100a_pair:
q_bhsd = q_bhsd[:, :, :logical_seq_len].contiguous()
k_bhsd = k_bhsd[:, :, :logical_seq_len].contiguous()
v_bhsd = v_bhsd[:, :, :logical_seq_len].contiguous()
out_bhsd, _ = block_sparse_attn_64_bhsd(
q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
)
if has_sm100a_pair and use_sm100a:
out_bhsd = out_bhsd[:, :, :logical_seq_len]
out = out_bhsd.transpose(1, 2).contiguous()
else:
out, _ = block_sparse_attn_256_bshd(
logical_query,
logical_key,
logical_value,
mask,
attn_metadata.variable_block_sizes,
)
if logical_gate is not None:
# Wan-style compression branch: dense attention over pooled tiles,
# broadcast to each tile's rows, scaled by the learned gate
# (zero-initialized for H3 => branch contributes nothing until
# finetuned; the model layer skips it entirely for all-zero gates).
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
v_pooled = _pool_tiles(logical_value, attn_metadata.variable_block_sizes, tile_elems)
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
batch, seq_len, 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_tiled = out.view(batch, n_tiles, tile_elems, heads, dim)
gate_tiled = logical_gate.view(batch, n_tiles, tile_elems, heads, dim)
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
return out
@@ -60,7 +60,7 @@ def record_probe(
gen = torch.Generator(device="cpu").manual_seed(step * 1000 + layer)
# sample among video rows in the PADDED/tiled domain that are non-pad
from fastvideo.attention.backends.video_sparse_attn_h3 import token_tile_and_valid
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes)
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
video_rows = torch.nonzero((token_tile >= P) & token_valid, as_tuple=False).flatten()
idx = video_rows[torch.randint(0, video_rows.numel(), (_TRUE_ROWS, ), generator=gen).to(query.device)]
+34 -5
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import os
from functools import wraps
import torch
import torch.nn as nn
@@ -20,7 +21,7 @@ def _attention_compile_disabled() -> bool:
Defaults to ``True`` (the historical behavior: attention runs eager via
``torch.compiler.disable``). Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=0``
to let attention be traced/compiled into the surrounding graph.
to let attention instances constructed under that environment be traced.
"""
val = os.environ.get("FASTVIDEO_DISABLE_ATTENTION_COMPILE")
if val is None:
@@ -28,11 +29,32 @@ def _attention_compile_disabled() -> bool:
return val.strip().lower() not in ("0", "false", "no", "off", "")
def _attention_compile_explicitly_disabled() -> bool:
"""Whether the environment explicitly requests the eager boundary.
Regional compile can override the historical default for one loaded
transformer, but it must still honor an explicit debugging escape hatch.
"""
return "FASTVIDEO_DISABLE_ATTENTION_COMPILE" in os.environ and _attention_compile_disabled()
def _maybe_compiler_disable(fn):
"""Apply ``torch.compiler.disable`` unless disabled via env var."""
if _attention_compile_disabled():
return torch.compiler.disable(fn)
return fn
"""Defer the eager/traceable choice to each attention instance.
A class-definition-time choice makes a process-wide default the only
option. The deferred wrapper keeps ordinary instances on the historical
eager boundary while allowing the regional loader to opt in only the
attention modules owned by the transformer it is compiling.
"""
disabled_fn = torch.compiler.disable(fn)
@wraps(fn)
def _dispatch(self, *args, **kwargs):
if self._compile_forward_enabled:
return fn(self, *args, **kwargs)
return disabled_fn(self, *args, **kwargs)
return _dispatch
class DistributedAttention(nn.Module):
@@ -77,6 +99,13 @@ class DistributedAttention(nn.Module):
self.num_kv_heads = num_kv_heads
self.backend = backend_name_to_enum(attn_backend.get_name())
self.dtype = dtype
# Preserve the historical compiler-disabled default. The regional
# inference loader may enable this one instance after validating the
# transformer's resolved backend; no process-global default changes.
self._compile_forward_enabled = not _attention_compile_disabled()
def _set_compile_forward_enabled(self, enabled: bool) -> None:
self._compile_forward_enabled = enabled
@_maybe_compiler_disable
def forward(
@@ -62,14 +62,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
hidden_size: int = 5120
intermediate_size: int = 25600
num_hidden_layers: int = 64
# H3 conditions on one intermediate hidden state and reads nothing above it,
# so the remaining layers are built, weight-loaded and then discarded: 14
# layers, 13.7 GB in bf16. Building exactly this many leaves that hidden
# state bit-identical, because the tuple records each layer's *input*, so
# entry N is the output of layer N-1. Set to None to keep the full stack.
# Must equal MINIMAX_H3_TEXT_ENCODER_LAYER in
# fastvideo/pipelines/basic/minimax_h3/packing.py; a test pins them together
# rather than importing across the models -> pipelines boundary.
output_hidden_state_index: int = 50
num_hidden_layers_override: int | None = 50
num_attention_heads: int = 64
num_key_value_heads: int = 8
@@ -116,7 +109,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
vision_initializer_range: float = 0.02
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
output_hidden_states: bool = True
output_hidden_states: bool = False
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
_fsdp_shard_conditions: list = field(default_factory=lambda: [
_is_language_transformer_layer,
@@ -127,15 +120,16 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
])
def __post_init__(self) -> None:
# Runs both at construction and after ``update_model_arch`` merges the
# checkpoint's config.json, so it also guards config-file overrides. A
# non-positive override would build no decoder layers at all, and a
# negative one would additionally make the surplus-key filter drop
# every ``language_model.layers.*`` checkpoint key, so the conditioner
# would "load" with no transformer stack and only fail at generation.
if self.num_hidden_layers_override is not None and self.num_hidden_layers_override < 1:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be a positive layer count "
f"or None for the full stack; got {self.num_hidden_layers_override}.")
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
if self.num_hidden_layers_override is not None:
if self.num_hidden_layers_override <= 0:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
if self.num_hidden_layers_override < self.output_hidden_state_index:
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
f"hidden_states[{self.output_hidden_state_index}], got "
f"{self.num_hidden_layers_override}.")
rope_scaling = dict(self.rope_scaling or {})
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
@@ -5,6 +5,7 @@ import torch
from torch.distributed import ProcessGroup
from fastvideo.distributed.device_communicators.base_device_communicator import (DeviceCommunicatorBase)
from fastvideo.distributed.device_communicators.ulysses_a2a import maybe_create_helper
class CudaCommunicator(DeviceCommunicatorBase):
@@ -25,6 +26,10 @@ class CudaCommunicator(DeviceCommunicatorBase):
device=self.device,
)
# Capability is agreed once; the persistent window arms on first use.
self.ulysses_a2a = maybe_create_helper(self.cpu_group, self.device_group, self.world_size, self.device,
self.pynccl_comm)
def all_reduce(self, input_, op: torch.distributed.ReduceOp | None = None):
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
@@ -38,6 +43,14 @@ class CudaCommunicator(DeviceCommunicatorBase):
torch.distributed.all_reduce(out, group=self.device_group, op=op)
return out
def all_to_all_4D(self, input_: torch.Tensor, scatter_dim: int = 2, gather_dim: int = 1) -> torch.Tensor:
"""All-to-all over the sequence parallel group, fused when available."""
if self.ulysses_a2a is not None:
output = self.ulysses_a2a.try_all_to_all_4D(input_, scatter_dim, gather_dim)
if output is not None:
return output
return super().all_to_all_4D(input_, scatter_dim, gather_dim)
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
"""Sends a tensor to the destination rank in a non-blocking way"""
"""NOTE: `dst` is the local rank of the destination rank."""
@@ -65,5 +78,9 @@ class CudaCommunicator(DeviceCommunicatorBase):
return tensor
def destroy(self) -> None:
if self.ulysses_a2a is not None:
# The helper still needs its NCCL communicator during teardown.
self.ulysses_a2a.close()
self.ulysses_a2a = None
if self.pynccl_comm is not None:
self.pynccl_comm = None
@@ -0,0 +1,413 @@
# SPDX-License-Identifier: Apache-2.0
"""Fused NVLink all-to-all for Ulysses sequence parallelism.
Drop-in replacement for DistributedAutograd.AllToAll4D when the group is a
load-store accessible NVLink mesh: same layout, byte-identical results, fewer
passes over local memory. Anything else falls back to the NCCL path.
"""
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup
from fastvideo import envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# The kernel is template-specialized on the world size, so only these dispatch.
SUPPORTED_WORLD_SIZES = (2, 4, 6, 8)
_DTYPE_CODES = {
torch.float16: 1,
torch.bfloat16: 2,
torch.float32: 3,
}
# Bound persistent registered memory per rank. Larger operands use NCCL instead
# of growing the window without limit.
MAX_WINDOW_BYTES = 1024**3
# (scatter_dim, gather_dim) -> kernel mode.
# 0: [B, S_local, H, D] -> [B, S_global, H_local, D]
# 1: [B, S_global, H_local, D] -> [B, S_local, H, D]
_MODE_FROM_DIMS = {(2, 1): 0, (1, 2): 1}
def is_enabled() -> bool:
"""Whether the fused path is opted in via FASTVIDEO_ULYSSES_A2A."""
return envs.FASTVIDEO_ULYSSES_A2A == "auto"
class _FusedUlyssesA2A(torch.autograd.Function):
"""Differentiable fused all-to-all.
The two directions are exact inverses, and Ulysses redistributes activations
rather than reducing them, so backward is the opposite mode with no scaling.
"""
@staticmethod
def forward(ctx, helper: "UlyssesA2AHelper", x: torch.Tensor, mode: int) -> torch.Tensor: # type: ignore[override]
ctx.helper = helper
ctx.mode = mode
return helper.run_armed(x, mode)
@staticmethod
def backward(ctx, grad_output: torch.Tensor): # type: ignore[override]
# Same numel and dtype as the forward output, so the window is already
# sized for it; only contiguity needs restoring.
grad_input = ctx.helper.run_armed(grad_output.contiguous(), 1 - ctx.mode)
return None, grad_input, None
class UlyssesA2AHelper:
"""Owns the fused all-to-all context for one sequence-parallel group.
Group capability is agreed during construction; the NCCL window is
registered on first use, once an operand size is known.
"""
def __init__(self, cpu_group: ProcessGroup, device_group: ProcessGroup, world_size: int, device: torch.device,
pynccl_comm):
self.cpu_group = cpu_group
self.device_group = device_group
self.world_size = world_size
self.device = device
self.pynccl_comm = pynccl_comm
self._handle: int | None = None
self._nbytes = 0
self._disabled_reason: str | None = None
if world_size not in SUPPORTED_WORLD_SIZES:
self._disabled_reason = (f"world size {world_size} is not one of "
f"{SUPPORTED_WORLD_SIZES}")
# -- lifecycle -----------------------------------------------------------
def _disable(self, reason: str) -> None:
if self._disabled_reason is None:
self._disabled_reason = reason
logger.info("Ulysses fused all-to-all disabled: %s", reason)
def _comm_ptr(self) -> int:
comm = self.pynccl_comm.comm
return int(getattr(comm, "value", comm))
def _can_attempt(self) -> tuple[bool, str]:
"""Whether this rank could use the fused path, without allocating anything."""
try:
from fastvideo_kernel import comm_ops
if not comm_ops.is_available():
return False, "fastvideo-kernel was built without the Ulysses a2a kernel"
if not comm_ops.lsa_covers_group(self._comm_ptr(), self.world_size):
return False, "the group is not a load-store-accessible (NVLink) mesh"
except Exception as e: # noqa: BLE001
return False, f"backend unavailable ({type(e).__name__}: {e})"
return True, ""
def _agree(self, ok: bool) -> bool:
"""Reduce a local yes/no to a group-wide verdict: True only if all agree."""
vote = torch.tensor([1 if ok else 0], dtype=torch.int32)
dist.all_reduce(vote, op=dist.ReduceOp.MIN, group=self.cpu_group)
return bool(vote.item())
def _allocate(self, nbytes: int) -> int:
"""Allocate locally; split out so allocation-failure tests can inject."""
from fastvideo_kernel import comm_ops
device_index = self.device.index
if device_index is None:
device_index = torch.cuda.current_device()
return comm_ops.allocate(nbytes, self.pynccl_comm.rank, self.world_size, device_index)
def _register_window(self, handle: int) -> None:
"""Register the user window collectively."""
from fastvideo_kernel import comm_ops
comm_ops.register_window(handle, self._comm_ptr())
def _create_dev_comm(self, handle: int) -> None:
"""Create the NCCL device communicator collectively."""
from fastvideo_kernel import comm_ops
comm_ops.create_dev_comm(handle)
def _dispose(self, handle: int, *, synchronize: bool) -> None:
from fastvideo_kernel import comm_ops
if synchronize:
# Kernel launches and copy-out are asynchronous. Do not deregister a
# window that a prior call on this device is still accessing.
torch.cuda.synchronize(self.device)
comm_ops.dispose(handle)
def _dispose_after_failure(self, handle: int | None) -> bool:
"""Best-effort group cleanup after a setup phase failed.
Every rank votes after attempting cleanup, including a rank that never
obtained a local allocation. This keeps the helper permanently disabled
if teardown was not unanimous instead of re-entering with split state.
"""
cleanup_ok = True
if handle is not None:
try:
self._dispose(handle, synchronize=False)
except Exception: # noqa: BLE001 - converted to a group verdict below
cleanup_ok = False
logger.warning("Ulysses partial-context cleanup failed", exc_info=True)
return self._agree(cleanup_ok)
def _call_signature(self, x: torch.Tensor, scatter_dim: int, gather_dim: int) -> tuple[tuple[int, ...], str]:
"""Return a rank-comparable call contract and any local decline reason."""
mode = _MODE_FROM_DIMS.get((scatter_dim, gather_dim))
dtype_code = _DTYPE_CODES.get(x.dtype, 0)
shape = tuple(int(dim) for dim in x.shape) if x.dim() == 4 else (0, 0, 0, 0)
status = 1
reason = ""
if self._disabled_reason is not None:
status, reason = -1, self._disabled_reason
elif not is_enabled():
status, reason = 0, "FASTVIDEO_ULYSSES_A2A is not auto"
elif x.is_cuda and torch.cuda.is_current_stream_capturing():
status, reason = 0, "the current CUDA stream is being captured"
elif mode is None:
status, reason = 0, "unsupported scatter/gather dimensions"
elif x.dim() != 4:
status, reason = 0, "input is not 4-D"
elif dtype_code == 0:
status, reason = 0, f"unsupported dtype {x.dtype}"
elif not x.is_cuda or x.device != self.device:
status, reason = 0, f"input device {x.device} does not match {self.device}"
elif not x.is_contiguous():
status, reason = 0, "input is not contiguous"
elif mode == 0 and shape[2] % self.world_size != 0:
status, reason = 0, "head count is not divisible by the group"
elif mode == 1 and shape[1] % self.world_size != 0:
status, reason = 0, "sequence length is not divisible by the group"
nbytes = int(x.numel() * x.element_size())
if status == 1 and nbytes == 0:
status, reason = 0, "input is empty"
elif status == 1 and nbytes > MAX_WINDOW_BYTES:
status, reason = 0, f"operand exceeds the {MAX_WINDOW_BYTES}-byte window cap"
# status, armed, mode, dtype, B, S, H, D, bytes, capacity. Comparing the
# whole vector prevents equal-size but differently-shaped ranks from
# entering the fused kernel with incompatible address math. CUDA device
# ordinals are deliberately absent: rank-local ordinals normally differ.
signature = (status, int(self._handle is not None), -1 if mode is None else mode, dtype_code, *shape, nbytes,
self._nbytes)
return signature, reason
def _agree_call(self, signature: tuple[int, ...]) -> tuple[bool, bool, bool]:
"""Agree on eligibility and the complete call signature across ranks.
Returns ``(use_fused, permanently_unavailable, lifecycle_consistent)``. This control
collective is intentionally eager-only; compiled regions use the NCCL
implementation before reaching here.
"""
# Host-side Gloo control keeps this agreement outside CUDA graph capture
# and avoids inserting a second NCCL collective ahead of the data path.
local = torch.tensor(signature, dtype=torch.int64)
gathered = torch.empty(self.world_size * local.numel(), dtype=local.dtype)
dist.all_gather_into_tensor(gathered, local, group=self.cpu_group)
contracts = gathered.view(self.world_size, local.numel())
identical = bool(torch.all(contracts == contracts[0]).item())
statuses = contracts[:, 0]
use_fused = identical and bool(torch.all(statuses == 1).item())
permanently_unavailable = bool(torch.any(statuses < 0).item())
lifecycle_consistent = (bool(torch.all(contracts[:, 1] == contracts[0, 1]).item())
and bool(torch.all(contracts[:, -1] == contracts[0, -1]).item()))
return use_fused, permanently_unavailable, lifecycle_consistent
def _build(self, nbytes: int) -> bool:
"""Collectively register the window. Returns True if it is armed."""
handle: int | None = None
allocation_reason = ""
try:
handle = self._allocate(nbytes)
except Exception as e: # noqa: BLE001 - converted to a group verdict below
allocation_reason = f"window allocation failed ({type(e).__name__}: {e})"
# Allocation is local, so vote before any rank enters registration.
if not self._agree(handle is not None):
cleanup_ok = self._dispose_after_failure(handle)
reason = allocation_reason or "a peer rank could not allocate the window"
if not cleanup_ok:
reason += "; partial-context cleanup failed on a peer"
self._disable(reason)
return False
assert handle is not None
window_registered = False
registration_reason = ""
try:
self._register_window(handle)
window_registered = True
except Exception as e: # noqa: BLE001 - converted to a group verdict below
registration_reason = f"window registration failed ({type(e).__name__}: {e})"
if not self._agree(window_registered):
cleanup_ok = self._dispose_after_failure(handle)
reason = registration_reason or "a peer rank could not register the window"
if not cleanup_ok:
reason += "; partial-context cleanup failed on a peer"
self._disable(reason)
return False
dev_comm_created = False
creation_reason = ""
try:
self._create_dev_comm(handle)
dev_comm_created = True
except Exception as e: # noqa: BLE001 - converted to a group verdict below
creation_reason = f"device communicator creation failed ({type(e).__name__}: {e})"
if not self._agree(dev_comm_created):
cleanup_ok = self._dispose_after_failure(handle)
reason = creation_reason or "a peer rank could not create the device communicator"
if not cleanup_ok:
reason += "; partial-context cleanup failed on a peer"
self._disable(reason)
return False
self._handle = handle
self._nbytes = nbytes
logger.info("Ulysses fused all-to-all armed: world_size=%d window=%.0f MiB", self.world_size, nbytes / 2**20)
return True
def close(self) -> bool:
"""Collectively destroy the device communicator and its window.
Returns whether all ranks completed teardown. An armed/unarmed split
cannot safely enter NCCL window deregistration, so that exceptional
state is leaked until process exit and permanently disabled instead of
risking a distributed deadlock.
"""
handle = self._handle
all_armed = self._agree(handle is not None)
all_unarmed = self._agree(handle is None)
if all_unarmed:
self._nbytes = 0
return True
if not all_armed:
self._handle = None
self._nbytes = 0
self._disable("ranks disagreed on whether a fused window was armed during teardown")
return False
assert handle is not None
synchronize_ok = True
try:
torch.cuda.synchronize(self.device)
except Exception: # noqa: BLE001 - converted to a group verdict below
synchronize_ok = False
logger.warning("Ulysses pre-teardown synchronization failed", exc_info=True)
if not self._agree(synchronize_ok):
self._disable("a peer rank could not synchronize before fused-window teardown")
return False
dispose_ok = True
try:
self._dispose(handle, synchronize=False)
except Exception: # noqa: BLE001 - teardown must not mask a real error
dispose_ok = False
logger.warning("Ulysses window deregistration failed", exc_info=True)
group_ok = self._agree(dispose_ok)
# The native disposer consumes the handle even when a cleanup call
# reports an error, so never retry a potentially dangling pointer.
self._handle = None
self._nbytes = 0
if not group_ok:
self._disable("fused-window teardown failed on a peer rank")
return group_ok
# -- collective ----------------------------------------------------------
def run_armed(self, x: torch.Tensor, mode: int) -> torch.Tensor:
"""Run one collective on an already-armed context."""
assert self._handle is not None, "run_armed called on an unarmed helper"
from fastvideo_kernel import comm_ops
w = self.world_size
if mode == 0:
B, S_local, H, D = x.shape
out = torch.empty(B, S_local * w, H // w, D, dtype=x.dtype, device=x.device)
else:
B, S_global, H_local, D = x.shape
S_local, H = S_global // w, H_local * w
out = torch.empty(B, S_local, H, D, dtype=x.dtype, device=x.device)
comm_ops.all_to_all(self._handle, x, out, B, S_local, H, D, mode)
return out
def try_all_to_all_4D(self, x: torch.Tensor, scatter_dim: int, gather_dim: int) -> torch.Tensor | None:
"""Fused collective, or None to let the caller use the NCCL path."""
if self._disabled_reason is not None:
return None
# Python lifecycle checks, votes, and pybind calls are not valid inside
# a fullgraph region. The inherited NCCL path is compiler-visible, so
# regional compile stays fullgraph by declining before any tensor read.
if torch.compiler.is_compiling():
return None
signature, reason = self._call_signature(x, scatter_dim, gather_dim)
use_fused, permanently_unavailable, lifecycle_consistent = self._agree_call(signature)
if not use_fused:
if not lifecycle_consistent:
self.close()
self._disable("ranks disagreed on the fused-window lifecycle")
if permanently_unavailable:
self._disable(reason or "a peer rank cannot use the fused path")
return None
mode = signature[2]
nbytes = signature[-2]
if self._handle is None:
if not self._build(nbytes):
return None
elif nbytes > self._nbytes:
logger.info("Ulysses window grow: %d -> %d bytes", self._nbytes, nbytes)
if not self.close():
return None
if not self._build(nbytes):
return None
return _FusedUlyssesA2A.apply(self, x, mode)
def maybe_create_helper(cpu_group: ProcessGroup | None, device_group: ProcessGroup | None, world_size: int,
device: torch.device | None, pynccl_comm) -> UlyssesA2AHelper | None:
"""Collectively create a helper only when every rank can use it."""
if (world_size <= 1 or cpu_group is None or device_group is None or device is None or device.type != "cuda"):
return None
if not dist.is_initialized():
return None
helper = None
reason = ""
if not is_enabled():
reason = "FASTVIDEO_ULYSSES_A2A is not auto"
elif world_size not in SUPPORTED_WORLD_SIZES:
reason = f"world size {world_size} is not one of {SUPPORTED_WORLD_SIZES}"
elif pynccl_comm is None or pynccl_comm.disabled:
reason = "the group has no usable PyNccl communicator"
else:
try:
candidate = UlyssesA2AHelper(cpu_group, device_group, world_size, device, pynccl_comm)
can_attempt, reason = candidate._can_attempt()
if can_attempt:
helper = candidate
except Exception as e: # noqa: BLE001 - converted to a group verdict below
reason = f"helper construction failed ({type(e).__name__}: {e})"
vote = torch.tensor([int(helper is not None)], dtype=torch.int32)
dist.all_reduce(vote, op=dist.ReduceOp.MIN, group=cpu_group)
if not bool(vote.item()):
if dist.get_rank(cpu_group) == 0:
logger.info("Ulysses fused all-to-all unavailable: %s", reason or "a peer rank declined")
return None
return helper
+4 -2
View File
@@ -635,14 +635,16 @@ class GroupCoordinator:
return self.device_communicator.recv(size, dtype, src)
def destroy(self) -> None:
# First: communicator teardown can be collective, so it needs the
# process groups alive.
if self.device_communicator is not None:
self.device_communicator.destroy()
if self.device_group is not None:
torch.distributed.destroy_process_group(self.device_group)
self.device_group = None
if self.cpu_group is not None:
torch.distributed.destroy_process_group(self.cpu_group)
self.cpu_group = None
if self.device_communicator is not None:
self.device_communicator.destroy()
if self.mq_broadcaster is not None:
self.mq_broadcaster = None
+57
View File
@@ -21,12 +21,20 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN: bool = False
FASTVIDEO_INFERENCE_TORCH_COMPILE: 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_ULYSSES_A2A: str = "off"
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
NVCC_THREADS: str | None = None
CMAKE_BUILD_TYPE: str | None = None
VERBOSE: bool = False
FASTVIDEO_NVTX_PROFILE: bool = False
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
@@ -217,10 +225,59 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# Use FA4's packed-varlen entry point for the long, single-document
# MiniMax-H3 dense DiT self-attention path. This changes floating-point
# reduction order relative to the fixed-length entry point, so it remains
# an explicit inference-only speed/quality opt-in.
"FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN":
lambda: os.getenv("FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN", "0") != "0",
# If set (=1), enable regional (per-transformer-block) fullgraph
# torch.compile for the DiT at inference — the inference-side counterpart
# of the training regional-compile port of hao-ai-lab/FastVideo#1718.
# Equivalent to FastVideoArgs.inference_torch_compile=True (e.g. via
# PipelineSelection.experimental={"inference_torch_compile": True}). VSA
# and other non-fullgraph-traceable attention backends degrade to eager
# with one warning; see _regional_compile_unsupported_reason in
# fastvideo/models/loader/fsdp_load.py.
"FASTVIDEO_INFERENCE_TORCH_COMPILE":
lambda: os.getenv("FASTVIDEO_INFERENCE_TORCH_COMPILE", "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", ""),
# Sequence-parallel all-to-all backend.
# - "off" (default): the NCCL path in DistributedAutograd.AllToAll4D
# - "auto": fused NVLink kernel when the group is a load-store accessible
# mesh of 2/4/6/8 ranks in eager execution, else the NCCL path.
"FASTVIDEO_ULYSSES_A2A":
lambda: os.getenv("FASTVIDEO_ULYSSES_A2A", "off").strip().lower(),
# 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":
+163 -14
View File
@@ -25,6 +25,17 @@ else:
logger = init_logger(__name__)
# Offload flags that trade device memory for host memory. All of them are a loss
# on a device where the two are the same physical pool. Keeping the policy
# centralized lets every loader and stage share one worker-local decision.
UNIFIED_MEMORY_OFFLOAD_FLAGS = (
"dit_layerwise_offload",
"dit_cpu_offload",
"text_encoder_cpu_offload",
"image_encoder_cpu_offload",
"vae_cpu_offload",
)
class ExecutionMode(str, Enum):
"""
@@ -146,6 +157,19 @@ class FastVideoArgs:
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
# Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the
# video VAE's temporal chunks (decode) and clips (reference encode) are
# round-robined across the sequence-parallel ranks and reassembled
# bit-exactly on the group's first rank instead of running serially on
# one rank while the others idle. ``__post_init__`` folds the
# FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars
# into these fields (parse-once, like attention_backend), and
# FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport
# collective ("gather" or "all_gather").
vae_parallel_decode: bool = False
vae_parallel_encode: bool = False
vae_parallel_decode_strategy: str | None = None
# Compilation
# ``enable_torch_compile`` covers the DiT path (transformer,
# transformer_2, and the LTX-2 stage-2 transformer_refine).
@@ -164,11 +188,25 @@ class FastVideoArgs:
torch_compile_kwargs_text_encoder: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_vae: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_audio_vae: dict[str, Any] = field(default_factory=dict)
# Regional (per-transformer-block) fullgraph torch.compile of the DiT at
# inference — the inference-side counterpart of the training regional
# compile ported from hao-ai-lab/FastVideo#1718. Applied by the loader
# right after the transformer loads, with fullgraph=True and inductor
# options {emulate_precision_casts: True} injected (no user kwargs
# needed). MiniMax-H3 VSA is supported only by its compile-safe sm_100a
# tile-64 inference route; other VSA routes degrade the transformer to
# eager with one warning. Dense FA2/FA3/FA4 inference uses compile-visible
# custom-op boundaries. Opt-in via FASTVIDEO_INFERENCE_TORCH_COMPILE=1 (folded in
# __post_init__) or PipelineSelection.experimental
# {"inference_torch_compile": true}. Distinct from ``enable_torch_compile``,
# which keeps the pipeline-level compile semantics.
inference_torch_compile: bool = False
disable_autocast: bool = False
# 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
@@ -271,6 +309,13 @@ class FastVideoArgs:
self._apply_ltx2_vae_overrides()
self._resolve_refine_args()
self._apply_transformer_quant()
if not self.inference_torch_compile:
# Parse-once adapter (same pattern as attention_backend below): the
# environment variable is an input read once here, so the loader
# only ever consults the typed field.
import fastvideo.envs as envs
if envs.FASTVIDEO_INFERENCE_TORCH_COMPILE:
self.inference_torch_compile = True
if self.attention_backend is not None:
# Fail fast on typos instead of silently auto-selecting later.
from fastvideo.attention.selector import coerce_attn_backend
@@ -286,8 +331,27 @@ class FastVideoArgs:
env_backend = envs.FASTVIDEO_ATTENTION_BACKEND
if env_backend is not None and backend_name_to_enum(env_backend) is not None:
self.attention_backend = env_backend
self._fold_vae_parallel_env()
self.check_fastvideo_args()
def _fold_vae_parallel_env(self) -> None:
"""Parse-once adapters for the sequence-parallel VAE env vars."""
import fastvideo.envs as envs
# Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES /
# DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args
# never imports model modules; a unit test pins the two in sync).
strategies = ("gather", "all_gather")
if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE:
self.vae_parallel_decode = True
if not self.vae_parallel_encode and envs.FASTVIDEO_VAE_PARALLEL_ENCODE:
self.vae_parallel_encode = True
if self.vae_parallel_decode_strategy is None:
self.vae_parallel_decode_strategy = envs.FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY or "gather"
if self.vae_parallel_decode_strategy not in strategies:
raise ValueError(f"vae_parallel_decode_strategy must be one of {strategies}, "
f"got {self.vae_parallel_decode_strategy!r}.")
def _apply_transformer_quant(self) -> None:
"""Pin the typed ``transformer_quant`` instance onto ``dit_config``.
@@ -591,6 +655,15 @@ class FastVideoArgs:
help=
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
)
parser.add_argument(
"--inference-torch-compile",
action=StoreBoolean,
default=FastVideoArgs.inference_torch_compile,
help="Regional fullgraph torch.compile of each DiT transformer block at inference "
"(port of the #1718 training-side regional compile). The loader injects fullgraph=True "
"and inductor options {emulate_precision_casts: true}; non-traceable attention backends "
"(VSA) degrade to eager with one warning. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent.",
)
parser.add_argument(
"--dit-cpu-offload",
@@ -631,6 +704,18 @@ class FastVideoArgs:
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
"Should be enabled in almost all cases",
)
parser.add_argument(
"--vae-parallel-decode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 VAE decode chunks across the SP ranks "
"and reassemble bit-exactly on the output rank (default: serial decode on the output rank)",
)
parser.add_argument(
"--vae-parallel-encode",
action=StoreBoolean,
help="With sequence parallelism, round-robin MiniMax-H3 reference-video VAE encode clips across "
"the SP ranks; every rank keeps the identical full encoding (default: serial encode on every rank)",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -644,6 +729,12 @@ class FastVideoArgs:
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
parser.add_argument(
"--VSA-tile-size",
type=int,
default=FastVideoArgs.VSA_tile_size,
help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path",
)
# Master port for distributed training/inference
parser.add_argument(
@@ -770,20 +861,6 @@ class FastVideoArgs:
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
from fastvideo.platforms import current_platform
if current_platform.is_mps():
self.use_fsdp_inference = False
self.dit_layerwise_offload = False
if self.dit_layerwise_offload:
if self.use_fsdp_inference:
logger.warning("dit_layerwise_offload is enabled, automatically disabling use_fsdp_inference.")
self.use_fsdp_inference = False
if self.dit_cpu_offload:
logger.warning("dit_layerwise_offload is enabled, automatically disabling dit_cpu_offload.")
self.dit_cpu_offload = False
# Validate mode and inference_mode consistency
assert isinstance(self.mode, ExecutionMode), f"Mode must be an ExecutionMode enum, got {type(self.mode)}"
assert self.mode in ExecutionMode.choices(), f"Invalid execution mode: {self.mode}"
@@ -800,6 +877,14 @@ class FastVideoArgs:
logger.warning("Mode is '%s' but inference_mode is False. Setting inference_mode to True.", self.mode)
self.inference_mode = True
# Inference policy must wait until a worker owns and binds its device:
# a unified-memory device disables layerwise offload before conflicts
# are resolved, preserving an explicit FSDP request. Training does not
# pass through the inference worker boundary, so retain its historical
# constructor-time normalization.
if not self.inference_mode:
self._resolve_device_offload_conflicts()
if not self.inference_mode:
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
@@ -834,6 +919,70 @@ class FastVideoArgs:
self.pipeline_config.vae_config.load_encoder = True
self.preprocess_config.check_preprocess_config()
def _resolve_device_offload_conflicts(self) -> None:
"""Resolve offload modes after device-local policy has been applied."""
from fastvideo.platforms import current_platform
if current_platform.is_mps():
self.use_fsdp_inference = False
self.dit_layerwise_offload = False
if self.dit_layerwise_offload:
if self.use_fsdp_inference:
logger.warning("dit_layerwise_offload is enabled, automatically disabling use_fsdp_inference.")
self.use_fsdp_inference = False
if self.dit_cpu_offload:
logger.warning("dit_layerwise_offload is enabled, automatically disabling dit_cpu_offload.")
self.dit_cpu_offload = False
def finalize_device_offload_policy(self, device_id: int = 0) -> bool:
"""Apply device-local memory policy, then resolve incompatible modes."""
has_unified_memory = self.disable_offload_on_unified_memory(device_id)
self._resolve_device_offload_conflicts()
return has_unified_memory
def disable_offload_on_unified_memory(self, device_id: int = 0, *, offload_flag: str | None = None) -> bool:
"""Disable host offload after a worker has selected its device.
CUDA's unified-memory probe reads runtime device properties and may
initialize a CUDA context. Callers must therefore use this only inside
a device-owning process, after selecting and binding ``device_id``.
Returning the classification lets direct component-loader callers
apply the same policy to explicit per-call overrides. When
``offload_flag`` is given, the return value says whether this policy
covers that component role.
"""
from fastvideo.platforms import current_platform
cached_device_id = getattr(self, "_unified_memory_device_id", None)
cached_result = getattr(self, "_unified_memory_result", None)
if cached_device_id != device_id or cached_result is None:
cached_result = current_platform.has_unified_memory(device_id)
self._unified_memory_device_id = device_id
self._unified_memory_result = cached_result
if not cached_result:
return False
enabled_flags = [flag for flag in UNIFIED_MEMORY_OFFLOAD_FLAGS if getattr(self, flag)]
if enabled_flags:
try:
device_name = current_platform.get_device_name(device_id)
except Exception:
# Device naming is diagnostic only. NVML can be unavailable on
# an integrated GPU (for example Jetson), and its physical-
# ordinal lookup cannot interpret CUDA_VISIBLE_DEVICES UUID/MIG
# selectors. Neither case should undo an authoritative driver
# classification.
device_name = current_platform.device_name
for flag in enabled_flags:
logger.info(
"Disabling %s: %s has unified memory, so moving weights to the host duplicates "
"them rather than freeing device memory.", flag, device_name)
setattr(self, flag, False)
return offload_flag is None or offload_flag in UNIFIED_MEMORY_OFFLOAD_FLAGS
_current_fastvideo_args = None
+3 -2
View File
@@ -65,8 +65,9 @@ DEFAULT_LOGGING_CONFIG = {
@lru_cache
def _print_info_once(logger: Logger, msg: str) -> None:
# Set the stacklevel to 2 to print the original caller's line info
logger.info(msg, stacklevel=2)
# The process-aware info wrapper owns stacklevel; passing it here would
# supply the keyword twice when the wrapper delegates to Logger.log.
logger.info(msg)
@lru_cache
+10
View File
@@ -24,6 +24,12 @@ from fastvideo.mlx_runtime.checkpoint import (
load_mlx_dit_checkpoint,
save_mlx_dit_checkpoint,
)
from fastvideo.mlx_runtime.checkpoint_compat import (
UnsupportedMLXCheckpointError,
discover_mlx_checkpoint,
raise_if_unsupported_mlx_checkpoint,
resolve_mlx_checkpoint,
)
from fastvideo.mlx_runtime.memory import (
AppliedMemoryLimits,
add_memory_limit_args,
@@ -81,6 +87,7 @@ __all__ = [
"MLXWanTransformerBlock",
"RefinePlan",
"TwoPassResult",
"UnsupportedMLXCheckpointError",
"UnsupportedMLXQuantizationError",
"add_memory_limit_args",
"apply_fast_spatial_upsample",
@@ -92,8 +99,11 @@ __all__ = [
"fastwan_shape",
"fastwan_shape_from_config",
"gib_to_bytes",
"discover_mlx_checkpoint",
"load_mlx_dit_checkpoint",
"load_or_enhance_prompt",
"raise_if_unsupported_mlx_checkpoint",
"resolve_mlx_checkpoint",
"mlx_dit_from_diffusers_safetensors",
"mlx_block_weights_from_diffusers_safetensors",
"mlx_block_weights_from_torch",
+2 -2
View File
@@ -197,8 +197,8 @@ def load_mlx_dit_checkpoint(checkpoint_dir: str | Path, *, compile: bool = False
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}).")
from fastvideo.mlx_runtime.checkpoint_compat import mlx_checkpoint_missing_hint
raise FileNotFoundError(mlx_checkpoint_missing_hint(checkpoint_dir))
manifest = json.loads(manifest_path.read_text())
version = manifest.get("format_version")
+181
View File
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
"""Reject NVIDIA FastWan-QAD checkpoints on the Apple Silicon MLX path.
FastMetal-QAD is the Apple Silicon release: DMD2 students trained on the affine
INT8 grid, shipped as packed ``mlx_dit.safetensors`` + ``mlx_dit.json``.
``FastVideo/FastWan-QAD-1.3B`` and ``FastVideo/FastWan-QAD-FP8-1.3B`` are
NVIDIA-only QAD checkpoints (NVFP4 / FP8). Loading them through the MLX
Diffusers path silently requantizes the wrong weights and produces videos that
ignore the prompt. Fail loudly instead.
"""
from __future__ import annotations
import json
import re
from pathlib import Path
FASTMETAL_COLLECTION_URL = "https://huggingface.co/collections/FastVideo/fastmetal"
FASTMETAL_BLOG_URL = "https://haoailab.com/blogs/fastmetal/"
FASTMETAL_MODEL_IDS = (
"FastVideo/FastMetal-1.3B-QAD",
"FastVideo/FastMetal-5B-QAD",
"FastVideo/FastMetal-14B-QAD",
)
MLX_DIT_MANIFEST = "mlx_dit.json"
MLX_DIT_WEIGHTS = "mlx_dit.safetensors"
CUDA_QAD_OVERLAY_DIR = "generator_inference_transformer"
# Directory / HF-cache names for the NVIDIA FastWan-QAD family. The INT8 Apple
# snapshots used an older FastWan-QAD-INT8-* name; ``int8`` in the path is
# excluded below so those packed MLX checkpoints still load.
_NVIDIA_FASTWAN_QAD_RE = re.compile(
r"(?:^|[/\\._-]|--)fastwan-qad-(?:fp8-)?1\.3b(?:-sa2)?(?:$|[/\\._-]|--)",
re.IGNORECASE,
)
_NVIDIA_QUANT_MARKERS = (
"nvfp4",
"nv_fp4",
"sageattention3",
"attn_qat_infer",
)
class UnsupportedMLXCheckpointError(ValueError):
"""Raised when an NVIDIA FastWan-QAD (or similarly incompatible) tree is used on MLX."""
def is_mlx_dit_checkpoint(path: str | Path) -> bool:
"""Return True if ``path`` is a packed FastMetal / MLX DiT directory."""
checkpoint_dir = Path(path)
return (checkpoint_dir / MLX_DIT_MANIFEST).is_file() and (checkpoint_dir / MLX_DIT_WEIGHTS).is_file()
def discover_mlx_checkpoint(*candidates: str | Path | None) -> Path | None:
"""Return the first candidate that is a packed MLX DiT directory."""
for candidate in candidates:
if candidate is None:
continue
path = Path(candidate)
if is_mlx_dit_checkpoint(path):
return path
return None
def resolve_mlx_checkpoint(explicit: str | Path | None, *search_roots: str | Path | None) -> Path | None:
"""Prefer an explicit ``--mlx-checkpoint``, otherwise scan search roots."""
if explicit is not None:
return Path(explicit)
return discover_mlx_checkpoint(*search_roots)
def mlx_checkpoint_missing_hint(checkpoint_dir: str | Path) -> str:
"""Extra FileNotFoundError text when a directory is not a packed MLX DiT."""
nvidia_reason = nvidia_fastwan_qad_reason(checkpoint_dir)
prefix = (f"Not an MLX DiT checkpoint directory: {checkpoint_dir} "
f"(expected {MLX_DIT_MANIFEST} and {MLX_DIT_WEIGHTS}).")
if nvidia_reason is not None:
return prefix + "\n\n" + _nvidia_fastwan_qad_message(Path(checkpoint_dir), nvidia_reason)
return prefix + "\n\n" + _fastmetal_howto()
def raise_if_unsupported_mlx_checkpoint(*paths: str | Path | None) -> None:
"""Raise if any path is an NVIDIA FastWan-QAD tree that must not run on MLX.
Packed FastMetal / MLX DiT directories are always allowed, including the
older FastWan-QAD-INT8 directory name, because those already contain
``mlx_dit.json``.
"""
for path in paths:
if path is None:
continue
checkpoint = Path(path)
if is_mlx_dit_checkpoint(checkpoint):
continue
reason = nvidia_fastwan_qad_reason(checkpoint)
if reason is None:
continue
raise UnsupportedMLXCheckpointError(_nvidia_fastwan_qad_message(checkpoint, reason))
def nvidia_fastwan_qad_reason(path: str | Path) -> str | None:
"""Return a short reason if ``path`` looks like NVIDIA FastWan-QAD, else None."""
checkpoint = Path(path)
haystack = _path_haystack(checkpoint)
if "int8" in haystack and is_mlx_dit_checkpoint(checkpoint):
return None
if _NVIDIA_FASTWAN_QAD_RE.search(haystack) and "int8" not in haystack:
return "NVIDIA FastWan-QAD checkpoint name (NVFP4/FP8, not FastMetal INT8)"
if (checkpoint / CUDA_QAD_OVERLAY_DIR).is_dir() or (checkpoint.parent / CUDA_QAD_OVERLAY_DIR).is_dir():
return f"CUDA QAD overlay directory ({CUDA_QAD_OVERLAY_DIR}/)"
for config_path in _config_candidates(checkpoint):
marker = _quant_marker_in_file(config_path)
if marker is not None:
return f"{config_path.name} contains NVIDIA quantization marker {marker!r}"
return None
def _path_haystack(path: Path) -> str:
try:
resolved = path.resolve()
except OSError:
resolved = path
return str(resolved).replace("\\", "/").lower()
def _config_candidates(path: Path) -> list[Path]:
roots = [path.parent, path.parent.parent] if path.is_file() else [path, path.parent, path / "transformer"]
seen: set[Path] = set()
files: list[Path] = []
for root in roots:
for candidate in (root / "config.json", root / "model_index.json", root / "transformer" / "config.json"):
if candidate in seen or not candidate.is_file():
continue
seen.add(candidate)
files.append(candidate)
return files
def _quant_marker_in_file(config_path: Path) -> str | None:
try:
payload = config_path.read_text(encoding="utf-8")
except OSError:
return None
lowered = payload.lower()
for marker in _NVIDIA_QUANT_MARKERS:
if marker in lowered:
return marker
try:
parsed = json.loads(payload)
except json.JSONDecodeError:
return None
quant = parsed.get("quantization_config") if isinstance(parsed, dict) else None
if isinstance(quant, dict):
blob = json.dumps(quant).lower()
for marker in ("fp8", "float8", "nvfp4", "fp4"):
if marker in blob:
return marker
return None
def _fastmetal_howto() -> str:
models = ", ".join(FASTMETAL_MODEL_IDS)
return ("The Apple Silicon MLX runtime requires FastMetal-QAD INT8 checkpoints, "
f"not CUDA FastWan-QAD weights.\n"
f" Download: hf download FastVideo/FastMetal-1.3B-QAD --local-dir ./FastMetal-1.3B-QAD\n"
" FastMetal repos ship mlx_dit.json (not transformer/config.json).\n"
f" Run: python examples/inference/basic/mlx_wan_prompt_to_video.py "
f"--model-root ./FastMetal-1.3B-QAD --mlx-checkpoint ./FastMetal-1.3B-QAD\n"
f" 5B uses examples/inference/basic/mlx_wan22_generate.py with FastMetal-5B-QAD.\n"
f" Models: {models}\n"
f" Guide: {FASTMETAL_BLOG_URL}\n"
f" Collection: {FASTMETAL_COLLECTION_URL}\n"
"FastVideo/FastWan-QAD-1.3B and FastVideo/FastWan-QAD-FP8-1.3B are "
"NVIDIA NVFP4/FP8 QAD checkpoints.")
def _nvidia_fastwan_qad_message(path: Path, reason: str) -> str:
return (f"Refusing to load {path} on the Apple Silicon MLX runtime ({reason}).\n\n" + _fastmetal_howto())
+4
View File
@@ -885,6 +885,10 @@ def mlx_dit_from_diffusers_safetensors(
import mlx.core as mx
from safetensors import safe_open
from fastvideo.mlx_runtime.checkpoint_compat import raise_if_unsupported_mlx_checkpoint
raise_if_unsupported_mlx_checkpoint(checkpoint_path, config_path)
config = json.loads(Path(config_path).read_text())
total_blocks = int(config["num_layers"])
if num_blocks is None:
+206 -27
View File
@@ -10,6 +10,7 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo import envs
from fastvideo.attention import DistributedAttention
from fastvideo.attention.layer import DistributedAttention_VSA
from fastvideo.attention.selector import get_attn_backend
@@ -23,12 +24,50 @@ from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.minimax_h3_fusions import (
HAVE_TRITON,
fused_qknorm_rope,
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
minimax_h3_swiglu,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
from fastvideo.utils import get_compute_dtype
logger = init_logger(__name__)
MINIMAX_H3_MODALITY_NUM = 3
_CFG = MiniMaxH3Config()
_MINIMAX_H3_FUSION_NAMES = frozenset({"modulate", "qknorm_rope", "swiglu"})
def _enabled_minimax_h3_fusions(value: str | None = None) -> frozenset[str]:
"""Parse the independently switchable inference fusion set."""
raw = envs.FASTVIDEO_MINIMAX_H3_FUSIONS if value is None else value
normalized = raw.strip().lower()
if normalized in {"", "0", "none"}:
return frozenset()
if normalized in {"1", "all"}:
return _MINIMAX_H3_FUSION_NAMES
enabled = frozenset(item.strip() for item in normalized.split(",") if item.strip())
unknown = enabled - _MINIMAX_H3_FUSION_NAMES
if unknown:
supported = ",".join(sorted(_MINIMAX_H3_FUSION_NAMES))
raise ValueError(f"Unknown MiniMax H3 fusion(s) {sorted(unknown)}; expected a subset of {supported}.")
return enabled
def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool:
"""Triton kernels are inference-only and trace as opaque custom ops.
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()
class MiniMaxH3RotaryPosEmbed(nn.Module):
@@ -62,6 +101,7 @@ class MiniMaxH3FeedForward(nn.Module):
ffn_dim: int,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
fuse_swiglu: bool = False,
) -> None:
super().__init__()
self.fc_in = ReplicatedLinear(
@@ -78,11 +118,15 @@ class MiniMaxH3FeedForward(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.fc_out",
)
self.fuse_swiglu = fuse_swiglu
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.fc_in(hidden_states)
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
if self.fuse_swiglu and _can_run_minimax_h3_fusion(hidden_states):
hidden_states = minimax_h3_swiglu(hidden_states)
else:
hidden_states, gate = hidden_states.chunk(2, dim=-1)
hidden_states = hidden_states * F.silu(gate)
hidden_states, _ = self.fc_out(hidden_states)
return hidden_states
@@ -99,6 +143,8 @@ class MiniMaxH3Attention(nn.Module):
supported_attention_backends: tuple[AttentionBackendEnum, ...],
quant_config: QuantizationConfig | None,
prefix: str,
fuse_qknorm_rope: bool = False,
fa4_packed_varlen: bool = False,
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
@@ -134,6 +180,7 @@ class MiniMaxH3Attention(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.to_out",
)
self.fuse_qknorm_rope = fuse_qknorm_rope
# VSA carries a learned gate on its pooled-compression branch. The H3
# checkpoint has no such weight, so the loader zero-initializes it
# (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
@@ -150,6 +197,7 @@ class MiniMaxH3Attention(nn.Module):
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=prefix,
fa4_packed_varlen=fa4_packed_varlen,
)
self.to_gate_compress: ReplicatedLinear | None = None
# None = unchecked; the first forward tests the loaded weight once and
@@ -175,12 +223,31 @@ class MiniMaxH3Attention(nn.Module):
"""
if torch.is_grad_enabled():
return True
if self._gate_compress_active is None:
if torch.compiler.is_compiling():
raise RuntimeError(
"MiniMax H3 VSA compression gate was not resolved before torch.compile; "
"call prepare_for_compile() after loading weights.")
self._resolve_gate_compress_for_compile()
assert self._gate_compress_active is not None
return self._gate_compress_active
def _resolve_gate_compress_for_compile(self) -> None:
"""Resolve the inference-only compression-gate branch eagerly.
Regional compilation wraps each transformer block with
``fullgraph=True``. Resolving the loaded weight before capture keeps
the GPU-to-host bool conversion and the cache mutation out of every
compiled block graph. Grad-enabled forwards ignore this cached value
in ``_gate_active`` so training can still learn from a zero gate.
"""
if self.to_gate_compress is None:
return
if self._gate_compress_active is None:
weight = self.to_gate_compress.weight
# bool() on a DTensor reduction resolves collectively, so every
# rank caches the same answer.
self._gate_compress_active = bool((weight != 0).any())
return self._gate_compress_active
@staticmethod
def _apply_rotary_emb(
@@ -211,11 +278,18 @@ class MiniMaxH3Attention(nn.Module):
query = query.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
key = key.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
value = value.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
if (self.fuse_qknorm_rope and rotary_emb is not None and _can_run_minimax_h3_fusion(query)):
cos, sin = rotary_emb
cos = cos.to(query.dtype)
sin = sin.to(query.dtype)
query = fused_qknorm_rope(query, self.norm_q.weight, cos, sin, self.norm_q.eps)
key = fused_qknorm_rope(key, self.norm_k.weight, cos, sin, self.norm_k.eps)
else:
query = self.norm_q(query)
key = self.norm_k(key)
if rotary_emb is not None:
query = self._apply_rotary_emb(query, rotary_emb)
key = self._apply_rotary_emb(key, rotary_emb)
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
# branch cannot express. Apply it above, then pass no RoPE here.
@@ -397,6 +471,10 @@ 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,
fa4_packed_varlen: bool = False,
) -> None:
super().__init__()
self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
@@ -408,6 +486,8 @@ class MiniMaxH3TransformerBlock(nn.Module):
supported_attention_backends,
quant_config,
prefix=f"{prefix}.attn",
fuse_qknorm_rope=fuse_qknorm_rope,
fa4_packed_varlen=fa4_packed_varlen,
)
self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
self.ff = MiniMaxH3FeedForward(
@@ -415,6 +495,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
ffn_dim,
quant_config=quant_config,
prefix=f"{prefix}.ff",
fuse_swiglu=fuse_swiglu,
)
self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
time_embed_dim,
@@ -423,6 +504,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
prefix=f"{prefix}.adaln_proj",
apply_silu=adaln_apply_silu,
)
self.fuse_modulate = fuse_modulate
def forward(
self,
@@ -435,19 +517,39 @@ class MiniMaxH3TransformerBlock(nn.Module):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
t.to(hidden_states.dtype) for t in self.adaln_proj(temb))
residual = hidden_states
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
use_modulate_fusion = self.fuse_modulate and _can_run_minimax_h3_fusion(hidden_states)
if use_modulate_fusion:
norm_hidden_states = fused_rmsnorm_modulate(
hidden_states,
self.norm1.weight,
scale_msa,
shift_msa,
adaln_indices,
self.norm1.eps,
)
else:
norm_hidden_states = self.norm1(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attention_output
residual = hidden_states
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
if use_modulate_fusion:
hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate(
hidden_states,
attention_output,
gate_msa,
self.norm2.weight,
scale_mlp,
shift_mlp,
adaln_indices,
self.norm2.eps,
)
else:
hidden_states = hidden_states + gate_msa.index_select(0, adaln_indices) * attention_output
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
feed_forward_output = self.ff(norm_hidden_states)
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
class MiniMaxH3Transformer3DModel(BaseDiT):
@@ -493,6 +595,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
super().__init__(config, hf_config)
arch = config.arch_config
self.enabled_fusions = _enabled_minimax_h3_fusions()
if self.enabled_fusions:
if HAVE_TRITON:
logger.info(
"MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled forwards "
"fall back to eager; torch.compile captures opaque custom-op boundaries).",
",".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 "
@@ -590,6 +703,10 @@ 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,
fa4_packed_varlen=envs.FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN,
) for index in range(arch.num_layers)
])
self.norm_out = MiniMaxH3AdaLayerNormOut(
@@ -616,6 +733,65 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
)
self.__post_init__()
def prepare_for_compile(self) -> None:
"""Pipeline hook, called once right before torch.compile wraps the blocks.
Resolve each loaded VSA compression gate eagerly. Generic and training
compile retain their established attention dispatch; only the
inference loader's separate ``prepare_for_regional_compile`` hook may
preselect the inference-only sm_100a path.
The inference-only Triton fusions expose fake-backed custom operators,
so Dynamo can keep them active as opaque nodes inside each fullgraph
block instead of tracing into their launcher implementation.
"""
gate_states: list[bool] = []
for block in self.transformer_blocks:
attention = block.attn
if attention.to_gate_compress is not None:
attention._resolve_gate_compress_for_compile()
assert attention._gate_compress_active is not None
gate_states.append(attention._gate_compress_active)
if gate_states:
logger.info(
"Resolved MiniMax H3 VSA compression gates before torch.compile: %d active, %d inactive",
sum(gate_states),
len(gate_states) - sum(gate_states),
)
if self.enabled_fusions:
logger.info(
"MiniMax H3 inference fusions remain active under torch.compile through custom-op boundaries: %s",
",".join(sorted(self.enabled_fusions)),
)
def prepare_for_regional_compile(self) -> str | None:
"""Resolve state used only by inference regional fullgraph compile."""
self.prepare_for_compile()
prepared_vsa_impls = 0
unsupported_reasons: set[str] = set()
for block in self.transformer_blocks:
attention = block.attn
prepare_vsa = getattr(attention.distributed_attention.attn_impl, "prepare_for_regional_compile", None)
if not callable(prepare_vsa):
continue
# Post-load FP8 conversion may replace to_q.weight with packed
# buffers. Either representation identifies the local device.
query_state = next(attention.to_q.parameters(), None)
if query_state is None:
query_state = next(attention.to_q.buffers(), None)
if query_state is None:
raise RuntimeError("MiniMax H3 to_q has no materialized parameter or buffer for compile setup.")
unsupported = prepare_vsa(query_state.device)
if unsupported:
unsupported_reasons.add(str(unsupported))
prepared_vsa_impls += 1
if prepared_vsa_impls:
logger.info("Prepared %d MiniMax H3 VSA attention implementations for regional torch.compile",
prepared_vsa_impls)
if unsupported_reasons:
return "; ".join(sorted(unsupported_reasons))
return None
def materialize_non_persistent_buffers(
self,
device: torch.device,
@@ -734,14 +910,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
rotary_emb = (rotary_cos, rotary_sin)
for block in self.transformer_blocks:
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
# The eager driver owns profiling markers while each block's compiled
# forward owns the graph that the marker surrounds.
for block_index, block in enumerate(self.transformer_blocks):
with nvtx_range(f"minimax_h3.transformer_block.{block_index}"):
packed_hidden_states = block(
packed_hidden_states,
temb,
adaln_indices,
rotary_emb,
original_seq_len,
)
packed_hidden_states = self.norm_out(
packed_hidden_states,
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
"""Inference-only MiniMax H3 fusions adapted from NVlabs/Sana Sol-Engine.
Source: https://github.com/NVlabs/Sana/tree/sol-engine/models/minimax_h3/GB200
"""
from .modulation import (
fused_residual_gate_rmsnorm_modulate,
fused_rmsnorm_modulate,
)
from .qknorm_rope import HAVE_TRITON, fused_qknorm_rope
from .swiglu import minimax_h3_swiglu
__all__ = [
"HAVE_TRITON",
"fused_qknorm_rope",
"fused_residual_gate_rmsnorm_modulate",
"fused_rmsnorm_modulate",
"minimax_h3_swiglu",
]
@@ -0,0 +1,425 @@
# 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_impl(
x: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> torch.Tensor:
_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_impl(
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]:
_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)
# Dynamo must see the Triton launches as opaque nodes. In eager mode the
# public wrappers below keep their strict, actionable input validation; while
# compiling, these custom-op boundaries avoid tracing into Triton's launcher
# and let the surrounding H3 block remain one full graph.
@torch.library.custom_op(
"fastvideo::_minimax_h3_rmsnorm_modulate",
mutates_args=(),
device_types="cuda",
)
def _fused_rmsnorm_modulate_op(
x: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> torch.Tensor:
return _fused_rmsnorm_modulate_impl(x, weight, scale, shift, index, eps)
@torch.library.register_fake("fastvideo::_minimax_h3_rmsnorm_modulate")
def _fused_rmsnorm_modulate_fake(
x: torch.Tensor,
weight: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
index: torch.Tensor,
eps: float,
) -> torch.Tensor:
del weight, scale, shift, index, eps
return x.new_empty(x.shape)
@torch.library.custom_op(
"fastvideo::_minimax_h3_residual_gate_rmsnorm_modulate",
mutates_args=(),
device_types="cuda",
)
def _fused_residual_gate_rmsnorm_modulate_op(
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]:
return _fused_residual_gate_rmsnorm_modulate_impl(
residual,
branch,
gate,
weight,
scale,
shift,
index,
eps,
)
@torch.library.register_fake("fastvideo::_minimax_h3_residual_gate_rmsnorm_modulate")
def _fused_residual_gate_rmsnorm_modulate_fake(
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]:
del branch, gate, weight, scale, shift, index, eps
return residual.new_empty(residual.shape), residual.new_empty(residual.shape)
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).
"""
_require_forward_only(x, weight, scale, shift)
if torch.compiler.is_compiling():
return torch.ops.fastvideo._minimax_h3_rmsnorm_modulate(x, weight, scale, shift, index, eps)
return _fused_rmsnorm_modulate_impl(x, weight, scale, shift, index, eps)
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.
"""
_require_forward_only(residual, branch, gate, weight, scale, shift)
if torch.compiler.is_compiling():
return torch.ops.fastvideo._minimax_h3_residual_gate_rmsnorm_modulate(
residual,
branch,
gate,
weight,
scale,
shift,
index,
eps,
)
return _fused_residual_gate_rmsnorm_modulate_impl(
residual,
branch,
gate,
weight,
scale,
shift,
index,
eps,
)
@@ -0,0 +1,216 @@
# 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_impl(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
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)
@torch.library.custom_op(
"fastvideo::_minimax_h3_qknorm_rope",
mutates_args=(),
device_types="cuda",
)
def _fused_qknorm_rope_op(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
return _fused_qknorm_rope_impl(x, weight, cos, sin, eps)
@torch.library.register_fake("fastvideo::_minimax_h3_qknorm_rope")
def _fused_qknorm_rope_fake(
x: torch.Tensor,
weight: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
eps: float,
) -> torch.Tensor:
del weight, cos, sin, eps
return x.new_empty(x.shape)
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.
"""
if torch.is_grad_enabled() and any(
isinstance(tensor, torch.Tensor) and tensor.requires_grad for tensor in (x, weight, cos, sin)):
raise RuntimeError("fused_qknorm_rope is inference-only and does not implement autograd")
if torch.compiler.is_compiling():
return torch.ops.fastvideo._minimax_h3_qknorm_rope(x, weight, cos, sin, eps)
return _fused_qknorm_rope_impl(x, weight, cos, sin, eps)
__all__ = ["HAVE_TRITON", "fused_qknorm_rope"]
@@ -0,0 +1,126 @@
# 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_impl(x: torch.Tensor) -> torch.Tensor:
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)
@torch.library.custom_op(
"fastvideo::_minimax_h3_swiglu",
mutates_args=(),
device_types="cuda",
)
def _minimax_h3_swiglu_op(x: torch.Tensor) -> torch.Tensor:
return _minimax_h3_swiglu_impl(x)
@torch.library.register_fake("fastvideo::_minimax_h3_swiglu")
def _minimax_h3_swiglu_fake(x: torch.Tensor) -> torch.Tensor:
return x.new_empty((*x.shape[:-1], x.shape[-1] // 2))
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.
"""
if torch.is_grad_enabled() and x.requires_grad:
raise RuntimeError("MiniMax H3 fused SwiGLU is forward-only and does not implement autograd")
if torch.compiler.is_compiling():
return torch.ops.fastvideo._minimax_h3_swiglu(x)
return _minimax_h3_swiglu_impl(x)
__all__ = ["HAVE_TRITON", "minimax_h3_swiglu"]
+8 -8
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from dataclasses import field
from typing import Any, Generic, TypeVar
import torch
from torch import nn
@@ -8,11 +9,16 @@ from torch import nn
from fastvideo.configs.models.encoders import (BaseEncoderOutput, ImageEncoderConfig, TextEncoderConfig)
from fastvideo.platforms import AttentionBackendEnum
TextEncoderOutputT = TypeVar("TextEncoderOutputT")
class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
"""Base for native encoders with a model-specific forward output contract."""
class TextEncoder(nn.Module, ABC):
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = TextEncoderConfig()._supported_attention_backends
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -23,13 +29,7 @@ class TextEncoder(nn.Module, ABC):
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
@abstractmethod
def forward(self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs) -> BaseEncoderOutput:
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
pass
@property
@@ -0,0 +1,470 @@
# SPDX-License-Identifier: Apache-2.0
"""Serialized block-FP8 execution for the MiniMax-H3 Qwen3-VL encoder."""
from typing import Any
import torch
from torch import nn
from torch.nn.parameter import Parameter
try:
import triton
import triton.language as tl
except ImportError:
triton = None
tl = None
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.linear import LinearBase, LinearMethodBase
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.layers.quantization.fp8_config import FP8_DTYPE
from fastvideo.models.utils import set_weight_attrs
class MiniMaxH3SerializedFP8Config(QuantizationConfig):
"""Serialized 128x128 block-FP8 contract for the H3 text encoder."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
if weight_block_size != (128, 128):
raise ValueError("MiniMax-H3 serialized FP8 requires weight_block_size=[128, 128], "
f"got {list(weight_block_size)}")
self.weight_block_size = weight_block_size
self.is_checkpoint_fp8_serialized = True
self.activation_scheme = "dynamic"
@classmethod
def get_name(cls) -> str:
return "fp8"
@classmethod
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
return [torch.bfloat16]
@classmethod
def get_min_capability(cls) -> int:
return 100
@staticmethod
def get_config_filenames() -> list[str]:
return []
@classmethod
def from_config(cls, config: dict[str, Any]) -> "MiniMaxH3SerializedFP8Config":
quant_method = str(config.get("quant_method", "")).lower()
if quant_method != "fp8":
raise ValueError(f"MiniMax-H3 only supports serialized FP8 text-encoder checkpoints, got {quant_method!r}")
if str(config.get("activation_scheme", "")).lower() != "dynamic":
raise ValueError("MiniMax-H3 serialized FP8 requires dynamic activation quantization")
if str(config.get("fmt", "e4m3")).lower() not in ("e4m3", "float8_e4m3fn"):
raise ValueError(f"MiniMax-H3 serialized FP8 requires E4M3 weights, got {config.get('fmt')!r}")
block_size = config.get("weight_block_size")
if not isinstance(block_size, list | tuple) or len(block_size) != 2:
raise ValueError("MiniMax-H3 serialized FP8 requires a two-dimensional weight_block_size")
ignored_layers = config.get("modules_to_not_convert", config.get("ignored_layers", []))
if not isinstance(ignored_layers, list | tuple):
raise ValueError("MiniMax-H3 serialized FP8 modules_to_not_convert must be a sequence")
language_exclusions = [
name for name in ignored_layers
if isinstance(name, str) and (name.startswith("language_model.") or ".language_model." in name)
]
if language_exclusions:
raise ValueError("MiniMax-H3 does not support partially quantized language stacks; "
f"ignored language layers: {language_exclusions[:3]}")
if not any(isinstance(name, str) and "visual" in name for name in ignored_layers):
raise ValueError("MiniMax-H3 serialized FP8 requires the vision stack to be listed in "
"modules_to_not_convert")
return cls((int(block_size[0]), int(block_size[1])))
def validate_runtime(self, device: torch.device) -> None:
if device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires a CUDA device; "
f"got {device.type!r}")
capability = torch.cuda.get_device_capability(device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < self.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{self.get_min_capability()} or newer, got sm{capability_number}")
if capability[0] not in (10, 12):
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 currently adapts SGLang's Blackwell "
f"FlashInfer path; got unsupported sm{capability_number}")
_require_sglang_per_token_group_fp8_quantization()
_get_flashinfer_groupwise_fp8_gemm()
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
if isinstance(layer, LinearBase) and ".language_model.layers." in prefix:
return MiniMaxH3SerializedFP8LinearMethod(self.weight_block_size)
return None
# Copyright 2024 SGLang Team
# Licensed under the Apache License, Version 2.0.
# Adapted from SGLang's per-token-group quantization kernels and Blackwell
# FlashInfer dispatch at commit f99c62063c7dcfcd06784b885dc08cb52cf23865:
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/kernels/ops/quantization/fp8_kernel.py
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/srt/layers/quantization/fp8_utils.py
if triton is not None:
@triton.jit
def _h3_per_token_group_quant_fp8_row_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
scale_ptr += group_id
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
@triton.jit
def _h3_per_token_group_quant_fp8_column_major(
input_ptr,
output_ptr,
scale_ptr,
group_size,
input_columns,
scale_column_stride,
eps,
fp8_min,
fp8_max,
BLOCK: tl.constexpr,
):
group_id = tl.program_id(0)
input_ptr += group_id.to(tl.int64) * group_size
output_ptr += group_id.to(tl.int64) * group_size
groups_per_row = input_columns // group_size
scale_column = group_id % groups_per_row
scale_row = group_id // groups_per_row
scale_ptr += scale_column * scale_column_stride + scale_row
offsets = tl.arange(0, BLOCK)
mask = offsets < group_size
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
scale = absmax / fp8_max
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
tl.store(output_ptr + offsets, quantized, mask=mask)
tl.store(scale_ptr, scale)
else:
_h3_per_token_group_quant_fp8_row_major = None
_h3_per_token_group_quant_fp8_column_major = None
def _require_sglang_per_token_group_fp8_quantization() -> None:
if (triton is None or _h3_per_token_group_quant_fp8_row_major is None
or _h3_per_token_group_quant_fp8_column_major is None):
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires Triton for SGLang-compatible "
"per-token-group activation quantization")
def _sglang_per_token_group_quant_fp8(
input_tensor: torch.Tensor,
group_size: int,
*,
column_major_scales: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""SGLang-compatible dynamic FP8 quantization for contiguous 2-D activations."""
_require_sglang_per_token_group_fp8_quantization()
if input_tensor.ndim != 2:
raise ValueError(f"per-token-group FP8 quantization expects 2-D input, got {input_tensor.ndim}-D")
if not input_tensor.is_contiguous():
raise ValueError("per-token-group FP8 quantization requires contiguous input")
if input_tensor.shape[-1] % group_size:
raise ValueError(f"activation width {input_tensor.shape[-1]} is not divisible by group_size={group_size}")
quantized = torch.empty_like(input_tensor, dtype=FP8_DTYPE)
rows, columns = input_tensor.shape
groups_per_row = columns // group_size
if column_major_scales:
scales = torch.empty(
(groups_per_row, rows),
device=input_tensor.device,
dtype=torch.float32,
).permute(1, 0)
else:
scales = torch.empty(
(rows, groups_per_row),
device=input_tensor.device,
dtype=torch.float32,
)
if rows:
num_groups = input_tensor.numel() // group_size
block = triton.next_power_of_2(group_size)
num_warps = min(max(block // 256, 1), 8)
if column_major_scales:
_h3_per_token_group_quant_fp8_column_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
columns,
scales.stride(1),
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
else:
_h3_per_token_group_quant_fp8_row_major[(num_groups,)](
input_tensor,
quantized,
scales,
group_size,
1e-10,
-448.0,
448.0,
BLOCK=block,
num_warps=num_warps,
num_stages=1,
)
return quantized, scales
def _get_flashinfer_groupwise_fp8_gemm():
try:
from flashinfer.gemm import gemm_fp8_nt_groupwise
except (AttributeError, ImportError) as error:
raise RuntimeError(
"MiniMax-H3 serialized blockwise FP8 requires "
"flashinfer.gemm.gemm_fp8_nt_groupwise (validated with flashinfer-python==0.6.8). "
"FastVideo will not re-quantize this checkpoint to tensorwise FP8.") from error
return gemm_fp8_nt_groupwise
def _get_flashinfer_groupwise_backend(device: torch.device) -> str:
capability = torch.cuda.get_device_capability(device)
if capability[0] >= 12:
return "cutlass"
if capability[0] == 10:
return "trtllm"
capability_number = capability[0] * 10 + capability[1]
raise RuntimeError(f"FlashInfer groupwise FP8 requires a Blackwell GPU, got sm{capability_number}")
def _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
input_tensor: torch.Tensor,
weight: torch.Tensor,
block_size: tuple[int, int],
weight_scale: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
input_2d = input_tensor.view(-1, input_tensor.shape[-1])
output_shape = [*input_tensor.shape[:-1], weight.shape[0]]
backend = _get_flashinfer_groupwise_backend(input_tensor.device)
if input_2d.dtype != torch.bfloat16:
raise RuntimeError("MiniMax-H3 FlashInfer groupwise FP8 requires BF16 activations; "
f"got {input_2d.dtype}. The SGLang FP16 Triton GEMM fallback is not enabled for H3.")
if backend == "trtllm" and input_2d.shape[1] < 256:
raise RuntimeError("MiniMax-H3 FlashInfer TRTLLM groupwise FP8 requires K >= 256; "
f"got K={input_2d.shape[1]}. The SGLang Triton GEMM fallback is not enabled for H3.")
gemm_fp8_nt_groupwise = _get_flashinfer_groupwise_fp8_gemm()
block_n, block_k = block_size
q_input, x_scale = _sglang_per_token_group_quant_fp8(
input_2d,
block_k,
column_major_scales=(backend == "trtllm"),
)
if backend == "cutlass":
m, k = input_2d.shape
n = weight.shape[0]
if x_scale.shape == (m, k // block_k):
x_scale = x_scale.transpose(-1, -2).contiguous()
if weight_scale.shape == (n // block_n, k // block_k):
weight_scale = weight_scale.transpose(-1, -2).contiguous()
# FlashInfer documents that ``m`` "should be padded to a multiple of 4
# before calling this function", and the CUTLASS groupwise module
# rejects unaligned m at dispatch (``cutlass gemm.can_implement
# failed``). Prompt token counts are arbitrary, so pad the quantized
# activation rows and their per-token scales with zeros, then slice
# the padded rows off the GEMM output. The trtllm (sm100) route
# tolerates any m and stays unpadded.
padded_m = (m + 3) // 4 * 4
if padded_m != m:
padded_q_input = q_input.new_zeros((padded_m, q_input.shape[1]))
padded_q_input[:m] = q_input
q_input = padded_q_input
padded_x_scale = x_scale.new_zeros((x_scale.shape[0], padded_m))
padded_x_scale[:, :m] = x_scale
x_scale = padded_x_scale
expected_x_scale_shape = (k // block_k, padded_m)
expected_weight_scale_shape = (k // block_k, n // block_n)
if x_scale.shape != expected_x_scale_shape or weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scale layout mismatch: "
f"x_scale={tuple(x_scale.shape)}, weight_scale={tuple(weight_scale.shape)}, "
f"expected={expected_x_scale_shape}/{expected_weight_scale_shape}")
if x_scale.dtype != torch.float32 or weight_scale.dtype != torch.float32:
raise RuntimeError("FlashInfer CUTLASS block-FP8 scales must be float32")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale.contiguous(),
weight_scale.contiguous(),
out_dtype=input_2d.dtype,
backend="cutlass",
scale_major_mode="MN",
)
if padded_m != m:
output = output[:m]
else:
expected_x_scale_shape = (input_2d.shape[0], input_2d.shape[1] // block_k)
expected_weight_scale_shape = (weight.shape[0] // block_n, weight.shape[1] // block_k)
if x_scale.shape != expected_x_scale_shape or x_scale.stride(0) != 1:
raise RuntimeError("FlashInfer TRTLLM block-FP8 activation scale layout mismatch: "
f"shape={tuple(x_scale.shape)}, stride={x_scale.stride()}, "
f"expected column-major {expected_x_scale_shape}")
if weight_scale.shape != expected_weight_scale_shape:
raise RuntimeError("FlashInfer TRTLLM block-FP8 weight scale layout mismatch: "
f"shape={tuple(weight_scale.shape)}, expected={expected_weight_scale_shape}")
output = gemm_fp8_nt_groupwise(
q_input,
weight,
x_scale,
weight_scale,
out_dtype=input_2d.dtype,
backend="trtllm",
)
if bias is not None:
output += bias
return output.to(dtype=input_2d.dtype).view(*output_shape)
class MiniMaxH3SerializedFP8LinearMethod(LinearMethodBase):
"""Execute serialized 128x128 block-FP8 weights without re-quantizing them."""
def __init__(self, weight_block_size: tuple[int, int]) -> None:
super().__init__()
self.weight_block_size = weight_block_size
def create_weights(
self,
layer: torch.nn.Module,
input_size_per_partition: int,
output_partition_sizes: list[int],
input_size: int,
output_size: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
output_size_per_partition = sum(output_partition_sizes)
block_n, block_k = self.weight_block_size
tp_size = get_tp_world_size()
if tp_size > 1 and input_size // input_size_per_partition == tp_size:
if input_size_per_partition % block_k:
raise ValueError(f"Weight input_size_per_partition={input_size_per_partition} is not divisible "
f"by block_k={block_k}")
if tp_size > 1 and output_size // output_size_per_partition == tp_size:
for output_partition_size in output_partition_sizes:
if output_partition_size % block_n:
raise ValueError(f"Weight output_partition_size={output_partition_size} is not divisible "
f"by block_n={block_n}")
layer.logical_widths = output_partition_sizes
layer.input_size_per_partition = input_size_per_partition
layer.output_size_per_partition = output_size_per_partition
layer.orig_dtype = params_dtype
weight_loader = extra_weight_attrs.get("weight_loader")
weight = Parameter(
torch.empty(output_size_per_partition, input_size_per_partition, dtype=FP8_DTYPE),
requires_grad=False,
)
set_weight_attrs(weight, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
layer.register_parameter("weight", weight)
scale = Parameter(
torch.empty((output_size_per_partition + block_n - 1) // block_n,
(input_size_per_partition + block_k - 1) // block_k,
dtype=torch.float32),
requires_grad=False,
)
set_weight_attrs(scale, {
"input_dim": 1,
"output_dim": 0,
"weight_loader": weight_loader,
})
scale.data.fill_(torch.finfo(torch.float32).min)
layer.register_parameter("weight_scale_inv", scale)
layer.register_parameter("input_scale", None)
def process_weights_after_loading(self, layer: nn.Module) -> None:
weight = getattr(layer, "weight", None)
block_scales = getattr(layer, "weight_scale_inv", None)
if weight is None or block_scales is None:
raise ValueError("Serialized MiniMax-H3 FP8 linear is missing weight or weight_scale_inv")
if weight.dtype != FP8_DTYPE:
raise ValueError(f"Serialized MiniMax-H3 FP8 weight must be {FP8_DTYPE}, got {weight.dtype}")
if block_scales.dtype != torch.float32:
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must be float32, "
f"got {block_scales.dtype}")
block_n, block_k = self.weight_block_size
output_size, input_size = weight.shape
if output_size % block_n or input_size % block_k:
raise ValueError("Serialized MiniMax-H3 FP8 weight dimensions must be divisible by the 128x128 block size; "
f"got {tuple(weight.shape)}")
expected_scale_shape = (output_size // block_n, input_size // block_k)
if tuple(block_scales.shape) != expected_scale_shape:
raise ValueError("Serialized MiniMax-H3 FP8 scale shape mismatch: "
f"expected {expected_scale_shape}, got {tuple(block_scales.shape)}")
if not bool(torch.isfinite(block_scales).all()) or bool((block_scales <= 0).any()):
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must contain finite positive values")
layer.weight.data = weight.data
layer.weight_scale_inv.data = block_scales.data
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
if x.device.type != "cuda":
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 execution requires CUDA")
capability = torch.cuda.get_device_capability(x.device)
capability_number = capability[0] * 10 + capability[1]
if capability_number < MiniMaxH3SerializedFP8Config.get_min_capability():
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
f"sm{MiniMaxH3SerializedFP8Config.get_min_capability()} or newer, "
f"got sm{capability_number}")
if not x.is_contiguous():
x = x.contiguous()
return _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
x,
layer.weight,
self.weight_block_size,
layer.weight_scale_inv,
bias,
)
__all__ = [
"MiniMaxH3SerializedFP8Config",
"MiniMaxH3SerializedFP8LinearMethod",
]
@@ -8,13 +8,13 @@ import torch
import torch.nn.functional as F
from torch import nn
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
from fastvideo.distributed import get_tp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import MiniMaxH3SerializedFP8Config
from fastvideo.models.loader.weight_utils import default_weight_loader
@@ -227,19 +227,13 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
org_num_embeddings=config.vocab_size,
quant_config=quant_config,
)
# Build only as far as the consumer reads. The hidden-state tuple records
# each layer's input, so stopping after N layers still yields entry N,
# the output of layer N-1, unchanged. Everything above it exists only to
# feed `last_hidden_state`, which nothing consumes.
override = config.num_hidden_layers_override
self.num_layers = (config.num_hidden_layers
if override is None else min(config.num_hidden_layers, override))
self.output_hidden_state_index = config.output_hidden_state_index
self.layers = nn.ModuleList(
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
for index in range(self.num_layers))
# The final norm sits above the tapped layer, so a truncated stack drops
# it. Keeping it would overwrite the tapped entry with a normalised
# tensor and change conditioning without raising anything.
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
if self.num_layers == config.num_hidden_layers else None)
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
@@ -249,18 +243,14 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
inputs_embeds: torch.Tensor,
position_ids: torch.Tensor,
attention_mask: torch.Tensor | None,
output_hidden_states: bool,
visual_pos_masks: torch.Tensor | None,
deepstack_visual_embeds: list[torch.Tensor] | None,
) -> BaseEncoderOutput:
) -> torch.Tensor:
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
attention_mask = None
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
hidden_states = inputs_embeds
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
for layer_index, layer in enumerate(self.layers):
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
hidden_states = layer(hidden_states, position_embeddings, attention_mask)
if deepstack_visual_embeds is not None and layer_index < len(deepstack_visual_embeds):
if visual_pos_masks is None:
@@ -269,13 +259,9 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
visual = deepstack_visual_embeds[layer_index].to(hidden_states.device, hidden_states.dtype)
updated = hidden_states[mask].clone() + visual
hidden_states[mask] = updated
if self.norm is not None:
hidden_states = self.norm(hidden_states)
# Truncated or not, the last entry is appended here, so the tapped index
# lands in the same place either way.
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
if layer_index + 1 == self.output_hidden_state_index:
return hidden_states
raise RuntimeError(f"MiniMax-H3 text stack did not reach hidden_states[{self.output_hidden_state_index}]")
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
@@ -513,10 +499,18 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
return self.merger(hidden_states), deepstack_features
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
supports_hf_from_pretrained = False
supported_checkpoint_quantization_methods = frozenset({"fp8"})
@classmethod
def checkpoint_quantization_config_from_metadata(
cls,
metadata: dict[str, Any],
) -> MiniMaxH3SerializedFP8Config:
return MiniMaxH3SerializedFP8Config.from_config(metadata)
def __init__(self, config: MiniMaxH3Qwen3VLConfig) -> None:
super().__init__(config)
@@ -530,15 +524,12 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
@property
def num_hidden_layers(self) -> int:
"""The checkpoint architecture's nominal depth, matching its config.json.
When ``num_hidden_layers_override`` truncates the stack at the
conditioning tap, fewer layers exist; the built count is
``self.language_model.num_layers``, and the hidden-state tuple has
``num_layers + 1`` entries, not ``num_hidden_layers + 1``.
"""
return self.config.num_hidden_layers
@property
def num_built_hidden_layers(self) -> int:
return self.language_model.num_layers
def _get_rope_index(
self,
input_ids: torch.Tensor,
@@ -631,35 +622,39 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
f"tokens={int(mask.sum())}, features={features.shape[0]}")
return mask
def forward(
# no_grad, NOT inference_mode: with text_encoder_cpu_offload=True (the
# FastVideoArgs default) the loader FSDP2-shards this conditioner, and
# FSDP2's wait_for_unshard reads tensor._version via
# _unsafe_preserve_version_counter - inference tensors do not track
# version counters, so inference_mode crashes the first encode. no_grad
# frees the same activation memory and keeps prompt_embeds ordinary
# tensors (safe for any future backward through the conditioning).
@torch.no_grad()
def encode_ids(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
mm_token_type_ids: torch.Tensor | None = None,
**kwargs: Any,
) -> BaseEncoderOutput:
del mm_token_type_ids, kwargs
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("Exactly one of input_ids or inputs_embeds is required")
if inputs_embeds is None:
assert input_ids is not None
inputs_embeds = self.language_model.embed_tokens(input_ids)
if input_ids is None and (pixel_values is not None or pixel_values_videos is not None):
raise ValueError("Multimodal Qwen3-VL inputs require input_ids for placeholder matching")
) -> torch.Tensor:
if input_ids.ndim != 1:
raise ValueError(f"MiniMax-H3 slim forward expects 1-D input_ids, got shape={tuple(input_ids.shape)}")
if (pixel_values is None) != (image_grid_thw is None):
raise ValueError("pixel_values and image_grid_thw must be provided together")
if (pixel_values_videos is None) != (video_grid_thw is None):
raise ValueError("pixel_values_videos and video_grid_thw must be provided together")
input_ids = input_ids.unsqueeze(0)
inputs_embeds = self.language_model.embed_tokens(input_ids)
image_mask = None
video_mask = None
image_deepstack = None
video_deepstack = None
if pixel_values is not None:
if input_ids is None or image_grid_thw is None:
if image_grid_thw is None:
raise ValueError("pixel_values require input_ids and image_grid_thw")
image_features, image_deepstack = self._visual_features(pixel_values, image_grid_thw)
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -667,7 +662,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
"image")
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
if pixel_values_videos is not None:
if input_ids is None or video_grid_thw is None:
if video_grid_thw is None:
raise ValueError("pixel_values_videos require input_ids and video_grid_thw")
video_features, video_deepstack = self._visual_features(pixel_values_videos, video_grid_thw)
video_features = video_features.to(inputs_embeds.device, inputs_embeds.dtype)
@@ -695,50 +690,34 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
visual_mask = video_mask
deepstack_features = video_deepstack
if position_ids is None:
if input_ids is None:
sequence_length = inputs_embeds.shape[1]
position_ids = torch.arange(sequence_length,
device=inputs_embeds.device).view(1, 1,
-1).expand(3, inputs_embeds.shape[0], -1)
else:
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, attention_mask)
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
outputs = self.language_model(
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
hidden_states = self.language_model(
inputs_embeds,
position_ids,
attention_mask,
output_hidden_states,
None,
visual_mask,
deepstack_features,
)
outputs.attention_mask = attention_mask
return outputs
if hidden_states.ndim != 3 or hidden_states.shape[0] != 1:
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
return hidden_states[0]
def _is_above_the_tap(self, name: str) -> bool:
"""Whether this checkpoint key belongs to a layer we did not build.
A truncated language stack still ships every layer in the checkpoint, and
the unexpected-key check below is strict on purpose, so the surplus keys
have to be dropped here rather than by relaxing it.
"""
language_model = self.language_model
# The final norm is dropped exactly when the stack is truncated, so its
# absence is the signal.
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
if not index.isdigit():
return False
# Only drop indexes the full stack would have built. Anything at or
# above the checkpoint's own num_hidden_layers is corrupt and must
# still raise below, exactly as it does without truncation.
return language_model.num_layers <= int(index) < self.config.num_hidden_layers
def forward(
self,
input_ids: torch.Tensor,
*,
pixel_values: torch.Tensor | None = None,
image_grid_thw: torch.Tensor | None = None,
pixel_values_videos: torch.Tensor | None = None,
video_grid_thw: torch.Tensor | None = None,
) -> torch.Tensor:
return self.encode_ids(
input_ids,
pixel_values=pixel_values,
image_grid_thw=image_grid_thw,
pixel_values_videos=pixel_values_videos,
video_grid_thw=video_grid_thw,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
parameters = dict(self.named_parameters())
@@ -748,7 +727,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
if source_name == "lm_head.weight":
continue
name = source_name[6:] if source_name.startswith("model.") else source_name
if self._is_above_the_tap(name):
if self._is_omitted_checkpoint_key(name):
continue
if name not in parameters:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
@@ -758,7 +737,23 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
loaded.add(name)
return loaded
def _is_omitted_checkpoint_key(self, name: str) -> bool:
"""Return whether a valid checkpoint key belongs to an unbuilt layer."""
language_model = self.language_model
if language_model.norm is not None:
return False
if name == "language_model.norm.weight":
return True
prefix = "language_model.layers."
if not name.startswith(prefix):
return False
index = name[len(prefix):].split(".", 1)[0]
return (index.isdigit() and language_model.num_layers <= int(index) < self.config.num_hidden_layers)
EntryClass = MiniMaxH3Qwen3VLConditioner
__all__ = ["MiniMaxH3Qwen3VLConditioner"]
__all__ = [
"MiniMaxH3Qwen3VLConditioner",
"MiniMaxH3SerializedFP8Config",
]
+79 -17
View File
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from contextlib import nullcontext
from copy import deepcopy
from typing import cast
from typing import Any, cast
import torch
import torch.distributed as dist
@@ -30,9 +30,13 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.quantization import get_quantization_config
from fastvideo.logger import init_logger
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.hf_transformer_utils import get_diffusers_config
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
from fastvideo.models.loader.text_encoder_quantization import (
_configure_text_encoder_quantization,
_process_quantized_text_encoder_weights,
_resolve_text_encoder_checkpoint_path,
)
from fastvideo.models.loader.utils import set_default_torch_dtype
from fastvideo.models.loader.weight_utils import (
filter_duplicate_safetensors_files,
@@ -343,9 +347,23 @@ class TextEncoderLoader(ComponentLoader):
dtype: str = "fp16",
use_text_encoder_override: bool = False, # prevent subclasses from misusing
cpu_offload: bool | None = None,
offload_flag: str = "text_encoder_cpu_offload",
):
if cpu_offload is None:
cpu_offload = fastvideo_args.text_encoder_cpu_offload
runtime_device = get_local_torch_device()
device_id = runtime_device.index if runtime_device.index is not None else 0
requested_cpu_offload = getattr(fastvideo_args, offload_flag) if cpu_offload is None else cpu_offload
disable_cpu_offload = fastvideo_args.disable_offload_on_unified_memory(device_id,
offload_flag=offload_flag)
if requested_cpu_offload and disable_cpu_offload:
# Direct loader callers can choose a CPU target before the worker
# applies its device-local policy. Reset both the request and the
# target so the model is never constructed on the host first.
logger.info("Disabling %s on unified-memory device %d", offload_flag, device_id)
cpu_offload = False
target_device = runtime_device
else:
cpu_offload = requested_cpu_offload
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
from fastvideo.platforms import current_platform
@@ -353,16 +371,39 @@ class TextEncoderLoader(ComponentLoader):
if cpu_offload:
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
# Set quantization config if specified
if (use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None):
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
checkpoint_path = _resolve_text_encoder_checkpoint_path(
model_path,
fastvideo_args,
use_text_encoder_override,
)
checkpoint_quant_config = _configure_text_encoder_quantization(
model_config,
model_cls,
checkpoint_path,
)
if checkpoint_quant_config is not None:
if fastvideo_args.override_text_encoder_quant is not None:
raise ValueError("Serialized checkpoint quantization is selected from checkpoint metadata; "
"override_text_encoder_quant is an online conversion option and must be unset")
requested_dtype = PRECISION_TO_TYPE[dtype]
if requested_dtype not in checkpoint_quant_config.get_supported_act_dtypes():
raise ValueError(f"Serialized {checkpoint_quant_config.get_name()} text encoder does not support "
f"activation dtype {requested_dtype}")
checkpoint_quant_config.validate_runtime(runtime_device)
logger.info(
"Selected serialized %s text-encoder checkpoint execution from %s",
checkpoint_quant_config.get_name(),
checkpoint_path,
)
elif use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
if fastvideo_args.override_text_encoder_safetensors is None:
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
model_config.quant_config = quant_cls()
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
@@ -381,11 +422,20 @@ class TextEncoderLoader(ComponentLoader):
weights_to_load = {name for name, _ in model.named_parameters()}
if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None):
loaded_weights: set[str] = model.load_weights(
safetensors_weights_iterator(
[fastvideo_args.override_text_encoder_safetensors],
if os.path.isdir(checkpoint_path):
override_weights = self._get_all_weights(
model,
checkpoint_path,
to_cpu=bool(cpu_offload),
)
else:
if self.counter_before_loading_weights == 0.0:
self.counter_before_loading_weights = time.perf_counter()
override_weights = safetensors_weights_iterator(
[checkpoint_path],
to_cpu=use_cpu_offload,
)) # type: ignore
)
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
else:
loaded_weights: set[str] = model.load_weights(
self._get_all_weights(
@@ -400,6 +450,10 @@ class TextEncoderLoader(ComponentLoader):
self.counter_after_loading_weights - self.counter_before_loading_weights,
)
if checkpoint_quant_config is not None:
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
logger.info("Validated %d serialized blockwise FP8 text-encoder linears", processed_linears)
# Explicitly move model to target device after loading weights
model = model.to(target_device)
@@ -442,7 +496,7 @@ class TextEncoderLoader(ComponentLoader):
# that have loaded weights tracking currently.
# if loaded_weights is not None:
weights_not_loaded = weights_to_load - loaded_weights
if weights_not_loaded and model_config.quant_config is None:
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
@@ -514,6 +568,7 @@ class ImageEncoderLoader(TextEncoderLoader):
fastvideo_args,
encoder_precision,
cpu_offload=fastvideo_args.image_encoder_cpu_offload,
offload_flag="image_encoder_cpu_offload",
)
@@ -1057,7 +1112,12 @@ class TransformerLoader(ComponentLoader):
# so recording here makes the decision readable from the loaded
# transformer — and records the narrowed one for teacher/critic.
resolved = record_resolved_attention_backend(dit_config)
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
# Every worker records its resolved backend so distributed profile
# snapshots can prove that all ranks use the requested kernels.
logger.info("Worker %s transformer attention backend: %s",
os.environ.get("RANK", "0"),
resolved.name if resolved else "automatic selection",
local_main_process_only=False)
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
@@ -1080,6 +1140,8 @@ class TransformerLoader(ComponentLoader):
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
inference_regional_compile=fastvideo_args.inference_torch_compile,
inference_vsa_tile_size=fastvideo_args.VSA_tile_size,
)
total_params = sum(p.numel() for p in model.parameters())
+167 -5
View File
@@ -117,6 +117,23 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
torch.set_default_dtype(old_dtype)
def _prepare_model_for_compile(model: nn.Module, *, regional: bool) -> str | None:
"""Run a model compile hook, preferring its regional specialization."""
prepare = None
hook_name = "prepare_for_compile"
if regional:
prepare = getattr(model, "prepare_for_regional_compile", None)
hook_name = "prepare_for_regional_compile"
if not callable(prepare):
prepare = getattr(model, "prepare_for_compile", None)
hook_name = "prepare_for_compile"
if not callable(prepare):
return None
logger.info("Running %s for %s", hook_name, type(model).__name__)
unsupported = prepare()
return unsupported if isinstance(unsupported, str) and unsupported else None
# Supports optional torch.compile for FSDP-wrapped models during training
def maybe_load_fsdp_model(
model_cls: type[nn.Module],
@@ -136,6 +153,8 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
inference_regional_compile: bool = False,
inference_vsa_tile_size: int | None = None,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -231,13 +250,149 @@ def maybe_load_fsdp_model(
compile_in_loader = enable_torch_compile and training_mode
if compile_in_loader:
compile_kwargs = torch_compile_kwargs or {}
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
unsupported = _prepare_model_for_compile(model, regional=False)
if unsupported is not None:
logger.warning("Training torch.compile requested but disabled: %s. Model stays eager.", unsupported)
else:
compile_kwargs = torch_compile_kwargs or {}
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
elif inference_regional_compile and not training_mode:
# Inference-side counterpart of the #1718 training regional compile:
# per-block fullgraph compile right after the transformer loads, no
# user kwargs needed (fullgraph + emulate_precision_casts injected).
unsupported = _regional_compile_unsupported_reason(
init_params,
vsa_tile_size=inference_vsa_tile_size,
)
if unsupported is None:
unsupported = _prepare_model_for_compile(model, regional=True)
if unsupported is not None:
logger.warning(
"inference_torch_compile requested but disabled: %s. "
"Inference continues in eager mode.", unsupported)
else:
attention_count = _enable_regional_attention_compile(model)
logger.info("Enabled attention tracing for %d modules in %s", attention_count, type(model).__name__)
_compile_model_regions(model, torch_compile_kwargs or {})
return model
def _regional_compile_unsupported_reason(
init_params: dict[str, Any],
*,
vsa_tile_size: int | None = None,
) -> str | None:
"""Return why regional fullgraph compile cannot run, or None if it can.
Dense FA2, FA3, and FA4 inference all route through compile-visible
custom-op boundaries. FA3's raw autograd.Function carve-out applies only
to grad-enabled calls, outside this inference-only loader path.
The legacy VSA backend remains outside the fullgraph support envelope.
MiniMax H3's VSA backend is supported only through the inference-only
sm_100a tile-64 route; its regional hook resolves loaded compression
gates and probes the kernel before block capture.
"""
try:
from fastvideo.attention.layer import _attention_compile_explicitly_disabled
except Exception: # pragma: no cover - attention stack not importable
pass
else:
if _attention_compile_explicitly_disabled():
# The escape hatch wraps attention forwards in
# torch.compiler.disable, which is a hard dynamo error inside a
# fullgraph region ("Skip inlining `torch.compiler.disable()`d
# function"). Degrade to eager instead, matching the hatch's
# debugging intent.
return ("FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 keeps attention "
"forwards out of compiled graphs via torch.compiler."
"disable, which fullgraph regional compile cannot trace; "
"this model stays eager")
config = init_params.get("config")
resolved = getattr(config, "_resolved_attention_backend", None)
resolved_name = getattr(resolved, "name", "")
if resolved_name == "VIDEO_SPARSE_ATTN_H3":
if os.environ.get("FASTVIDEO_H3_VSA_PROBE"):
return ("FASTVIDEO_H3_VSA_PROBE records tensors and files from the VSA-H3 attention body, which "
"regional fullgraph compile cannot capture; this model stays eager")
if os.environ.get("FASTVIDEO_VSA_SM100A", "0") != "1":
return ("VIDEO_SPARSE_ATTN_H3 regional compile requires the compile-safe sm_100a route "
"(FASTVIDEO_VSA_SM100A=1); Triton/CuTe VSA stays eager")
if vsa_tile_size != 64:
return ("VIDEO_SPARSE_ATTN_H3 regional compile requires VSA_tile_size=64; "
f"got {vsa_tile_size!r}, so tile-256/CuTe VSA stays eager")
if resolved_name == "VIDEO_SPARSE_ATTN":
return (f"attention backend resolved to {resolved_name}, whose Triton "
"kernels, sequence-parallel collectives, and sync metadata "
"guard graph-break (incompatible with fullgraph regional "
"compile); this model stays eager")
return None
def _enable_regional_attention_compile(model: nn.Module) -> int:
"""Opt in distributed-attention instances owned by ``model`` only."""
from fastvideo.attention.layer import DistributedAttention
enabled_count = 0
for submodule in model.modules():
if isinstance(submodule, DistributedAttention):
submodule._set_compile_forward_enabled(True)
enabled_count += 1
return enabled_count
def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> int:
"""Compile repeated mathematical regions of a loaded model.
Only the selected module ``forward`` is replaced. This keeps activation
checkpoint wrappers structurally transparent while any module-level hooks
(FSDP pre/post, layerwise offload) execute outside the compiled region.
"""
compile_conditions = getattr(model, "_compile_conditions", None)
if not compile_conditions:
raise ValueError(f"{type(model).__name__} does not declare _compile_conditions")
if compile_kwargs.get("fullgraph", True) is not True:
raise ValueError("Regional compile requires fullgraph=True")
if "mode" in compile_kwargs:
# torch.compile forbids passing both `mode` and `options`, and
# regional compile always injects options (emulate_precision_casts)
# to match the training-side regional-compile configuration. Fail here
# with an actionable message instead of letting torch raise a
# mode/options conflict about an `options` key the user never wrote.
raise ValueError("Regional compile sets inductor options "
"(emulate_precision_casts) and cannot be combined "
"with torch_compile_kwargs['mode']. Remove 'mode' or "
"express its effect via torch_compile_kwargs['options'].")
kwargs = {**compile_kwargs, "fullgraph": True}
options = {"emulate_precision_casts": True}
options.update(kwargs.get("options") or {})
kwargs["options"] = options
compiled_count = 0
for name, submodule in list(model.named_modules()):
if not name:
continue
if any(condition(name, submodule) for condition in compile_conditions):
# Activation checkpoint wrappers are control-flow boundaries, not
# mathematical regions. Keep their saved-tensor/recompute logic
# eager and compile only the repeated block they own.
compile_target = getattr(submodule, "_checkpoint_wrapped_module", submodule)
compile_target.forward = torch.compile(compile_target.forward, **kwargs)
compiled_count += 1
if compiled_count == 0:
raise ValueError(f"No submodules in {type(model).__name__} matched _compile_conditions")
logger.info(
"Enabled regional torch.compile for %d submodules in %s with kwargs=%s",
compiled_count,
type(model).__name__,
kwargs,
)
return compiled_count
def shard_model(
model,
*,
@@ -390,7 +545,14 @@ def load_model_from_full_model_state_dict(
sharded_sd = {}
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(full_sd_iterator,
param_names_mapping) # type: ignore
for target_param_name, full_tensor in custom_param_sd.items():
# Drain rather than iterate. Production safetensors values may retain
# memory-mapped shard storage, while mapped or merged parameters can own
# ordinary allocations. Keeping the dict retains all of that source
# storage until loading finishes; popping releases each reference as soon
# as its conversion completes and lowers the host/unified-memory working
# set.
for target_param_name in list(custom_param_sd):
full_tensor = custom_param_sd.pop(target_param_name)
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
# Some checkpoints include extra entries that are not part of the
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
"""Checkpoint-serialized quantization lifecycle for native text encoders."""
import json
import os
from itertools import chain
from typing import Any
import torch
import torch.nn as nn
from safetensors.torch import safe_open
from fastvideo.configs.models import EncoderConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.layers.linear import LinearBase, UnquantizedLinearMethod
from fastvideo.layers.quantization.base_config import QuantizationConfig
from fastvideo.models.encoders.base import TextEncoder
def _resolve_text_encoder_checkpoint_path(
model_path: str,
fastvideo_args: FastVideoArgs,
use_text_encoder_override: bool,
) -> str:
override = fastvideo_args.override_text_encoder_safetensors if use_text_encoder_override else None
checkpoint_path = override or model_path
if not os.path.exists(checkpoint_path):
raise FileNotFoundError(f"Text-encoder checkpoint does not exist: {checkpoint_path}")
if not os.path.isdir(checkpoint_path) and not os.path.isfile(checkpoint_path):
raise ValueError(f"Text-encoder checkpoint must be a file or directory: {checkpoint_path}")
return checkpoint_path
def _read_text_encoder_checkpoint_quantization_config(checkpoint_path: str) -> dict[str, Any] | None:
checkpoint_dir = checkpoint_path if os.path.isdir(checkpoint_path) else os.path.dirname(checkpoint_path)
config_path = os.path.join(checkpoint_dir, "config.json")
if os.path.isfile(config_path):
try:
with open(config_path, encoding="utf-8") as config_file:
checkpoint_config = json.load(config_file)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid text-encoder checkpoint config: {config_path}") from error
quantization_config = checkpoint_config.get("quantization_config")
if quantization_config is not None:
if not isinstance(quantization_config, dict):
raise ValueError(f"quantization_config in {config_path} must be an object")
return quantization_config
if not os.path.isfile(checkpoint_path) or not checkpoint_path.endswith(".safetensors"):
return None
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint_file:
metadata = checkpoint_file.metadata() or {}
for key in ("quantization_config", "_quantization_metadata"):
serialized = metadata.get(key)
if serialized is None:
continue
try:
quantization_config = json.loads(serialized)
except json.JSONDecodeError as error:
raise ValueError(f"Invalid {key} metadata in {checkpoint_path}") from error
if not isinstance(quantization_config, dict):
raise ValueError(f"{key} metadata in {checkpoint_path} must decode to an object")
return quantization_config
return None
def _configure_text_encoder_quantization(
model_config: EncoderConfig,
model_cls: type[nn.Module],
checkpoint_path: str,
) -> QuantizationConfig | None:
if not issubclass(model_cls, TextEncoder):
return None
checkpoint_quantization = _read_text_encoder_checkpoint_quantization_config(checkpoint_path)
if checkpoint_quantization is None:
return None
quant_method = str(checkpoint_quantization.get("quant_method", "")).lower()
if not quant_method:
raise ValueError(f"Quantized text-encoder checkpoint {checkpoint_path} does not declare quant_method")
supported_methods = getattr(model_cls, "supported_checkpoint_quantization_methods", frozenset())
if quant_method not in supported_methods:
supported = ", ".join(sorted(supported_methods)) or "none"
raise ValueError(f"Text encoder {model_cls.__name__} does not support serialized {quant_method!r} "
f"checkpoints (supported: {supported})")
factory = getattr(model_cls, "checkpoint_quantization_config_from_metadata", None)
if not callable(factory):
raise ValueError(f"Text encoder {model_cls.__name__} advertises serialized {quant_method!r} support "
"without a checkpoint quantization factory")
quant_config = factory(checkpoint_quantization)
model_config.quant_config = quant_config
return quant_config
def _module_tensor_device(module: nn.Module) -> torch.device | None:
devices = {
tensor.device
for tensor in chain(
module.parameters(recurse=False),
module.buffers(recurse=False),
)
}
if len(devices) > 1:
raise ValueError(f"Quantized text-encoder module {type(module).__name__} spans multiple devices: {devices}")
return next(iter(devices), None)
def _process_quantized_text_encoder_weights(model: nn.Module, process_device: torch.device) -> int:
"""Run quantized post-load hooks one linear at a time on ``process_device``."""
processed = 0
for module in model.modules():
if not isinstance(module, LinearBase) or isinstance(module.quant_method, UnquantizedLinearMethod):
continue
if module.quant_method is None:
continue
original_device = _module_tensor_device(module)
try:
module.to(process_device)
module.quant_method.process_weights_after_loading(module)
finally:
if original_device is not None:
module.to(original_device)
processed += 1
if processed == 0:
raise ValueError("Serialized quantized text-encoder checkpoint selected, but no quantized linear layers exist")
return processed
@@ -392,9 +392,16 @@ class MiniMaxH3AudioBigVGANDecoder(nn.Module):
return torch.clamp(hidden_states, min=-1.0, max=1.0)
def _is_minimax_h3_audio_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the audio decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3AudioBigVGANDecoder)
class MiniMaxH3AudioVAE(nn.Module):
"""DAC encoder plus BigVGAN decoder for mono 32 kHz waveforms."""
_compile_conditions = [_is_minimax_h3_audio_vae_decoder]
def __init__(self, config: MiniMaxH3AudioVAEConfig):
super().__init__()
self.config = config
@@ -0,0 +1,381 @@
# SPDX-License-Identifier: Apache-2.0
"""Sequence-parallel chunk scheduling for the MiniMax-H3 video VAE.
The H3 video VAE decodes a video as a series of temporal-chunk decoder
forwards whose outputs are joined by a short deterministic frame blend
(``AutoencoderKLMiniMaxH3._decode_chunks``), and encodes videos as fully
independent ``clip_length``-frame encoder forwards. Neither the chunk decode
nor the clip encode has any cross-chunk data dependency — only the *joining*
of decoded chunks (overlap blending, frame trimming) is sequential. This
module round-robins the chunk/clip forwards across the ranks of a
sequence-parallel group and replays the serial joining logic on the
assembling rank, reproducing the serial result bit for bit.
Bit-exactness contract:
- every rank holds an identical copy of the inputs (the H3 DiT all-gathers
its outputs, and reference pixels are prepared identically on all ranks);
- a chunk decoded on any rank is bitwise the tensor the serial loop would
produce (identical weights, inputs, and deterministic kernels on identical
GPUs), and NCCL transports it bitwise;
- every serialization point of the serial algorithm (overlap blending, frame
trimming, pixel denormalization, output-buffer copies, moment
concatenation and token-drop trimming) runs on the assembling rank in
serial order via the same VAE methods the serial path uses.
Collective safety: all group ranks must call these functions together with
identically shaped inputs. Work proceeds in rounds of one collective each;
ranks without a chunk in the final round contribute a placeholder tensor, so
participation is uniform by construction and no rank-dependent branch guards
a collective.
Caveat — compiled decoders (``enable_torch_compile_vae``): inductor autotunes
kernel configs per process at first call, so a compiled decoder is only
deterministic WITHIN a process, not across processes. Chunks decoded on other
ranks then differ from the serial rank's decode of the same chunk exactly as
two serial runs in different processes would. Direct decoder tensors measured
on GB200 at 124f had max absolute error 0.00268358 (0.684/255), mean absolute
error 4.213e-05 (0.0107/255), and 24.59% nonzero values; the first chunk was
bit-identical. A separate decoded-MP4 comparison reached 63/255 on <0.5% of
pixels, but that includes lossy MP4 encoding and is not the decoder-tensor
error envelope. With the eager decoder — the pipeline default — parallel
output is bitwise equal to serial ``decode_to_pixels``.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
from fastvideo.models.vaes.minimax_h3_video import (
AutoencoderKLMiniMaxH3,
AutoencoderKLOutput,
DiagonalGaussianDistribution,
)
from fastvideo.profiler import nvtx_range
if TYPE_CHECKING:
from fastvideo.distributed.parallel_state import GroupCoordinator
# Collective used to move decoded chunk segments to the assembling rank.
# "gather" moves each segment once (destination-only); "all_gather" also
# leaves every rank with every segment. Both are exact; the default is the
# faster one measured on GB200 NVL72 (see the PR notes).
DECODE_GATHER_STRATEGIES = ("gather", "all_gather")
DEFAULT_DECODE_GATHER_STRATEGY = "gather"
def parallel_chunk_indices(num_chunks: int, world_size: int, rank_in_group: int) -> list[int]:
"""Round-robin chunk ownership: chunk ``i`` belongs to rank ``i % world_size``."""
if num_chunks < 0:
raise ValueError(f"num_chunks must be non-negative, got {num_chunks}.")
if world_size < 1:
raise ValueError(f"world_size must be positive, got {world_size}.")
if not 0 <= rank_in_group < world_size:
raise ValueError(f"rank_in_group {rank_in_group} out of range for world_size {world_size}.")
return list(range(rank_in_group, num_chunks, world_size))
def _num_rounds(num_chunks: int, world_size: int) -> int:
return -(-num_chunks // world_size)
def _decode_segment(vae: AutoencoderKLMiniMaxH3, z_padded: torch.Tensor, chunk_index: int) -> torch.Tensor:
"""Decode one temporal chunk's clip and keep the frames the join consumes.
The serial loop uses two spans of each decoded clip: the chunk body
``clip[:, :, frame_pre_padding:chunk_num_frames]`` and (when
``token_drop > 0``) the blend tail
``clip[:, :, chunk_num_frames + frame_pre_padding:]``. Everything from
``frame_pre_padding`` on covers both, so one contiguous slice per chunk
travels over the wire. ``.contiguous()`` also detaches the segment from
any decoder-owned storage (e.g. a compiled decoder's reuse pools) before
the next chunk decode can overwrite it.
"""
start = chunk_index * vae.tokens_chunk_size
with nvtx_range(f"minimax_h3.vae.parallel_chunk.{chunk_index}"):
clip = vae._decode_clip(z_padded[:, :, start:start + vae.tokens_chunk_size + vae.token_overlap])
return clip[:, :, vae.frame_pre_padding:].contiguous()
class _ChunkAssembler:
"""Replay the serial chunk-joining semantics of ``_decode_chunks`` +
``_decode_to_pixels`` on gathered chunk segments, in chunk order.
On CUDA the joining kernels and output copies run on a dedicated side
stream: they depend only on already-gathered segments, so running them
off the main stream keeps the assembling rank's next chunk decode (and
therefore every other rank's next collective) off the assembly's tail.
Stream placement cannot change values — the ops and their order are
identical — so bit-exactness with the serial path is unaffected.
"""
def __init__(self, vae: AutoencoderKLMiniMaxH3, output: torch.Tensor, output_num_frames: int,
non_blocking: bool, device: torch.device) -> None:
self._vae = vae
self._output = output
self._output_num_frames = output_num_frames
self._non_blocking = non_blocking
self._body_frames = vae.tokens_chunk_size * vae.temporal_compression_ratio - vae.frame_pre_padding
self._overlap: torch.Tensor | None = None
self._frame_start = 0
self._stream = torch.cuda.Stream(device) if device.type == "cuda" else None
def push(self, segment: torch.Tensor) -> None:
"""Consume the next chunk's segment (``clip[:, :, frame_pre_padding:]``)."""
if self._stream is None:
self._push(segment)
return
# The segment is produced on the current (collective) stream; hand it
# to the assembly stream and pin its storage until assembly reads it.
self._stream.wait_stream(torch.cuda.current_stream(segment.device))
segment.record_stream(self._stream)
with torch.cuda.stream(self._stream):
self._push(segment)
def _push(self, segment: torch.Tensor) -> None:
vae = self._vae
chunk = segment[:, :, :self._body_frames]
if self._overlap is not None:
chunk = vae._blend(self._overlap, chunk, vae.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], self._output_num_frames - self._frame_start)
chunk = chunk[:, :, :num_frames]
# The tail past the body (and its pre-padding gap) is the next
# chunk's blend overlap — the serial loop's ``next_overlap``.
self._overlap = segment[:, :, self._body_frames + vae.frame_pre_padding:] if vae.config.token_drop > 0 else None
if num_frames > 0:
self._emit(chunk)
def finalize(self) -> None:
"""Emit the final overlap tail exactly as the serial generator does."""
if self._overlap is not None and self._frame_start < self._output_num_frames:
tail = self._overlap[:, :, :self._output_num_frames - self._frame_start]
if self._stream is None:
self._emit(tail)
else:
with torch.cuda.stream(self._stream):
self._emit(tail)
if self._frame_start != self._output.shape[2]:
raise RuntimeError(
f"MiniMax-H3 decode wrote {self._frame_start} frames into an output buffer expecting "
f"{self._output.shape[2]}.")
def synchronize(self) -> None:
"""Drain assembly kernels and output copies before the buffer is read."""
if self._stream is not None:
self._stream.synchronize()
def _emit(self, chunk: torch.Tensor) -> None:
pixels = self._vae.denormalize_pixels(chunk.float()).clamp_(0, 1)
self._vae._copy_chunk_pixels(pixels, self._output, self._frame_start, self._non_blocking)
self._frame_start += pixels.shape[2]
def _broadcast_segment_meta(group: "GroupCoordinator",
segment: torch.Tensor | None) -> tuple[torch.dtype, tuple[int, ...]]:
"""Share the leader's real segment dtype/shape so placeholder tensors match.
The decoder's output dtype depends on the surrounding autocast context;
deriving it on the leader from an actually decoded segment (instead of
predicting it) keeps collective dtypes correct by construction.
"""
meta = (segment.dtype, tuple(segment.shape)) if segment is not None else None
meta = group.broadcast_object(meta, src=0)
if meta is None:
raise RuntimeError("MiniMax-H3 parallel VAE meta broadcast returned no leader metadata.")
return meta
def decode_to_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str = DEFAULT_DECODE_GATHER_STRATEGY,
) -> torch.Tensor | None:
"""Chunk-parallel ``decode_to_pixels`` across a sequence-parallel group.
All group ranks call this together with identical ``z``. Temporal chunks
are decoded round-robin across the group and their segments move to the
group's first rank, which assembles bitwise the serial
``decode_to_pixels`` result into ``output``. Only the first rank passes
``output`` (validated exactly like the serial API); other ranks pass
``None`` and receive ``None``.
"""
if strategy not in DECODE_GATHER_STRATEGIES:
raise ValueError(f"Unknown parallel-decode strategy {strategy!r}; expected one of {DECODE_GATHER_STRATEGIES}.")
is_leader = group.rank_in_group == 0
if is_leader:
if output is None:
raise ValueError("The first sequence-parallel rank must provide the CPU output buffer.")
expected_shape = vae.decoded_pixel_shape(z.shape)
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
raise ValueError(
"`output` must be a CPU float32 tensor with shape "
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
elif output is not None:
raise ValueError("Only the first sequence-parallel rank may provide an output buffer.")
if group.world_size == 1:
return vae.decode_to_pixels(z, output)
try:
if vae.use_slicing and z.shape[0] > 1:
for batch_index, z_slice in enumerate(z.split(1)):
slice_output = output[batch_index:batch_index + 1] if output is not None else None
_decode_single_parallel(vae, z_slice, slice_output, group, strategy)
else:
_decode_single_parallel(vae, z, output, group, strategy)
finally:
# Drain the leader's async chunk copies before the caller (or an
# exception handler) can read or release the pinned buffer.
if output is not None and vae._streams_chunk_copies(z, output):
torch.cuda.current_stream(z.device).synchronize()
return output
def _decode_single_parallel(
vae: AutoencoderKLMiniMaxH3,
z: torch.Tensor,
output: torch.Tensor | None,
group: "GroupCoordinator",
strategy: str,
) -> None:
pad_tokens, num_chunks, output_num_frames = vae._temporal_decode_plan(z.shape[2])
if pad_tokens > 0:
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
world_size = group.world_size
rank = group.rank_in_group
# Every rank decodes its round-0 chunk BEFORE the metadata rendezvous so
# the first decodes run concurrently (a rank that waited on the broadcast
# first would idle a full chunk-decode behind the leader). The leader
# owns chunk 0 under round-robin assignment, so its segment supplies real
# dtype/shape for placeholder rounds instead of guessing autocast state.
first_segment = _decode_segment(vae, z, rank) if rank < num_chunks else None
segment_dtype, segment_shape = _broadcast_segment_meta(group, first_segment if rank == 0 else None)
assembler = None
if output is not None:
non_blocking = vae._streams_chunk_copies(z, output)
assembler = _ChunkAssembler(vae, output, output_num_frames, non_blocking, z.device)
try:
segment_frames = segment_shape[2]
for round_index in range(_num_rounds(num_chunks, world_size)):
chunk_index = round_index * world_size + rank
if chunk_index >= num_chunks:
segment = torch.zeros(segment_shape, dtype=segment_dtype, device=z.device)
elif round_index == 0 and first_segment is not None:
segment = first_segment
else:
segment = _decode_segment(vae, z, chunk_index)
with nvtx_range(f"minimax_h3.vae.parallel_{strategy}.{round_index}"):
if strategy == "gather":
gathered = group.gather(segment, dst=0, dim=2)
else:
gathered = group.all_gather(segment, dim=2)
if assembler is None or gathered is None:
continue
for slot in range(world_size):
if round_index * world_size + slot >= num_chunks:
break
assembler.push(gathered.narrow(2, slot * segment_frames, segment_frames))
if assembler is not None:
assembler.finalize()
finally:
# Drain assembly-stream copies into ``output`` even on the error path
# so an exception cannot leave an in-flight DMA into a buffer the
# caller may release.
if assembler is not None:
assembler.synchronize()
def _encode_clip_moments(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor, clip_index: int) -> torch.Tensor:
"""Encode one ``clip_length``-frame clip exactly as ``_encode_pixels`` does."""
clip_length = vae.config.clip_length
frame_start = clip_index * clip_length
with nvtx_range(f"minimax_h3.vae.parallel_encode_clip.{clip_index}"):
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
device=vae.pixel_mean.device,
dtype=torch.float32,
)
if pixels.dtype == torch.uint8:
clip = clip / 255.0
if clip.shape[2] < clip_length:
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
clip = torch.cat([clip, pad_frames], dim=2)
clip = vae.normalize_pixels(clip)
return vae._encode_clip(clip).contiguous()
def encode_pixels_parallel(
vae: AutoencoderKLMiniMaxH3,
pixels: torch.Tensor,
group: "GroupCoordinator",
) -> AutoencoderKLOutput:
"""Clip-parallel ``encode_pixels`` across a sequence-parallel group.
Encoder clips have no cross-clip dependency (no overlap, no blending), so
ranks encode disjoint clips and all-gather the per-clip moment tensors.
Every rank returns the identical full posterior — preserving the serial
contract that all ranks hold the same encoded latents — bitwise equal to
``vae.encode_pixels(pixels)``. Moments are latent-sized (a few MB per
clip), so the all-gather is negligible next to the clip forwards.
"""
if pixels.ndim != 5 or pixels.shape[1] != vae.config.in_channels or pixels.shape[2] <= 0:
raise ValueError(
f"`pixels` must have shape [B, {vae.config.in_channels}, T, H, W] with T > 0, "
f"got {tuple(pixels.shape)}.")
if pixels.device.type != "cpu":
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
if group.world_size == 1:
return vae.encode_pixels(pixels)
if vae.use_slicing and pixels.shape[0] > 1:
moments = torch.cat([_encode_single_parallel(vae, pixel_slice, group) for pixel_slice in pixels.split(1)])
else:
moments = _encode_single_parallel(vae, pixels, group)
return AutoencoderKLOutput(latent_dist=DiagonalGaussianDistribution(moments))
def _encode_single_parallel(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor,
group: "GroupCoordinator") -> torch.Tensor:
clip_length = vae.config.clip_length
num_clips = -(-pixels.shape[2] // clip_length)
world_size = group.world_size
rank = group.rank_in_group
# Same first-work-then-rendezvous ordering as the decode path: encode the
# round-0 clip before the metadata broadcast so first encodes overlap.
first_moments = _encode_clip_moments(vae, pixels, rank) if rank < num_clips else None
moment_dtype, moment_shape = _broadcast_segment_meta(group, first_moments if rank == 0 else None)
moment_tokens = moment_shape[2]
parts: list[torch.Tensor] = []
for round_index in range(_num_rounds(num_clips, world_size)):
clip_index = round_index * world_size + rank
if clip_index >= num_clips:
moments = torch.zeros(moment_shape, dtype=moment_dtype, device=vae.pixel_mean.device)
elif round_index == 0 and first_moments is not None:
moments = first_moments
else:
moments = _encode_clip_moments(vae, pixels, clip_index)
gathered = group.all_gather(moments, dim=2)
for slot in range(world_size):
if round_index * world_size + slot >= num_clips:
break
parts.append(gathered.narrow(2, slot * moment_tokens, moment_tokens))
encoded = torch.cat(parts, dim=2)
if vae.config.token_drop > 0:
encoded = encoded[:, :, :-vae.config.token_drop]
return encoded
__all__ = [
"DECODE_GATHER_STRATEGIES",
"DEFAULT_DECODE_GATHER_STRATEGY",
"decode_to_pixels_parallel",
"encode_pixels_parallel",
"parallel_chunk_indices",
]
+192 -58
View File
@@ -15,7 +15,10 @@ import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
from fastvideo.attention import get_attn_backend
from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.profiler import nvtx_range
class DiagonalGaussianDistribution:
@@ -292,6 +295,7 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
class MiniMaxH3VideoAttention(nn.Module):
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
"""Build projections and the selected dense FastVideo attention implementation."""
super().__init__()
self.heads = heads
self.dim_head = dim_head
@@ -303,12 +307,38 @@ class MiniMaxH3VideoAttention(nn.Module):
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
self.attn_impl = None
from fastvideo.platforms import current_platform
if current_platform.is_cuda_alike():
attention_backend = get_attn_backend(
dim_head,
# FlashAttention executes the FP32 VAE activations in BF16 and
# restores FP32 output, so resolve against the kernel dtype.
torch.bfloat16,
supported_attention_backends=(
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
),
)
self.attn_impl = attention_backend.get_impl_cls()(
num_heads=heads,
head_size=dim_head,
softmax_scale=dim_head**-0.5,
num_kv_heads=heads,
causal=False,
# The FASTVIDEO_NVFP4_FA4 env opt-in targets the DiT; this VAE
# is FP32-pinned (_keep_in_fp32_modules), so force-disable FP4
# Q/K quantization for its attention regardless of the env.
nvfp4_fa4=False,
)
def forward(
self,
hidden_states: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""Apply dense self-attention to one spatial VAE token sequence."""
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
@@ -329,9 +359,17 @@ class MiniMaxH3VideoAttention(nn.Module):
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
if self.attn_impl is not None and query.device.type != "cpu":
# VAE decoding has no diffusion-step metadata, so call the selected
# backend implementation directly with dense BSHD tensors.
hidden_states = self.attn_impl.forward(query, key, value, None)
hidden_states = hidden_states.flatten(2, 3)
else:
# Keep CPU construction and execution available without requiring
# an accelerator attention backend.
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
hidden_states = F.scaled_dot_product_attention(query, key, value)
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
return self.to_out[0](hidden_states)
@@ -434,6 +472,7 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Decode one latent spatial input through the H3 video transformer."""
batch_size, num_channels, num_frames, height, width = hidden_states.shape
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).reshape(
batch_size,
@@ -483,6 +522,11 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
)
def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool:
"""Select the video decoder that serves the H3 VAE ``decode`` path."""
return name == "decoder" and isinstance(submodule, MiniMaxH3VideoViTDecoder3d)
class AutoencoderKLMiniMaxH3(nn.Module):
"""MiniMax-H3 causal encoder and ViT decoder with exact release geometry."""
@@ -490,6 +534,10 @@ class AutoencoderKLMiniMaxH3(nn.Module):
_no_split_modules = ["MiniMaxH3VideoResnetBlock3d", "MiniMaxH3VideoTransformerBlock"]
_repeated_blocks = ["MiniMaxH3VideoTransformerBlock"]
_keep_in_fp32_modules = ["encoder", "decoder", "quant_conv", "post_quant_conv"]
_compile_conditions = [_is_minimax_h3_video_vae_decoder]
# ``prepare_for_compile`` flips this per instance when the pipeline opts
# into ``enable_torch_compile_vae``; default instances stay fully eager.
_tile_helpers_compiled = False
def __init__(self, config: MiniMaxH3VideoVAEConfig) -> None:
super().__init__()
@@ -655,12 +703,40 @@ class AutoencoderKLMiniMaxH3(nn.Module):
slice_rest[dim] = slice(blend_extent, None)
return torch.cat([blended, b[tuple(slice_rest)]], dim=dim)
def prepare_for_compile(self) -> None:
"""Compile the fixed-shape tile helpers for the opt-in VAE compile path.
``ComposedPipelineBase._maybe_compile_pipeline_module`` calls this hook
only when ``enable_torch_compile_vae`` is set, right before the decoder
is compiled through ``_compile_conditions``. The spatial tile grid and
the per-tile decoder-input projection have fixed shapes, so
``mode="reduce-overhead"`` records one CUDA graph per geometry and
replays it for every tile and temporal chunk. Keeping this behind the
opt-in means default (eager) users pay neither the inductor/triton
toolchain requirement and first-decode compile latency nor the
permanent cudagraph memory pools, and multi-resolution callers never
churn ``dynamic=False`` recompiles they did not ask for.
"""
if self._tile_helpers_compiled:
return
# The fixed spatial tile grid reuses one compiled blend-and-concatenate graph.
self._stitch_tiles = torch.compile(self._stitch_tiles, backend="inductor", mode="reduce-overhead", dynamic=False)
# Each fixed-shape latent tile reuses one compiled decoder-input projection.
self._project_decoder_tile = torch.compile(
self._project_decoder_tile,
backend="inductor",
mode="reduce-overhead",
dynamic=False,
)
self._tile_helpers_compiled = True
def _stitch_tiles(
self,
tiles: list[list[torch.Tensor]],
height_overlaps: list[int],
width_overlaps: list[int],
) -> torch.Tensor:
"""Blend decoded tile overlaps and concatenate the spatial canvas."""
result_rows = []
for row_index, row in enumerate(tiles):
result_row = []
@@ -677,6 +753,10 @@ class AutoencoderKLMiniMaxH3(nn.Module):
result_rows.append(torch.cat(result_row, dim=-1))
return torch.cat(result_rows, dim=-2)
def _project_decoder_tile(self, tile: torch.Tensor) -> torch.Tensor:
"""Project one spatial latent tile into the decoder input channels."""
return self.post_quant_conv(tile)
def _encode_clip(self, x: torch.Tensor) -> torch.Tensor:
if not self.use_tiling:
return self.quant_conv(self.encoder(x))
@@ -700,36 +780,73 @@ class AutoencoderKLMiniMaxH3(nn.Module):
rows.append(row)
latent_y_overlaps = [overlap // self.spatial_compression_ratio for overlap in y_overlaps]
latent_x_overlaps = [overlap // self.spatial_compression_ratio for overlap in x_overlaps]
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps)
stitched = self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps)
if self._tile_helpers_compiled:
# Under the opt-in mode="reduce-overhead" compile the stitched
# canvas is a CUDA-graph static buffer that the next _stitch_tiles
# replay overwrites. Callers (_encode/_encode_pixels/
# encode_keyframe) collect per-clip results across replays before
# concatenating, so hand them a caller-owned tensor instead of
# cudagraph-pooled storage. Eager instances return the fresh
# torch.cat result directly.
stitched = stitched.clone()
return stitched
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
if not self.use_tiling:
return self.decoder(self.post_quant_conv(z))
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
for y_position, y_length in zip(y_indices, y_lengths):
row = []
for x_position, x_length in zip(x_indices, x_lengths):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
row.append(self.decoder(self.post_quant_conv(tile)))
rows.append(row)
return self._stitch_tiles(rows, y_overlaps, x_overlaps)
"""Decode one temporal clip, with optional overlapping spatial tiles."""
with nvtx_range("minimax_h3.vae.decode_clip"):
if not self.use_tiling:
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"):
projected_clip = self.post_quant_conv(z)
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"):
return self.decoder(projected_clip)
height = z.shape[-2] * self.spatial_compression_ratio
width = z.shape[-1] * self.spatial_compression_ratio
with nvtx_range("minimax_h3.vae.decode_clip.split_tiles"):
y_indices, y_lengths, y_overlaps = self._split_tiles(
height,
self.tile_sample_min_height,
self.tile_sample_min_overlap_height,
)
x_indices, x_lengths, x_overlaps = self._split_tiles(
width,
self.tile_sample_min_width,
self.tile_sample_min_overlap_width,
)
ratio = self.spatial_compression_ratio
rows = []
# The eager tile driver owns NVTX so each marker remains outside
# the compiled decoder graph.
with nvtx_range("minimax_h3.vae.decode_clip.decode_tiles"):
for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths)):
row = []
for column_index, (x_position, x_length) in enumerate(zip(x_indices, x_lengths)):
with nvtx_range(f"minimax_h3.vae.decode_clip.tile.{row_index}.{column_index}"):
tile = z[
...,
y_position // ratio:y_position // ratio + y_length // ratio,
x_position // ratio:x_position // ratio + x_length // ratio,
]
projected_tile = self._project_decoder_tile(tile)
with nvtx_range("minimax_h3.vae.decode_clip.tile.decoder_forward"):
decoded_tile = self.decoder(projected_tile)
row.append(decoded_tile)
rows.append(row)
with nvtx_range("minimax_h3.vae.decode_clip.stitch_tiles"):
stitched = self._stitch_tiles(rows, y_overlaps, x_overlaps)
if self._tile_helpers_compiled:
# Same CUDA-graph output-ownership contract as
# _encode_clip: _decode collects chunks across
# _stitch_tiles replays before torch.cat, so the pooled
# canvas must not escape this driver. (The streaming
# _decode_to_pixels path copies each chunk out before the
# next decode; the clone keeps it correct too at one D2D
# copy per chunk.)
stitched = stitched.clone()
return stitched
def _encode(self, x: torch.Tensor) -> torch.Tensor:
clip_length = self.config.clip_length
@@ -809,21 +926,27 @@ class AutoencoderKLMiniMaxH3(nn.Module):
output_frame_start = 0
overlap = None
for index in range(num_chunks):
start = index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
next_overlap = None
if self.config.token_drop > 0:
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
if overlap is not None:
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
for chunk_index in range(num_chunks):
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}"):
start = chunk_index * tokens_chunk_size
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.0"):
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
if overlap is not None:
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
chunk = chunk[:, :, :num_frames]
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
if num_frames > 0:
yield chunk[:, :, :num_frames]
output_frame_start += num_frames
next_overlap = None
if self.config.token_drop > 0:
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.1"):
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
# Yield after the ranges close so consumer-side CPU copies do not inflate decoder timing.
overlap = next_overlap
if num_frames > 0:
output_frame_start += num_frames
yield chunk
if overlap is not None and output_frame_start < output_num_frames:
yield overlap[:, :, :output_num_frames - output_frame_start]
@@ -849,31 +972,42 @@ class AutoencoderKLMiniMaxH3(nn.Module):
"""Whether finalized chunks copy to ``output`` asynchronously on the current CUDA stream."""
return z.device.type == "cuda" and output.is_pinned()
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
@staticmethod
def _copy_chunk_pixels(pixels: torch.Tensor, output: torch.Tensor, frame_start: int, non_blocking: bool) -> None:
"""Copy one finalized fp32 pixel chunk into the CPU ``output`` buffer.
Device-to-host copies run per (batch, channel) plane: the temporal
slice of ``output`` is strided across channels, but each plane is
contiguous on both sides, so every transfer stays a direct memcpy
instead of staging through a pageable CPU temporary. With a pinned
``output`` the copies are additionally asynchronous and overlap the
next chunk's decode; ``decode_to_pixels`` synchronizes once before
returning.
``output`` and ``non_blocking=True`` the copies are additionally
asynchronous on the current CUDA stream; callers synchronize once
before releasing the buffer.
"""
target = output[:, :, frame_start:frame_start + pixels.shape[2]]
if pixels.device.type == "cuda":
pixels = pixels.contiguous()
for batch_index in range(pixels.shape[0]):
for channel_index in range(pixels.shape[1]):
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
non_blocking=non_blocking)
else:
target.copy_(pixels)
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
Each finalized chunk streams through ``_copy_chunk_pixels`` (direct
per-plane memcpys; asynchronous with a pinned ``output``) so the
copies overlap the next chunk's decode; ``decode_to_pixels``
synchronizes once before returning.
"""
non_blocking = self._streams_chunk_copies(z, output)
output_frame_start = 0
for chunk in self._decode_chunks(z):
num_frames = chunk.shape[2]
pixels = self.denormalize_pixels(chunk.float()).clamp_(0, 1)
target = output[:, :, output_frame_start:output_frame_start + num_frames]
if z.device.type == "cuda":
pixels = pixels.contiguous()
for batch_index in range(pixels.shape[0]):
for channel_index in range(pixels.shape[1]):
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
non_blocking=non_blocking)
else:
target.copy_(pixels)
self._copy_chunk_pixels(pixels, output, output_frame_start, non_blocking)
output_frame_start += num_frames
if output_frame_start != output.shape[2]:
raise RuntimeError(
@@ -12,9 +12,9 @@ from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_IMAGE_PAD_TOKEN,
MINIMAX_H3_TEXT_ENCODER_LAYER,
MINIMAX_H3_TEXT_TAG,
MINIMAX_H3_VIDEO_PAD_TOKEN,
MINIMAX_H3_VIDEO_TAG,
@@ -42,25 +42,6 @@ def _token_ids(tokenized: Any) -> list[int]:
return [int(token_id) for token_id in input_ids]
def _create_mm_token_type_ids(processor: Any, token_ids: list[int]) -> list[list[int]]:
"""Build Qwen3-VL modality IDs across old and new Transformers releases."""
create_ids = getattr(processor, "create_mm_token_type_ids", None)
if callable(create_ids):
return create_ids([token_ids])
modality_ids = [0] * len(token_ids)
for modality, modality_type in (("image", 1), ("video", 2), ("audio", 3)):
special_ids = getattr(processor, f"{modality}_token_ids", None)
if special_ids is None:
special_id = getattr(processor, f"{modality}_token_id", None)
special_ids = [] if special_id is None else [special_id]
resolved_ids = {int(special_id) for special_id in special_ids if special_id is not None}
for index, token_id in enumerate(token_ids):
if token_id in resolved_ids:
modality_ids[index] = modality_type
return [modality_ids]
def build_ref2va_presentation(
tokenizer: Any,
prompt: str,
@@ -155,20 +136,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
device: torch.device,
**vision_inputs: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
hidden_state_index = MINIMAX_H3_TEXT_ENCODER_LAYER
input_ids = torch.tensor([token_ids], dtype=torch.long, device=device)
mm_token_type_ids = torch.as_tensor(
_create_mm_token_type_ids(self.processor, token_ids),
dtype=torch.long,
device=device,
)
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
dtype = self.conditioner.dtype
outputs = self.conditioner(
input_ids=input_ids,
attention_mask=torch.ones_like(input_ids),
mm_token_type_ids=mm_token_type_ids,
use_cache=False,
output_hidden_states=True,
prompt_embeds = self.conditioner(
input_ids,
**{
name:
None if value is None else value.to(
@@ -178,10 +149,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
for name, value in vision_inputs.items()
},
)
if outputs.hidden_states is None or len(outputs.hidden_states) <= hidden_state_index:
raise ValueError(f"Qwen3-VL did not return `hidden_states[{hidden_state_index}]`.")
if prompt_embeds.ndim != 2 or prompt_embeds.shape[0] != len(token_ids):
raise ValueError(f"MiniMax-H3 slim text encoder returned unexpected shape={tuple(prompt_embeds.shape)}")
return (
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
torch.tensor(token_tags, dtype=torch.long),
)
@@ -286,6 +257,7 @@ class MiniMaxH3ConditioningStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Encode one H3 prompt presentation and attach its packed text features."""
device = get_local_torch_device()
first_param = next(self.conditioner.parameters(), None)
moved_for_forward = (fastvideo_args.text_encoder_cpu_offload and first_param is not None
@@ -293,10 +265,13 @@ class MiniMaxH3ConditioningStage(PipelineStage):
if moved_for_forward:
self.conditioner.to(device)
try:
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
# Keep both H3 prompt-presentation modes under one text-encoding
# range so Nsight Systems exposes their complete conditioning cost.
with nvtx_range("minimax_h3.text_encoding"):
if self.ref2va:
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
else:
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
finally:
if moved_for_forward:
self.conditioner.to("cpu")
@@ -7,10 +7,13 @@ from typing import Any
import torch
from fastvideo.distributed import get_local_torch_device, get_world_group, model_parallel_is_initialized
from fastvideo.distributed import get_local_torch_device, get_sp_group, get_world_group, model_parallel_is_initialized
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE
from fastvideo.models.vaes.minimax_h3_parallel import DEFAULT_DECODE_GATHER_STRATEGY, decode_to_pixels_parallel
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
from fastvideo.profiler import nvtx_range
from fastvideo.pipelines.basic.minimax_h3.packing import (
MiniMaxH3PackedLayout,
unpack_audio_tokens,
@@ -23,6 +26,8 @@ from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.utils import is_pin_memory_available
logger = init_logger(__name__)
def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
@@ -31,6 +36,23 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
return layout
def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]:
"""Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages.
The existing serial path keeps its global-rank-zero output ownership.
Parallel decode assembles once per sequence-parallel group, on that
group's first rank. ``parallel`` is only true when every group rank will
run the decode body — the collectives inside require uniform
participation, so no rank-dependent branch may guard them.
"""
if not model_parallel_is_initialized():
return None, True, False
sp_group = get_sp_group()
if bool(want_parallel) and sp_group.world_size > 1:
return sp_group, sp_group.is_first_rank, True
return sp_group, get_world_group().is_first_rank, False
class MiniMaxH3VideoDecodingStage(PipelineStage):
"""Drop visual condition rows, unpatchify, and decode the target video."""
@@ -55,11 +77,14 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
# Distributed executors consume rank 0's ForwardBatch. Keep a
"""Decode H3 video latents into normalized CPU pixels."""
placeholder = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
sp_group, is_output_rank, parallel = _decode_participation(fastvideo_args, fastvideo_args.vae_parallel_decode)
if not is_output_rank and not parallel:
# Consumers read the output rank's ForwardBatch. Keep a
# verifier-compatible placeholder on other ranks and avoid
# duplicating the full VAE decode and CPU output buffer.
batch.output = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
batch.output = placeholder
return batch
layout = _layout(batch)
@@ -79,19 +104,33 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
try:
latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32))
if fastvideo_args.output_type == "latent":
batch.output = latents.detach().float().cpu()
# No collectives on this path, so uniform participation is
# trivial: every rank returns here.
batch.output = latents.detach().float().cpu() if is_output_rank else placeholder
return batch
output = torch.empty(
self.vae.decoded_pixel_shape(latents.shape),
device="cpu",
dtype=torch.float32,
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
)
# The published decode recipe uses FP16 autocast over FP32 weights.
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
self.vae.decode_to_pixels(latents, output)
batch.output = output
output = None
if is_output_rank:
output = torch.empty(
self.vae.decoded_pixel_shape(latents.shape),
device="cpu",
dtype=torch.float32,
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
)
# Attribute the streamed decoder computation while retaining
# per-chunk device-to-host transfer and pinned-buffer reuse.
with (
nvtx_range("minimax_h3.vae"),
torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"),
):
if parallel:
strategy = fastvideo_args.vae_parallel_decode_strategy or DEFAULT_DECODE_GATHER_STRATEGY
logger.info("MiniMax-H3 VAE decode: sequence-parallel chunks across %d ranks (%s)",
sp_group.world_size, strategy)
decode_to_pixels_parallel(self.vae, latents, output, sp_group, strategy=strategy)
else:
self.vae.decode_to_pixels(latents, output)
batch.output = output if is_output_rank else placeholder
return batch
finally:
if fastvideo_args.vae_cpu_offload:
@@ -121,6 +160,9 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 audio latents into a stereo CPU waveform."""
# Audio decode is sub-second, so preserve the serial path's global
# rank-zero ownership.
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32)
batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate
@@ -144,7 +186,10 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
self._clear_runtime(batch)
return batch
decoded = self.audio_vae.decode(latents).sample.float()
# The range isolates waveform synthesis from packing and runtime
# cleanup so the audio decoder has one stable timeline boundary.
with nvtx_range("minimax_h3.audio_vae"):
decoded = self.audio_vae.decode(latents).sample.float()
if decoded.ndim != 3 or decoded.shape[0] != 2 or decoded.shape[1] != 1:
raise ValueError("MiniMax-H3 audio VAE must decode stereo channels as two mono batch items; "
f"got {tuple(decoded.shape)}.")
@@ -11,8 +11,8 @@ from fastvideo.attention.selector import component_attention_backend, get_attn_b
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.profiler import profiler_region
from fastvideo.hooks.activation_trace import trace_step
from fastvideo.profiler import nvtx_range, profiler_region
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_KEYFRAME_NOISE_AUG,
MiniMaxH3PackedLayout,
@@ -89,6 +89,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Denoise the packed H3 video and audio streams over one shared schedule."""
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
if not isinstance(layout, MiniMaxH3PackedLayout):
raise ValueError("MiniMax-H3 packed layout is missing before denoising.")
@@ -145,9 +146,15 @@ class MiniMaxH3DenoisingStage(PipelineStage):
vsa_exempt = vsa_mode == "exempt"
vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ()))
vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0))
# Run-level tile geometry (256 default, 64 = native Triton path),
# plumbed like the run-level sparsity; the builder validates the
# value against VSA_H3_TILE_SHAPES.
vsa_tile_size = int(fastvideo_args.VSA_tile_size)
try:
with profiler_region("inference_denoising"):
# The stage range groups the complete denoising loop while the
# indexed model ranges retain timing detail for every H3 block.
with profiler_region("inference_denoising"), nvtx_range("minimax_h3.dit"):
for index, (video_timestep,
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
unique_timesteps, timestep_indices = row_timestep_plan[index]
@@ -167,6 +174,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
device=device,
exempt=vsa_exempt,
dense_layers=vsa_dense_layers,
tile_size=vsa_tile_size,
)
# Under torch.compile(mode="reduce-overhead") each denoising
# step must be marked, or cudagraph trees flag cross-step

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