Compare commits
21
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
be1df9fec5 | ||
|
|
e9bbaca07d | ||
|
|
9bfa585448 | ||
|
|
b2062556a9 | ||
|
|
c9c5585758 | ||
|
|
9212f4f218 | ||
|
|
6388db815b | ||
|
|
7a4285189f | ||
|
|
a837fe841a | ||
|
|
f9e3680f11 | ||
|
|
98f761ec45 | ||
|
|
c041318f2c | ||
|
|
604e0205a4 | ||
|
|
13213395b4 | ||
|
|
46afee5998 | ||
|
|
c488fa1211 | ||
|
|
d3cff517cd | ||
|
|
2f3d407406 | ||
|
|
bcffa4026e | ||
|
|
6d6a10be7a | ||
|
|
73dd105f3d |
@@ -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
|
||||
|
||||
@@ -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
@@ -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())
|
||||
Executable
+5
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+87
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+35
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+52
@@ -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"
|
||||
Executable
+6
@@ -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
|
||||
Executable
+205
@@ -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[@]}"
|
||||
Executable
+5
@@ -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
|
||||
Executable
+6
@@ -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
|
||||
Executable
+6
@@ -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
|
||||
Executable
+6
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
Executable
+5
@@ -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
|
||||
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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())
|
||||
@@ -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
|
||||
|
||||
@@ -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',
|
||||
});
|
||||
}
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}')"
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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/)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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/
|
||||
|
||||
@@ -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
@@ -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 && \
|
||||
|
||||
@@ -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 |
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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).
|
||||
|
||||
@@ -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!
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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 "============================================================")
|
||||
|
||||
@@ -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", ®ister_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_
|
||||
@@ -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))
|
||||
+11
-4
@@ -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"
|
||||
|
||||
@@ -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)]
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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())
|
||||
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
Reference in New Issue
Block a user