Compare commits
62
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4e1603634d | ||
|
|
f1eeb6303f | ||
|
|
990d2c2410 | ||
|
|
1772226bf1 | ||
|
|
1fe8e64e23 | ||
|
|
46f18d8cd2 | ||
|
|
2e8db18d94 | ||
|
|
4190c7203f | ||
|
|
8f1443f47b | ||
|
|
42ed546a66 | ||
|
|
d6e020402e | ||
|
|
620e100af4 | ||
|
|
0fde316a19 | ||
|
|
0d99e47e16 | ||
|
|
27f6f0aacd | ||
|
|
c97fb6b3b3 | ||
|
|
3464cb8b03 | ||
|
|
e114fba53f | ||
|
|
093f5e699c | ||
|
|
e24bc12c59 | ||
|
|
4d04c1b01c | ||
|
|
3fb150bbe0 | ||
|
|
320f8a1f8d | ||
|
|
ccfcc3042b | ||
|
|
8f0493637e | ||
|
|
f5ce12f17a | ||
|
|
17cb6737c1 | ||
|
|
f39dbe482c | ||
|
|
5b5608cb37 | ||
|
|
1d4a6037eb | ||
|
|
8ac6526cdc | ||
|
|
08364b2080 | ||
|
|
f81de3926f | ||
|
|
04633096e2 | ||
|
|
b6be3d0c8a | ||
|
|
d5acc7bbae | ||
|
|
9035927da1 | ||
|
|
42d1c79694 | ||
|
|
59000cb933 | ||
|
|
e6b15fc4de | ||
|
|
818daea816 | ||
|
|
d1bd1d8da4 | ||
|
|
d210a076e6 | ||
|
|
7e96218d04 | ||
|
|
6d0ddb44bc | ||
|
|
14c142ba9c | ||
|
|
3fa3f4e333 | ||
|
|
5c7cd391ac | ||
|
|
e7d6c40860 | ||
|
|
12a7cb53ff | ||
|
|
c17d33bf33 | ||
|
|
2aaeee2ab8 | ||
|
|
eb3a394224 | ||
|
|
f673423b51 | ||
|
|
eb0a41528a | ||
|
|
140bd1a6cf | ||
|
|
11f5a8e582 | ||
|
|
71b3cb8c34 | ||
|
|
c85f6a477f | ||
|
|
40d4930d73 | ||
|
|
f9be085243 | ||
|
|
36b53ff350 |
@@ -119,7 +119,7 @@ FastVideo-WorldModel/
|
||||
## Build & Test Commands
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev] # Editable install
|
||||
uv pip install -e ".[dev]" # Editable install
|
||||
pre-commit run --all-files # Lint/format/spell
|
||||
pytest tests/ # Top-level tests
|
||||
pytest fastvideo/tests/ -v # Package tests
|
||||
|
||||
@@ -6,3 +6,5 @@
|
||||
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "reseed-ssim-references", "description": "Re-seed (overwrite) HF reference videos for an existing fastvideo/tests/ssim/ test and a single model id on Modal L40S. Always backs up current refs first, regenerates on Modal, pauses for the user to eyeball before-vs-after, then uploads with --force scoped to --model-id. Sister skill to seed-ssim-references; use when intentional code change has invalidated existing refs", "path": "reseed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "add-model", "description": "Add a new model (or variant) to FastVideo: DiT + configs + pipeline + presets + registry + tests. Walks through FastVideo's single stage-based pipeline architecture with exact file paths and registration hooks.", "path": "add-model/SKILL.md", "status": "draft", "trust": "low"}
|
||||
|
||||
@@ -12,7 +12,7 @@ automates the boilerplate of setting environment variables, picking the right
|
||||
entrypoint, and applying defaults from the closest example script.
|
||||
|
||||
## Prerequisites
|
||||
- The repo is cloned and `fastvideo` is installed (`uv pip install -e .[dev]`).
|
||||
- The repo is cloned and `fastvideo` is installed (`uv pip install -e ".[dev]"`).
|
||||
- Dataset is preprocessed (see `docs/training/data_preprocess.md`).
|
||||
- `WANDB_API_KEY` is set in the environment (or `WANDB_MODE=offline` for local).
|
||||
- GPU resources are available (multi-GPU requires NCCL).
|
||||
|
||||
@@ -0,0 +1,343 @@
|
||||
---
|
||||
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.
|
||||
---
|
||||
|
||||
# Re-seed SSIM Reference Videos
|
||||
|
||||
## Purpose
|
||||
|
||||
Replace the existing SSIM reference videos for a single `(test_file, model_id)`
|
||||
pair on the HF dataset (`FastVideo/ssim-reference-videos`). This is **destructive**
|
||||
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).
|
||||
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.
|
||||
|
||||
Pairs with `seed-ssim-references`, which is the inverse (first-time seeding
|
||||
only, refuses to overwrite). Re-seeding is intentionally a separate, more
|
||||
ceremonial operation because mistakenly clobbering production refs is much
|
||||
harder to recover from than failing closed.
|
||||
|
||||
## When to use
|
||||
|
||||
- An intentional code change (model port fix, kernel upgrade, attention
|
||||
backend swap, hyperparameter change in the test itself) has shifted the
|
||||
expected SSIM output and the existing refs no longer represent the new
|
||||
ground truth.
|
||||
- A test is failing in CI **for the right reason** (the new code is correct,
|
||||
the old refs are stale).
|
||||
|
||||
## When not to use
|
||||
|
||||
- A test is failing for the **wrong** reason (the port is buggy, not the
|
||||
refs). Fix the port; re-seeding hides the bug.
|
||||
- A brand-new test that has no refs on HF yet. Use `seed-ssim-references`.
|
||||
- "Just to clean up drift" without a concrete code change to point at. The
|
||||
PR description has to justify *why* refs changed; without a concrete
|
||||
change, there's nothing to write.
|
||||
|
||||
## Inputs
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `test_file` | Yes | Path to the SSIM test, e.g. `fastvideo/tests/ssim/test_matrixgame_similarity.py`. Validated against `fastvideo/tests/ssim/test_*_similarity.py`. |
|
||||
| `model_id` | Yes | Single model id from the test's `*_MODEL_TO_PARAMS`, e.g. `Matrix-Game-2.0-Diffusers-Base`. Re-seed runs are **per model**. For multi-model tests, invoke the skill once per model. |
|
||||
| `intent_rationale` | Yes | One-line explanation of *why* refs are being regenerated (e.g. "Relax FA-2 head_size whitelist to include 80 — matrix_game now uses FLASH_ATTN instead of TORCH_SDPA"). Recorded in the backup directory and reused in the PR description. |
|
||||
|
||||
Hardcoded:
|
||||
|
||||
- Modal GPU: **L40S** (matches CI; re-seeding from another SKU produces refs
|
||||
that L40S CI cannot match).
|
||||
- Quality tier: **`default`**. `full_quality` is a separate, deliberate
|
||||
operation.
|
||||
- HF repo: `FastVideo/ssim-reference-videos` (override via
|
||||
`FASTVIDEO_SSIM_REFERENCE_HF_REPO`).
|
||||
- Device folder: `L40S_reference_videos`.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
The user has confirmed:
|
||||
|
||||
- `modal` CLI authenticated.
|
||||
- `hf` CLI authenticated, **and** `HF_API_KEY` (or `HUGGINGFACE_HUB_TOKEN` /
|
||||
`HF_TOKEN`) exported with **write** access to
|
||||
`FastVideo/ssim-reference-videos`.
|
||||
- The current branch's code is the change that motivated the re-seed (i.e.
|
||||
`git rev-parse HEAD` is the commit that intentionally invalidated refs).
|
||||
|
||||
Fail fast if any of these are missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Validate inputs and confirm intent
|
||||
|
||||
- Verify `test_file` exists and matches `fastvideo/tests/ssim/test_*_similarity.py`.
|
||||
- Grep the file for `*_MODEL_TO_PARAMS` and assert `model_id` is one of its
|
||||
keys. If the file has only a single hardcoded model, accept that model id
|
||||
as the only valid value.
|
||||
- Print the rationale and ask the user to type **`confirm reseed`** (not just
|
||||
`y` — make it deliberate):
|
||||
|
||||
> About to RE-SEED references for model `<model_id>` from test `<test_file>`.
|
||||
> This will OVERWRITE existing refs on
|
||||
> `FastVideo/ssim-reference-videos/reference_videos/default/L40S_reference_videos/<model_id>/`
|
||||
> after backup + Modal regen + eyeball.
|
||||
>
|
||||
> Reason: `<intent_rationale>`
|
||||
> HEAD: `<git rev-parse --short=12 HEAD>`
|
||||
>
|
||||
> Reply `confirm reseed` to proceed, anything else to abort.
|
||||
|
||||
Stop until the user types exactly `confirm reseed`. Anything else aborts
|
||||
with no side effects.
|
||||
|
||||
### 2. Back up existing refs
|
||||
|
||||
Always required. The backup is the only graceful path back if anything goes
|
||||
wrong later.
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(echo "<model_id>" | tr '/' '_')
|
||||
BACKUP_DIR="ssim_reseed_backup/${TIMESTAMP}_${SHORT_COMMIT}_${MODEL_SAFE}"
|
||||
mkdir -p "$BACKUP_DIR"
|
||||
|
||||
hf download \
|
||||
--repo-type dataset FastVideo/ssim-reference-videos \
|
||||
--include "reference_videos/default/L40S_reference_videos/<model_id>/**" \
|
||||
--local-dir "$BACKUP_DIR"
|
||||
|
||||
mp4_count=$(find "$BACKUP_DIR" -name "*.mp4" | wc -l)
|
||||
echo "Backup mp4 count: $mp4_count"
|
||||
[ "$mp4_count" -gt 0 ] || {
|
||||
echo "ERROR: backup is empty for <model_id>. Either the model id is wrong"
|
||||
echo "or there are no existing refs (use seed-ssim-references instead)."
|
||||
exit 1
|
||||
}
|
||||
|
||||
# Provenance — used in the PR description
|
||||
cat > "$BACKUP_DIR/PROVENANCE.txt" <<EOF
|
||||
test_file: <test_file>
|
||||
model_id: <model_id>
|
||||
head_commit: $(git rev-parse HEAD)
|
||||
timestamp_utc: $(date -u +%FT%TZ)
|
||||
reason: <intent_rationale>
|
||||
EOF
|
||||
```
|
||||
|
||||
If the `hf download` produces zero mp4s, abort — the user has either picked a
|
||||
non-existent `model_id` or there are no refs yet (in which case
|
||||
`seed-ssim-references` is the right tool).
|
||||
|
||||
### 3. Regenerate on Modal L40S
|
||||
|
||||
Mirror CI's exact env recipe so the regenerated refs are byte-comparable to
|
||||
what CI will produce on the same commit. Two differences from CI:
|
||||
|
||||
1. **Pass the same env prefix CI uses** (`IMAGE_VERSION`, `BUILDKITE_*`) — see
|
||||
`.buildkite/pipeline.yml:1-3` and `.buildkite/scripts/pr_test.sh:62-83`.
|
||||
Without this, `ssim_test.py:17-18` resolves a different GHCR image tag
|
||||
(default is `latest`, CI is `py3.12-latest`), and `ssim_test.py:38-46`
|
||||
bakes different values into the image's frozen env block. **Mismatched
|
||||
image or env is the most common source of SSIM drift between reseed and
|
||||
CI runs.**
|
||||
2. **Do not pass `--skip-reference-download`**. Letting the test fetch the
|
||||
existing refs and run the full SSIM compare gives "before" SSIM numbers
|
||||
for the PR description, and the test still produces the new mp4s
|
||||
regardless of whether the comparison passes or fails.
|
||||
|
||||
```bash
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
|
||||
IMAGE_VERSION="py3.12-latest" \
|
||||
BUILDKITE_REPO="$(git config --get remote.origin.url)" \
|
||||
BUILDKITE_COMMIT="$(git rev-parse HEAD)" \
|
||||
BUILDKITE_PULL_REQUEST="${BUILDKITE_PULL_REQUEST:-false}" \
|
||||
modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--git-repo="$(git config --get remote.origin.url)" \
|
||||
--git-commit="$(git rev-parse HEAD)" \
|
||||
--hf-api-key="$HF_API_KEY" \
|
||||
--test-files="<test_file>" \
|
||||
--sync-generated-to-volume \
|
||||
--generated-volume-subdir="$SUBDIR" \
|
||||
--no-fail-fast
|
||||
```
|
||||
|
||||
Capture the printed `modal volume get ...` hint — its `<SUBDIR>` matches
|
||||
`$SUBDIR` and is needed for step 4. Capture the SSIM numbers from the test
|
||||
output (or from the JSON next to the generated mp4) for the PR description.
|
||||
|
||||
### 4. Download generated videos
|
||||
|
||||
```bash
|
||||
modal volume get --force hf-model-weights \
|
||||
ssim_generated_videos/default/"$SUBDIR"/generated_videos \
|
||||
./generated_videos_modal/default
|
||||
```
|
||||
|
||||
After this, the new mp4s live at:
|
||||
|
||||
```
|
||||
./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4
|
||||
```
|
||||
|
||||
`--force` is required when `./generated_videos_modal/default` already exists
|
||||
from a prior run; safe on the first run too.
|
||||
|
||||
### 5. PAUSE — user reviews quality side-by-side
|
||||
|
||||
Print the diff and the comparison:
|
||||
|
||||
```bash
|
||||
echo "=== File list diff (backup vs new) ==="
|
||||
diff -u \
|
||||
<(find "$BACKUP_DIR/reference_videos/default/L40S_reference_videos/<model_id>" -name "*.mp4" \
|
||||
| sed "s|$BACKUP_DIR/reference_videos/default/L40S_reference_videos/||" | sort) \
|
||||
<(find ./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id> -name "*.mp4" \
|
||||
| sed "s|./generated_videos_modal/default/generated_videos/L40S_reference_videos/||" | sort) \
|
||||
|| true
|
||||
|
||||
echo
|
||||
echo "=== SSIM numbers from this run (paste into PR) ==="
|
||||
find ./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id> -name "*_ssim.json" -exec cat {} \;
|
||||
```
|
||||
|
||||
Then stop and tell the user:
|
||||
|
||||
> Old refs backed up to `$BACKUP_DIR`.
|
||||
> New videos in `./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/`.
|
||||
>
|
||||
> Open both in a video player. Confirm the new videos:
|
||||
> 1. Look correct (no obvious artifacts, no black/static frames).
|
||||
> 2. Are *intentionally* different from the backup in the way described
|
||||
> in `<intent_rationale>` (e.g. slight numerical drift only, not a
|
||||
> different scene / different motion / corrupted output).
|
||||
>
|
||||
> Reply **`upload`** to overwrite HF, anything else to abort.
|
||||
> Aborting leaves the backup and new videos on disk for inspection — nothing
|
||||
> on HF changes.
|
||||
|
||||
Do not proceed until the user types exactly `upload`. If they abort, leave
|
||||
everything on disk and stop here.
|
||||
|
||||
### 6. Copy into the local reference layout
|
||||
|
||||
Same as `seed-ssim-references` step 5:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--generated-dir ./generated_videos_modal/default/generated_videos/L40S_reference_videos
|
||||
```
|
||||
|
||||
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
|
||||
### 7. Upload with `--force`, scoped to `--model-id`
|
||||
|
||||
The `--force` flag is what makes this skill different from `seed-ssim-references`.
|
||||
Always pair it with `--model-id` so a typo cannot accidentally overwrite a
|
||||
neighboring model's refs.
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>" \
|
||||
--force
|
||||
```
|
||||
|
||||
The CLI's overwrite guard refuses without `--force`; with `--force` it
|
||||
overwrites only files under
|
||||
`reference_videos/default/L40S_reference_videos/<model_id>/`.
|
||||
|
||||
### 8. Report success and retention guidance
|
||||
|
||||
Print:
|
||||
|
||||
- The HF path that was overwritten (`<repo>/reference_videos/default/L40S_reference_videos/<model_id>/`).
|
||||
- The local backup directory path.
|
||||
- The new SSIM numbers from step 5.
|
||||
- This restore command, in case the PR review surfaces a problem after
|
||||
upload:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>" \
|
||||
--reference-dir "$BACKUP_DIR/reference_videos/default/L40S_reference_videos" \
|
||||
--force
|
||||
```
|
||||
|
||||
- This PR-description checklist (see `fastvideo/tests/ssim/AGENTS.md` →
|
||||
*Updating Reference Videos*):
|
||||
1. Source commit that produced the new refs (HEAD at re-seed time).
|
||||
2. Test command and GPU SKU (`L40S`).
|
||||
3. Before/after SSIM numbers.
|
||||
4. The `<intent_rationale>` from step 1.
|
||||
5. A note that the backup lives at `$BACKUP_DIR` and should be retained
|
||||
until CI on the PR is green.
|
||||
|
||||
Do **not** auto-rerun the SSIM test — the user does that as part of the PR.
|
||||
|
||||
## Failure modes and how to handle them
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before step 2.
|
||||
- **Backup is empty (zero mp4s).** Stop before step 3 — the model id is
|
||||
wrong or the refs don't exist yet (use `seed-ssim-references`).
|
||||
- **Modal run fails before generation.** No mp4s on the volume. Don't
|
||||
upload. Investigate the failure (test crash, OOM, partition exhaustion),
|
||||
fix, then retry from step 3. Backup is still intact.
|
||||
- **Quality regressed (visual or metric).** User aborts at step 5. Backup
|
||||
retained. New videos retained on disk for inspection. Nothing on HF
|
||||
changed. Either fix the underlying code change or abandon the re-seed.
|
||||
- **User confirmed `upload` but later realized the new refs are wrong.**
|
||||
Run the restore command from step 8 with the backup `--reference-dir`.
|
||||
This is exactly why the backup exists.
|
||||
- **Multi-model test, only one model is being re-seeded.** Run the skill
|
||||
once per model id. The `--model-id` scope on upload guarantees the others
|
||||
are untouched.
|
||||
|
||||
## Design notes (for future skill maintainers)
|
||||
|
||||
- Per-`model_id` scope is mandatory. The dataset houses many model subtrees;
|
||||
re-seeding the wrong one is hard to undo without backup.
|
||||
- `default` tier only; `full_quality` is a separate, deliberate operation
|
||||
with different params and ~doubled runtime, and isn't what CI gates on.
|
||||
- The skill deliberately does **not** pass `--skip-reference-download` to
|
||||
Modal so we get pre-reseed SSIM numbers for the PR. The `seed`-skill
|
||||
passes it because no refs exist yet; for re-seed, refs do exist and
|
||||
exposing the comparison is informative.
|
||||
- The two-token confirm (`confirm reseed`, then `upload`) is intentional.
|
||||
Re-seeding is high-blast-radius and should not be one-keystroke.
|
||||
- The backup directory is plain mp4s + `PROVENANCE.txt`. No HF metadata is
|
||||
preserved; the restore path uses `reference_videos_cli.py upload
|
||||
--reference-dir` which doesn't need it.
|
||||
|
||||
## References
|
||||
|
||||
- `.agents/skills/seed-ssim-references/SKILL.md` — the first-time seed
|
||||
skill this one parallels. Read it for the Modal flag rationale shared
|
||||
between the two flows.
|
||||
- `fastvideo/tests/ssim/AGENTS.md` — directory rules, including the PR
|
||||
expectations for any reference-video change (rationale, before/after
|
||||
SSIM, source commit/model/backend).
|
||||
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
|
||||
(with `--model-id`, `--force`), `download`. The overwrite guard at
|
||||
`upload_reference_videos` is the safety net this skill leans on.
|
||||
- `fastvideo/tests/modal/ssim_test.py` — Modal orchestrator;
|
||||
`--sync-generated-to-volume`, `--generated-volume-subdir`,
|
||||
`--skip-reference-download`, `--no-fail-fast`.
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|------|--------|
|
||||
| 2026-05-02 | Initial version. Sister skill to `seed-ssim-references`, scoped to single `(test_file, model_id)` re-seeds, with mandatory backup and two-token confirm. |
|
||||
@@ -1,26 +1,41 @@
|
||||
---
|
||||
name: seed-ssim-references
|
||||
description: Seed HF reference videos for a single newly-added SSIM test. Runs the test on Modal L40S, downloads the generated mp4s via `modal volume get`, pauses for the user to eyeball quality, then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
|
||||
description: Seed HF reference artefacts for a single newly-added SSIM test (pixel `.mp4` for `run_text_to_video_similarity_test`-style tests, or latent `.pt` for `run_text_to_latent_similarity_test`-style tests). Runs the test on Modal L40S, downloads the generated artefacts via `modal volume get`, pauses for the user to verify (visual eyeball for mp4, numerics dump for pt), then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
|
||||
---
|
||||
|
||||
# Seed SSIM Reference Videos
|
||||
# Seed SSIM Reference Artefacts (mp4 or pt)
|
||||
|
||||
## Purpose
|
||||
|
||||
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
|
||||
reference videos exist on the HF dataset (`FastVideo/ssim-reference-videos`).
|
||||
reference artefacts exist on the HF dataset
|
||||
(`FastVideo/ssim-reference-videos`). The dataset hosts two kinds of artefacts
|
||||
side-by-side per `(model_id, backend, prompt)`:
|
||||
|
||||
- **`.mp4`** — pixel ground-truth for tests that call
|
||||
`run_text_to_video_similarity_test` / `run_image_to_video_similarity_test`
|
||||
in `inference_similarity_utils.py`. Compared via SSIM.
|
||||
- **`.pt`** — pre-VAE latent bundle (fp16 full latent + fp32 slice +
|
||||
metadata + `slice_spec` + `format_version`) for tests that call
|
||||
`run_text_to_latent_similarity_test` in `latent_similarity_utils.py`.
|
||||
Compared via cosine distance on the slice and the full tensor.
|
||||
|
||||
This skill:
|
||||
|
||||
1. Runs the test on Modal's L40S pool to generate the videos.
|
||||
2. Downloads them to the local repo via `modal volume get`.
|
||||
3. Pauses so the user can eyeball the mp4s and confirm quality.
|
||||
4. Uploads only the new test's files to HF, with a guard that refuses to
|
||||
1. Detects which artefact type the test produces (pixel vs latent).
|
||||
2. Runs the test on Modal's L40S pool to generate the artefacts.
|
||||
3. Downloads them to the local repo via `modal volume get`.
|
||||
4. Pauses so the user can verify quality:
|
||||
- **mp4**: visual eyeball in a video player.
|
||||
- **pt**: numerics dump (shape, slice stats, NaN/Inf check, metadata).
|
||||
5. Uploads only the new test's files to HF, with a guard that refuses to
|
||||
overwrite anything already present.
|
||||
|
||||
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 mp4 without crashing. The skill does not re-test locally; it
|
||||
goes straight to Modal L40S (which is what CI uses).
|
||||
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).
|
||||
|
||||
## When to use
|
||||
|
||||
@@ -69,7 +84,7 @@ Fail fast if the token env var is missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Ask for the test file
|
||||
### 1. Ask for the test file, then detect artefact type
|
||||
|
||||
If the user didn't name one, ask: *"Which SSIM test file do you want to seed
|
||||
references for? (e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`)"*.
|
||||
@@ -80,6 +95,22 @@ Validate:
|
||||
- File defines a `*_MODEL_TO_PARAMS` dict — grep it to extract the set of
|
||||
model ids. Those ids drive step 5.
|
||||
|
||||
Detect artefact type by inspecting the file's imports / helper call:
|
||||
|
||||
- **latent** (`.pt`) — file imports `run_text_to_latent_similarity_test`
|
||||
from `fastvideo.tests.ssim.latent_similarity_utils` (or any other helper
|
||||
that ends with `_latent_similarity_test`).
|
||||
- **pixel** (`.mp4`) — file imports
|
||||
`run_text_to_video_similarity_test` / `run_image_to_video_similarity_test`
|
||||
from `fastvideo.tests.ssim.inference_similarity_utils`, OR uses the
|
||||
legacy custom-inline helper pattern (see `test_gamecraft`,
|
||||
`test_longcat`, etc.). Default to pixel when both heuristics fail.
|
||||
|
||||
Record `ARTEFACT_TYPE ∈ {pixel, latent}` for use in step 4. Steps 2, 3, 5,
|
||||
and 6 are artefact-type-agnostic — `_iter_reference_files`,
|
||||
`copy_generated_to_reference`, and `upload_reference_videos` already walk
|
||||
both `.mp4` and `.pt` (see `reference_videos_cli.py`).
|
||||
|
||||
If either check fails, stop and tell the user what's wrong.
|
||||
|
||||
### 2. Run the test on Modal L40S
|
||||
@@ -92,9 +123,19 @@ TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
```
|
||||
|
||||
Then launch the Modal run:
|
||||
Then launch the Modal run. The `IMAGE_VERSION` and `BUILDKITE_*` env-prefix
|
||||
**must** match what CI exports in `.buildkite/scripts/pr_test.sh`, otherwise
|
||||
`fastvideo/tests/modal/ssim_test.py` resolves a different GHCR image tag
|
||||
(default is `latest`, CI is `py3.12-latest`) and bakes different values into
|
||||
the image's frozen env block (`ssim_test.py:17-18, 38-46`). Mismatched image
|
||||
or env produces SSIM drift that doesn't show up until the same commit runs
|
||||
in CI.
|
||||
|
||||
```bash
|
||||
IMAGE_VERSION="py3.12-latest" \
|
||||
BUILDKITE_REPO="$(git config --get remote.origin.url)" \
|
||||
BUILDKITE_COMMIT="$(git rev-parse HEAD)" \
|
||||
BUILDKITE_PULL_REQUEST="${BUILDKITE_PULL_REQUEST:-false}" \
|
||||
modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--git-repo="$(git config --get remote.origin.url)" \
|
||||
--git-commit="$(git rev-parse HEAD)" \
|
||||
@@ -106,6 +147,19 @@ modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--no-fail-fast
|
||||
```
|
||||
|
||||
Env prefix rationale (parity with CI; see `.buildkite/pipeline.yml:1-3` and
|
||||
`.buildkite/scripts/pr_test.sh:62-83`):
|
||||
- `IMAGE_VERSION=py3.12-latest`: pins the Modal image tag to the same one CI
|
||||
uses. Without this, `ssim_test.py:17` falls back to `latest`, which on
|
||||
GHCR is built from `Dockerfile.python3.10` — different Python, torch, and
|
||||
flash-attn wheel than CI's `py3.12-latest` (`infra-build-image.yml:51-67`,
|
||||
`_template-build-image.yml:65-101`).
|
||||
- `BUILDKITE_REPO`/`BUILDKITE_COMMIT`/`BUILDKITE_PULL_REQUEST`: mirror what
|
||||
Buildkite exports. `ssim_test.py:38-46` bakes these into the image's
|
||||
`.env(...)` block; mismatched values can perturb in-container code paths
|
||||
that branch on PR-vs-non-PR. `false` for `BUILDKITE_PULL_REQUEST` matches
|
||||
Buildkite's "non-PR build" sentinel.
|
||||
|
||||
Flag rationale:
|
||||
- `--skip-reference-download`: no refs exist yet, so conftest must not try to
|
||||
pull them.
|
||||
@@ -143,17 +197,59 @@ get` preserves that trailing `generated_videos/` segment.
|
||||
|
||||
### 4. PAUSE — user reviews quality
|
||||
|
||||
Print the list of downloaded mp4s and their paths, then stop. Tell the user:
|
||||
Type-aware verification.
|
||||
|
||||
**For `ARTEFACT_TYPE = pixel`** — list the downloaded mp4s and ask the user to
|
||||
open them in a video player:
|
||||
|
||||
> "Generated videos downloaded to `./generated_videos_modal/default/generated_videos/L40S_reference_videos/`. Please open them and confirm the quality looks correct. Reply **`upload`** to continue, or anything else to abort."
|
||||
|
||||
**For `ARTEFACT_TYPE = latent`** — `.pt` files are not human-watchable. Print
|
||||
a numerics dump for each `.pt` so the user can sanity-check shape, distribution,
|
||||
and metadata:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from pathlib import Path
|
||||
ROOT = Path("./generated_videos_modal/default/generated_videos/L40S_reference_videos")
|
||||
for p in sorted(ROOT.rglob("*.pt")):
|
||||
d = torch.load(p, map_location="cpu", weights_only=False)
|
||||
s = d["expected_slice"]
|
||||
L = d["latent"].float()
|
||||
print(f"=== {p.relative_to(ROOT)} ===")
|
||||
print(f" format_version: {d['format_version']}")
|
||||
print(f" shape: {d['shape']}")
|
||||
print(f" dtype_original: {d['dtype_original']}")
|
||||
print(f" slice_spec: {d['slice_spec']}")
|
||||
print(f" slice shape={tuple(s.shape)} mean={s.mean():+.4f} std={s.std():.4f} min={s.min():+.4f} max={s.max():+.4f}")
|
||||
print(f" latent shape={tuple(L.shape)} mean={L.mean():+.4f} std={L.std():.4f} min={L.min():+.4f} max={L.max():+.4f}")
|
||||
print(f" finite: latent NaN={torch.isnan(L).any().item()} Inf={torch.isinf(L).any().item()}; "
|
||||
f"slice NaN={torch.isnan(s).any().item()} Inf={torch.isinf(s).any().item()}")
|
||||
print(f" metadata: {d['metadata']}\n")
|
||||
```
|
||||
|
||||
Sanity criteria:
|
||||
- `format_version == 1` (matches `LATENT_REFERENCE_FORMAT_VERSION`).
|
||||
- `shape` matches what the model produces (e.g. LTX-2 distilled =
|
||||
`[1, 128, T_lat, H_lat, W_lat]`; Stable Audio Open 1.0 = `[1, 64, 1024]`).
|
||||
- `slice_spec.kind` matches a registered kind (`corner_3x3_first_frame`
|
||||
for video, `audio_first_8_timesteps` for audio).
|
||||
- No `NaN`/`Inf`. `mean ≈ 0`, `std ≈ 1` (denoised latents stay close to
|
||||
the initial Gaussian distribution; very wide deviations suggest
|
||||
numerical drift).
|
||||
- `metadata.prompt` matches the test's prompt.
|
||||
|
||||
Then ask:
|
||||
|
||||
> "Numerics look right? Reply **`upload`** to continue, or anything else to abort."
|
||||
|
||||
Do not proceed until the user explicitly says `upload`. If they abort, leave
|
||||
everything on disk so they can inspect further — no cleanup.
|
||||
|
||||
### 5. Copy into the local reference layout
|
||||
|
||||
Scoped copy — only the new test's mp4s. Loop over each `<model_id>` extracted
|
||||
in step 1:
|
||||
Scoped copy — only the new test's artefacts. Single command works for both
|
||||
artefact types because `_iter_reference_files` walks `.mp4` and `.pt`:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
@@ -163,12 +259,13 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
```
|
||||
|
||||
(The `--generated-dir` points at the device-folder root inside the
|
||||
downloaded tree; `copy-local` walks all `<model>/<backend>/*.mp4`
|
||||
downloaded tree; `copy-local` walks all `<model>/<backend>/*.{mp4,pt}`
|
||||
underneath it. Since the Modal run was scoped to a single test file via
|
||||
`--test-files`, only that test's model(s) are present — so the copy is
|
||||
implicitly per-test.)
|
||||
|
||||
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
Result for pixel: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
Result for latent: same path with `.pt` extension.
|
||||
|
||||
### 6. Upload to HF — scoped per model_id, with overwrite guard
|
||||
|
||||
@@ -201,33 +298,54 @@ it will auto-download the refs they just uploaded.
|
||||
## Failure modes and how to handle them
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before step 2. The Modal run needs it (passed
|
||||
via `--hf-api-key`), and step 6 needs it for upload.
|
||||
- **Modal run fails before generation.** No mp4s on the volume — nothing to
|
||||
download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
|
||||
via `--hf-api-key`), and step 6 needs it for upload. If the user
|
||||
ran `hf auth login` instead of exporting an env var, read the cached
|
||||
token via `huggingface_hub.get_token()` and forward it to Modal as
|
||||
`--hf-api-key="$CACHED_TOKEN"`.
|
||||
- **Modal run fails before generation.** No artefacts on the volume — nothing
|
||||
to download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
|
||||
and retry from step 2.
|
||||
- **`./generated_videos_modal/default/L40S_reference_videos/` missing after
|
||||
`modal volume get`.** The run didn't produce videos (most likely the test
|
||||
crashed before writing, or `REQUIRED_GPUS` exceeded the partition capacity
|
||||
— see Modal logs).
|
||||
`modal volume get`.** The run didn't produce artefacts (most likely the
|
||||
test crashed before writing, or `REQUIRED_GPUS` exceeded the partition
|
||||
capacity — see Modal logs).
|
||||
- **Latent test crashed with FSDP / inference_mode error
|
||||
(`RuntimeError: Inference tensors do not track version counter`).** The
|
||||
test must pass `init_kwargs_override={"use_fsdp_inference": False}` when
|
||||
`sp_size == 1` — see `test_stable_audio_similarity.py` for the pattern.
|
||||
Fix in the test, push, retry.
|
||||
- **Upload guard fires (files already exist).** The test name / model id
|
||||
collides with something already on HF. Verify the user actually wants to
|
||||
replace existing refs; if so, re-run the upload with `--force`. If not,
|
||||
rename the model id in `*_MODEL_TO_PARAMS` and re-seed.
|
||||
- **Quality looks wrong in step 4.** Abort. The mp4s stay on disk for
|
||||
- **Quality looks wrong in step 4.** Abort. The artefacts stay on disk for
|
||||
inspection. The fix is usually in the test's params (resolution, steps,
|
||||
seed) — edit the test, then re-run the skill.
|
||||
- For latent: also check `slice_spec.kind` matches the latent rank
|
||||
(`corner_3x3_first_frame` requires 5-D, `audio_first_8_timesteps`
|
||||
requires 3-D); a rank/kind mismatch raises in `_extract_expected_slice`.
|
||||
|
||||
## Design notes (for future skill maintainers)
|
||||
|
||||
- The skill deliberately runs on Modal, **not** locally, because the CI
|
||||
runner is L40S. Seeding from a different GPU SKU produces refs that CI's
|
||||
L40S runs can't match (SSIM drifts across SKUs).
|
||||
L40S runs can't match (pixel SSIM drifts across SKUs; latent cosine has
|
||||
tighter cross-SKU bf16 drift but the configured tolerances assume
|
||||
same-SKU seed → same-SKU verify).
|
||||
- The skill is default-tier only. `full_quality` refs are seeded by a
|
||||
separate, deliberate operation — they double runtime and aren't what CI
|
||||
gates on.
|
||||
- The overwrite guard in `reference_videos_cli.py upload` is default-on
|
||||
specifically because this skill exists. Re-seeding is a distinct operation
|
||||
that requires explicit `--force`.
|
||||
- Both artefact types share the same Modal flow: the orchestrator sets
|
||||
`--skip-reference-download` + `--no-fail-fast`, runs pytest, the test's
|
||||
helper writes the artefact (`.mp4` via `imageio` for pixel,
|
||||
`save_latent_reference` → `torch.save` for latent) BEFORE the
|
||||
missing-reference assertion raises. `_sync_generated_videos_to_volume` in
|
||||
`ssim_test.py` does a `shutil.copytree` of the whole `generated_videos/`
|
||||
tree, picking up `.mp4`, `.pt`, and the `*_ssim.json` / `*_latent.json`
|
||||
metric files alongside.
|
||||
|
||||
## References
|
||||
|
||||
@@ -236,10 +354,17 @@ it will auto-download the refs they just uploaded.
|
||||
`--skip-reference-download`, `--no-fail-fast`.
|
||||
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
|
||||
(with `--model-id`, `--force`), `download`, `ensure` subcommands.
|
||||
Extension allowlist is `REFERENCE_EXTENSIONS = VIDEO_EXTENSIONS +
|
||||
LATENT_EXTENSIONS` (`.pt`).
|
||||
- `fastvideo/tests/ssim/README.md` — reference layout, HF repo conventions.
|
||||
- `fastvideo/tests/ssim/inference_similarity_utils.py` —
|
||||
`run_text_to_video_similarity_test` + `_build_init_kwargs`: what each test
|
||||
config passes to `VideoGenerator.from_pretrained`.
|
||||
- `fastvideo/tests/ssim/inference_similarity_utils.py` — pixel helpers
|
||||
(`run_text_to_video_similarity_test`,
|
||||
`run_image_to_video_similarity_test`, `build_init_kwargs`).
|
||||
- `fastvideo/tests/ssim/latent_similarity_utils.py` — latent helper
|
||||
(`run_text_to_latent_similarity_test`), slice spec dispatch
|
||||
(`_extract_expected_slice`), reference schema
|
||||
(`save_latent_reference` / `load_latent_reference`),
|
||||
`LATENT_REFERENCE_FORMAT_VERSION`.
|
||||
|
||||
## Changelog
|
||||
|
||||
@@ -248,3 +373,4 @@ it will auto-download the refs they just uploaded.
|
||||
| 2026-04-17 | Initial version (Modal sync-to-volume flow). |
|
||||
| 2026-04-21 | Rewrite: single-test scope, explicit user-review pause, per-`model_id` upload, HF overwrite guard. Dropped `scripts/seed_ssim.sh`. |
|
||||
| 2026-04-21 | Post-first-run fixes: `modal volume get` needs `--force` when parent exists; download tree has an extra `generated_videos/` level so `--generated-dir` must reflect it. |
|
||||
| 2026-05-01 | Latent (`*.pt`) artefact support: artefact-type detection in step 1, type-aware verification (visual eyeball for mp4, numerics dump for pt) in step 4, FSDP+inference_mode failure-mode added, design notes for the unified Modal flow. Triggered by PR #1253 (LTX-2 latent migration + Stable Audio latent test). |
|
||||
|
||||
@@ -15,8 +15,21 @@ log "Project root: $PROJECT_ROOT"
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
python3 -m pip install modal
|
||||
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "uv not found, bootstrapping..."
|
||||
if ! curl -LsSf https://astral.sh/uv/install.sh | sh; then
|
||||
log "Error: Failed to bootstrap uv via astral.sh installer."
|
||||
exit 1
|
||||
fi
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "Error: uv still not on PATH after bootstrap."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
# --break-system-packages preserves prior `pip install --user` semantics on PEP 668 agents.
|
||||
uv pip install --system --break-system-packages modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
@@ -82,7 +95,7 @@ upload_performance_artifacts() {
|
||||
|
||||
_upload_dashboard() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_*${SHORT_SHA}*" | head -n 1)
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_${SHORT_SHA}_*" | head -n 1)
|
||||
log "TARGET dashboard: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
@@ -96,7 +109,7 @@ upload_performance_artifacts() {
|
||||
|
||||
_upload_perf_summary() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "perf_*${SHORT_SHA}*" | head -n 1)
|
||||
target=$(find "$LOCAL_DIR" -name "perf_${SHORT_SHA}_*" | head -n 1)
|
||||
log "TARGET perf summary: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
|
||||
@@ -13,8 +13,21 @@ log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "uv not found, bootstrapping..."
|
||||
if ! curl -LsSf https://astral.sh/uv/install.sh | sh; then
|
||||
log "Error: Failed to bootstrap uv via astral.sh installer."
|
||||
exit 1
|
||||
fi
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "Error: uv still not on PATH after bootstrap."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
# --break-system-packages preserves prior `pip install --user` semantics on PEP 668 agents.
|
||||
uv pip install --system --break-system-packages pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
|
||||
@@ -37,10 +37,11 @@ jobs:
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-mkdocs.txt
|
||||
run: uv pip install --system -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
@@ -56,10 +56,11 @@ jobs:
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install build dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install build twine wheel
|
||||
run: uv pip install --system build twine wheel
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
|
||||
@@ -131,11 +131,13 @@ jobs:
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
uv pip install --system typing-extensions==4.12.2
|
||||
uv pip install --system --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
@@ -145,20 +147,20 @@ jobs:
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
|
||||
uv pip install --system setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
cd fastvideo-kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
# Release builds are produced on GPU-less runners, so force-enable TK and target Hopper.
|
||||
export TORCH_CUDA_ARCH_LIST="9.0a"
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
|
||||
|
||||
|
||||
# Build standard wheel (no local version suffix) for PyPI
|
||||
python -m build --wheel --outdir dist
|
||||
|
||||
|
||||
# Fix the wheel to be manylinux compliant
|
||||
pip install auditwheel
|
||||
uv pip install --system auditwheel
|
||||
# Point auditwheel at torch libs, but do not vendor them into the wheel.
|
||||
TORCH_LIB_DIR=$(python - <<'PY'
|
||||
import os
|
||||
@@ -211,10 +213,13 @@ jobs:
|
||||
pattern: 'fastvideo_kernel-py*'
|
||||
merge-multiple: true
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
pip install build scikit-build-core cmake ninja
|
||||
|
||||
uv pip install --system build scikit-build-core cmake ninja
|
||||
|
||||
cd fastvideo-kernel
|
||||
# We don't need full CUDA/Torch to just package the source (sdist)
|
||||
python -m build --sdist --outdir dist
|
||||
|
||||
@@ -92,3 +92,11 @@ preprocess_output_text/
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
|
||||
# Local clones of upstream repos used only for parity testing.
|
||||
/stable-audio-tools/
|
||||
/daVinci-MagiHuman/
|
||||
|
||||
# Converted model weights (produced by scripts/checkpoint_conversion/*).
|
||||
# Tens of GB; should live on HF, not in git.
|
||||
/converted_weights/
|
||||
|
||||
@@ -7,21 +7,13 @@ exclude: |
|
||||
fastvideo-kernel/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
assets/prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
examples/.*|
|
||||
\.agents/.*|
|
||||
.github/workflows/publish-fastvideo.yml|
|
||||
.github/workflows/_template-build-image.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
.github/workflows/_template-build-image.yml
|
||||
)
|
||||
repos:
|
||||
- repo: https://github.com/google/yapf
|
||||
|
||||
@@ -11,7 +11,7 @@
|
||||
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
- `uv pip install -e .[dev]`: editable install with lint/test extras.
|
||||
- `uv pip install -e ".[dev]"`: editable install with lint/test extras.
|
||||
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
|
||||
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
|
||||
- `pytest tests/`: run top-level test suite.
|
||||
@@ -23,7 +23,8 @@
|
||||
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
|
||||
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
|
||||
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
|
||||
- Target line length is 80.
|
||||
- Lint via `pre-commit run --files <changed paths>` (or `pre-commit run --all-files` for a full sweep) before committing. Do not shell out to `yapf`/`ruff`/`codespell`/`mypy` directly — pre-commit chains them with the project's config and respects the `.pre-commit-config.yaml` excludes (e.g. `fastvideo/tests/` is intentionally skipped). If pre-commit reports `(no files to check)` for your paths, that exclude is deliberate — don't bypass it.
|
||||
- Target line length is 120 (configured in `pyproject.toml` for ruff, yapf, and isort).
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
@@ -54,3 +55,31 @@ This repository is agent-friendly. Before doing any work, read:
|
||||
If you are exploring a new procedure that has no existing SOP, document your
|
||||
progress in `.agents/exploration/` and flag it for review at the end of your
|
||||
session.
|
||||
|
||||
## Per-Directory AGENTS.md
|
||||
|
||||
Local guidance lives next to the code. Read the in-scope file before editing:
|
||||
|
||||
| Directory | What it covers |
|
||||
|-----------|----------------|
|
||||
| `fastvideo/AGENTS.md` | Core package map, public API, registry-driven model dispatch |
|
||||
| `fastvideo/configs/AGENTS.md` | Arch + pipeline config dataclasses, `param_names_mapping` |
|
||||
| `fastvideo/models/AGENTS.md` | DiT / VAE / encoder / scheduler / loader layout (pre-commit excluded) |
|
||||
| `fastvideo/layers/AGENTS.md` | Tensor-parallel linear/attention layer rules for ports |
|
||||
| `fastvideo/attention/AGENTS.md` | Backend registry + env-var override |
|
||||
| `fastvideo/pipelines/AGENTS.md` | Stage ABC, `basic/<model>/`, `preprocess/`, presets |
|
||||
| `fastvideo/training/AGENTS.md` | Legacy monolithic pipelines (frozen for existing models) |
|
||||
| `fastvideo/train/AGENTS.md` | New modular trainer (methods × models × callbacks, YAML) |
|
||||
| `fastvideo/tests/AGENTS.md` | Test taxonomy, conftest, pre-commit-excluded path |
|
||||
| `fastvideo/tests/ssim/AGENTS.md` | GPU SSIM regression authoring + reference video sync |
|
||||
| `scripts/checkpoint_conversion/AGENTS.md` | Adding a converter for a new HF/official checkpoint |
|
||||
|
||||
## Critical: Two Training Stacks Coexist
|
||||
|
||||
- `fastvideo/training/` — legacy, monolithic per-model `*_training_pipeline.py` and
|
||||
`*_distillation_pipeline.py`. Still authoritative for shipped models.
|
||||
- `fastvideo/train/` — new modular framework (composable methods × models × callbacks
|
||||
driven by YAML). Preferred for new training work.
|
||||
|
||||
Pick the matching stack before editing. Do not migrate a pipeline between them
|
||||
without an explicit ask — the conventions and config surfaces differ.
|
||||
|
||||
@@ -128,7 +128,7 @@ class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "openai/clip-vit-base-patch32"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
raise ImportError("Please install transformers: uv pip install transformers")
|
||||
super().__init__(device)
|
||||
self.processor = CLIPProcessor.from_pretrained(model_name)
|
||||
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
|
||||
@@ -171,7 +171,7 @@ class VideoMAEFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "MCG-NJU/videomae-base"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
raise ImportError("Please install transformers: uv pip install transformers")
|
||||
super().__init__(device)
|
||||
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
@@ -57,7 +57,7 @@ class I3DFeatureExtractor(nn.Module):
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"pip install huggingface_hub") from e
|
||||
f"uv pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless transformers huggingface_hub
|
||||
uv pip install -q opencv-python-headless transformers huggingface_hub
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
pip install -q opencv-python-headless
|
||||
uv pip install -q opencv-python-headless
|
||||
+2
-2
@@ -38,10 +38,10 @@ cp -r /path/to/FastVideo/comfyui /path/to/ComfyUI/custom_nodes/FastVideo
|
||||
|
||||
#### Install dependencies:
|
||||
|
||||
Currently, the only dependency is `fastvideo`, which can be installed using pip.
|
||||
Currently, the only dependency is `fastvideo`, which can be installed with `uv`.
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
#### Install missing custom nodes:
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp310-cp310-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp310-cp310-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.11 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp311-cp311-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp311-cp311-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp312-cp312-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,7 +42,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
@@ -50,7 +50,7 @@ COPY . .
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -43,7 +43,7 @@ COPY . .
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[rocm] && \
|
||||
uv pip install --no-cache-dir -e ".[rocm]" && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
pip install -r requirements-mkdocs.txt
|
||||
uv pip install -r requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -0,0 +1,210 @@
|
||||
# Activation Trace Mode
|
||||
|
||||
!!! note
|
||||
This page covers Extension 0 (module forward hooks), which is the implemented
|
||||
tracing mechanism. Extensions 1-3 are design sketches for future work and are
|
||||
**not yet implemented**.
|
||||
|
||||
## Overview
|
||||
|
||||
Activation trace mode is a zero-overhead-when-off, env-gated mechanism for
|
||||
dumping per-layer activation statistics during FastVideo inference. Its primary
|
||||
use case is **parity debugging across model ports**: enable tracing on both
|
||||
FastVideo and the upstream reference implementation, then `diff` the resulting
|
||||
JSONL files to find the first divergent layer.
|
||||
|
||||
The mechanism is intentionally narrow. It doesn't replace general logging,
|
||||
profiling, or function tracing. It answers one question: "at which layer do
|
||||
FastVideo and the reference model first produce different numbers?"
|
||||
|
||||
## When to use
|
||||
|
||||
- Investigating numerical drift between FastVideo and an upstream reference.
|
||||
- Debugging mid-pipeline divergence (e.g., one block produces wrong output while earlier blocks match).
|
||||
- Validating that a refactor preserves bf16 noise-floor behavior across many layers.
|
||||
|
||||
## When NOT to use
|
||||
|
||||
| Goal | Use instead |
|
||||
|---|---|
|
||||
| General logging | `init_logger(__name__)` |
|
||||
| Per-stage timing | `FASTVIDEO_STAGE_LOGGING` |
|
||||
| Profiling kernel timings | `FASTVIDEO_TORCH_PROFILER_DIR` (see [Profiling](profiling.md)) |
|
||||
| Function-call tracing | `FASTVIDEO_TRACE_FUNCTION` (heavy) |
|
||||
|
||||
## Quickstart
|
||||
|
||||
```bash
|
||||
FASTVIDEO_TRACE_ACTIVATIONS=1 \
|
||||
FASTVIDEO_TRACE_LAYERS="^block\.layers\.[0-9]+$" \
|
||||
FASTVIDEO_TRACE_STATS="abs_mean,sum,max,shape" \
|
||||
FASTVIDEO_TRACE_OUTPUT="/tmp/fv_trace.jsonl" \
|
||||
python examples/inference/basic/basic_magi_human.py
|
||||
```
|
||||
|
||||
Each line in `/tmp/fv_trace.jsonl` is a JSON record:
|
||||
|
||||
```json
|
||||
{"module": "block.layers.0", "tensor": "out", "step": 0, "abs_mean": 1.234, "sum": -5.678, "max": 9.012, "shape": [1, 4096, 5120]}
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
| Env var | Default | Description |
|
||||
|---|---|---|
|
||||
| `FASTVIDEO_TRACE_ACTIVATIONS` | `False` | Master toggle. When unset or false, **zero overhead** in the production hot path. |
|
||||
| `FASTVIDEO_TRACE_LAYERS` | `""` (all) | Python regex filter applied to `model.named_modules()` names. Empty string matches all modules. |
|
||||
| `FASTVIDEO_TRACE_STATS` | `"abs_mean,sum"` | Comma-separated stats to compute. Available: `abs_mean`, `sum`, `min`, `max`, `mean`, `std`, `shape`, `dtype`. |
|
||||
| `FASTVIDEO_TRACE_OUTPUT` | `"/tmp/fv_trace_<pid>.jsonl"` | Output file path. `<pid>` is replaced with the process ID at runtime. |
|
||||
| `FASTVIDEO_TRACE_STEPS` | `""` (all) | Comma-separated denoising step indices to capture. Empty string captures all steps. |
|
||||
|
||||
## Workflow: parity-debug a model port
|
||||
|
||||
1. Set up a tightly-controlled comparison: a parity test or a small standalone
|
||||
script that loads both the FastVideo model and the upstream reference with
|
||||
identical inputs and seeds.
|
||||
|
||||
2. Run the FastVideo side with tracing on:
|
||||
|
||||
```bash
|
||||
FASTVIDEO_TRACE_ACTIVATIONS=1 \
|
||||
FASTVIDEO_TRACE_LAYERS="<your regex>" \
|
||||
FASTVIDEO_TRACE_OUTPUT="/tmp/fv_trace_fv.jsonl" \
|
||||
python <fv_runner.py>
|
||||
```
|
||||
|
||||
3. Run the upstream side. The upstream repo needs separate instrumentation. See
|
||||
"Hooking the upstream side" below.
|
||||
|
||||
4. Sort both files by `(module, step)` if needed, then diff:
|
||||
|
||||
```bash
|
||||
diff /tmp/fv_trace_fv.jsonl /tmp/fv_trace_upstream.jsonl
|
||||
```
|
||||
|
||||
5. The first divergent line identifies the first layer where FastVideo and the
|
||||
upstream produce different outputs. Start debugging there.
|
||||
|
||||
## Architecture (Extension 0: module forward hooks)
|
||||
|
||||
At pipeline initialization, `attach_activation_trace()` reads the env vars once.
|
||||
If `FASTVIDEO_TRACE_ACTIVATIONS` is unset or false, the function returns
|
||||
immediately and no hooks are registered. If tracing is on, it walks
|
||||
`model.named_modules()`, filters by the layer regex, and registers an
|
||||
`ActivationStatHook` on each matching module.
|
||||
|
||||
During the forward pass, each hook fires after its module completes, computes
|
||||
the requested stats on the output tensor, and appends a JSON record to the
|
||||
output file.
|
||||
|
||||
```
|
||||
ComposedPipelineBase
|
||||
└─ attach_activation_trace()
|
||||
├─ reads env vars (once at startup)
|
||||
├─ if off: returns None immediately
|
||||
└─ if on: walks named_modules()
|
||||
└─ registers ActivationStatHook on matching modules
|
||||
└─ on each forward: compute stats → append JSONL
|
||||
```
|
||||
|
||||
### Zero-overhead-when-off guarantee
|
||||
|
||||
- The env var check happens **once at startup** inside `attach_activation_trace()`.
|
||||
- If the env var is unset or false, the function returns `None` immediately.
|
||||
- No hooks are registered. No branches are added to the production forward path.
|
||||
- The only cost when tracing is off is one env var lookup at pipeline
|
||||
initialization, which takes under a microsecond.
|
||||
|
||||
### Hooking the upstream side
|
||||
|
||||
The upstream reference repo isn't part of FastVideo, so it can't read FastVideo
|
||||
env vars directly. Two options:
|
||||
|
||||
**Option 1: Inline patch** in your local clone of the upstream repo. Add
|
||||
`register_forward_hook` calls in the same shape as `ActivationStatHook`. Clean
|
||||
up afterward with `git stash` or `git checkout HEAD -- <file>`.
|
||||
|
||||
**Option 2: Wrapper script**. Write a small Python harness that imports the
|
||||
upstream model, walks its `named_modules()`, and attaches hooks externally.
|
||||
This is the same pattern used in
|
||||
`tests/local_tests/transformers/_debug_magi_human_block_parity.py`.
|
||||
|
||||
The `add-model-trace` skill at `~/.config/opencode/skill/add-model-trace/`
|
||||
provides a script template for this purpose.
|
||||
|
||||
## Future extensions (design only, not yet implemented)
|
||||
|
||||
### Extension 1: FX/Dynamo backend graph rewrite
|
||||
|
||||
**Granularity**: per-FX-node (every matmul, every add).
|
||||
|
||||
**Mechanism**: a `torch.compile` backend that takes the captured `GraphModule`
|
||||
and inserts logger nodes after each op. Compiles into a separate artifact from
|
||||
the production graph.
|
||||
|
||||
**Off semantics**: zero overhead. The production compile path is untouched.
|
||||
|
||||
**When to add**: if you need to trace inside a `torch.compile`'d graph and
|
||||
Extension 0 is too coarse.
|
||||
|
||||
**Build cost**: roughly 1-2 days. Reference:
|
||||
`torchao.quantization.pt2e._numeric_debugger`.
|
||||
|
||||
### Extension 2: AST source injection at import time
|
||||
|
||||
**Granularity**: per-line (between any two Python statements).
|
||||
|
||||
**Mechanism**: an importlib loader hook rewrites Python source AST at module
|
||||
import time, inserting `if TRACE: dump(...)` statements. The decision is made
|
||||
once at import.
|
||||
|
||||
**Off semantics**: zero overhead. If the env var is off at import time, source
|
||||
is loaded as-is.
|
||||
|
||||
**When to add**: if you need per-line granularity that even FX-node-level can't
|
||||
provide. This is almost never the right choice.
|
||||
|
||||
**Build cost**: roughly 1 week. Brittle and hard to debug.
|
||||
|
||||
### Extension 3: `__torch_dispatch__` / `TorchDispatchMode`
|
||||
|
||||
**Granularity**: per-op (every dispatcher call: matmul, add, view, etc.).
|
||||
|
||||
**Mechanism**: a `TorchDispatchMode` context manager that intercepts all ops at
|
||||
the dispatcher level.
|
||||
|
||||
**Off semantics**: zero overhead. PyTorch's dispatcher only invokes mode hooks
|
||||
when a mode is active.
|
||||
|
||||
**When on**: significant overhead. Every op pays a Python callback cost. Triton
|
||||
kernels bypass it.
|
||||
|
||||
**When to add**: useful for quantization or dtype debugging where module-level
|
||||
granularity isn't enough.
|
||||
|
||||
**Build cost**: roughly 1 day. Reference:
|
||||
`torch.utils._python_dispatch.TorchDispatchMode`.
|
||||
|
||||
## Comparison with similar tools
|
||||
|
||||
| Tool | Pattern | FastVideo equivalent |
|
||||
|---|---|---|
|
||||
| SGLang `--debug-tensor-dump-output-folder` | env-gated forward hooks at startup | Extension 0 (this) |
|
||||
| TransformerEngine `DumpTensors` | config-driven selective dumps | Extension 0 (env-driven) |
|
||||
| HuggingFace `output_hidden_states=True` | source-level boolean gating | Not used; Extension 0 avoids model code edits |
|
||||
| torchao numeric debugger | FX pass + node-level loggers | Extension 1 (future) |
|
||||
| W&B `wandb.watch()` | runtime forward hooks (always on once registered) | Extension 0 has a similar mechanism, but gated off by default |
|
||||
|
||||
## Implementation references
|
||||
|
||||
- Module: `fastvideo/hooks/activation_trace.py`
|
||||
- Env vars: `fastvideo/envs.py` (`FASTVIDEO_TRACE_ACTIVATIONS` and friends)
|
||||
- Pipeline integration: `fastvideo/pipelines/composed_pipeline_base.py`
|
||||
- Tests: `fastvideo/tests/hooks/test_activation_trace.py`
|
||||
- Companion skill (for ad-hoc port investigations): `~/.config/opencode/skill/add-model-trace/`
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|---|---|
|
||||
| 2026-05-01 | Initial Extension 0 (module forward hooks) implementation. Extensions 1-3 designed but not implemented. |
|
||||
@@ -99,7 +99,7 @@ cd /FastVideo
|
||||
**Install the package**
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
uv pip install -e ".[dev]"
|
||||
```
|
||||
|
||||
The Docker image already includes Flash Attention and most heavy dependencies, so this is fast.
|
||||
|
||||
@@ -49,7 +49,7 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Install FastVideo in editable mode and set up hooks:
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
uv pip install -e ".[dev]"
|
||||
|
||||
# Optional: FlashAttention (builds native kernels)
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
@@ -306,6 +306,30 @@ surfaces:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
num_frames_per_block:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
audio_channels:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
audio_end_in_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
audio_start_in_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
max_audio_duration_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
sample_size:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
sampling_rate:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
compatibility_only:
|
||||
batch_size: "Gen3C inference-only tuning field pending typed batching design."
|
||||
gradient_checkpointing: "Gen3C inference-only compatibility field pending typed batching design."
|
||||
@@ -380,6 +404,13 @@ surfaces:
|
||||
ltx2_stg_scale_audio: request.extensions.ltx2.stg_scale_audio
|
||||
ltx2_stg_blocks_video: request.extensions.ltx2.stg_blocks_video
|
||||
ltx2_stg_blocks_audio: request.extensions.ltx2.stg_blocks_audio
|
||||
audio_start_in_s: request.extensions.stable_audio.audio_start_in_s
|
||||
audio_end_in_s: request.extensions.stable_audio.audio_end_in_s
|
||||
init_audio: request.extensions.stable_audio.init_audio
|
||||
init_audio_strength: request.extensions.stable_audio.init_audio_strength
|
||||
init_noise_level: request.extensions.stable_audio.init_noise_level
|
||||
inpaint_audio: request.extensions.stable_audio.inpaint_audio
|
||||
inpaint_mask: request.extensions.stable_audio.inpaint_mask
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
|
||||
|
||||
@@ -243,7 +243,7 @@ for step in range(start_step, max_steps):
|
||||
|
||||
```bash
|
||||
# Install
|
||||
uv pip install -e .[dev]
|
||||
uv pip install -e ".[dev]"
|
||||
|
||||
# Run DMD2 distillation on Wan 2.1
|
||||
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
|
||||
|
||||
@@ -27,7 +27,7 @@ uv pip install fastvideo
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### From source
|
||||
@@ -41,11 +41,11 @@ uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
Alternative with Conda environment:
|
||||
Alternative with Conda environment (still drives installs through `uv`):
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
@@ -58,14 +58,16 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install FlashAttention:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -87,7 +89,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
@@ -101,7 +103,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
|
||||
@@ -57,8 +57,10 @@ uv pip install fastvideo
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -80,7 +82,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
- Install MoGe:
|
||||
|
||||
```bash
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
```
|
||||
|
||||
- If you hit `ImportError: libGL.so.1` (common on Ubuntu/headless nodes), you can try installing OpenCV runtime libs:
|
||||
|
||||
@@ -54,7 +54,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
We recommend always installing [Flash Attention 2](https://github.com/Dao-AILab/flash-attention):
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
uv pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://github.com/Dao-AILab/flash-attention?tab=readme-ov-file#flashattention-3-beta-release) by compiling it from source (takes about 10 minutes for me):
|
||||
@@ -63,7 +63,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
|
||||
|
||||
cd hopper
|
||||
pip install ninja
|
||||
uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
@@ -98,7 +98,7 @@ To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please co
|
||||
```bash
|
||||
git clone https://github.com/thu-ml/SageAttention.git
|
||||
cd sageattention
|
||||
python setup.py install # or pip install -e .
|
||||
python setup.py install # or uv pip install -e .
|
||||
```
|
||||
|
||||
### Sage Attention 3
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
uv pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
uv pip install vsa
|
||||
```
|
||||
|
||||
### Data-free Distillation
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
uv pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -7,7 +7,7 @@ and the GEN3C diffusion model.
|
||||
|
||||
Requirements:
|
||||
1. Install MoGe:
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
If you hit `ImportError: libGL.so.1`, install:
|
||||
sudo apt-get update && sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1
|
||||
2. Download and convert weights:
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal user-runnable example for the daVinci-MagiHuman base AV pipeline.
|
||||
|
||||
Produces an mp4 with both video (Wan 2.2 TI2V-5B VAE) and audio (Stable
|
||||
Audio Open 1.0 VAE, first-class FastVideo port in
|
||||
`fastvideo/models/vaes/oobleck.py`) muxed together via PyAV.
|
||||
|
||||
Prerequisites (one-off):
|
||||
|
||||
# Accept terms of use on the gated HF repos with your HF_TOKEN:
|
||||
# - https://huggingface.co/google/t5gemma-9b-9b-ul2
|
||||
# - https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
# All four cross-variant shared components (Wan 2.2 VAE, T5-Gemma
|
||||
# encoder + tokenizer, Stable Audio VAE) are lazy-loaded from their
|
||||
# canonical upstream HF repos on first build, so a single ~25 GB
|
||||
# cache is shared across every MagiHuman variant.
|
||||
|
||||
The umbrella HF repo `FastVideo/MagiHuman-Diffusers` holds all four
|
||||
variants (base / distill / sr_540p / sr_1080p) under sibling subfolders
|
||||
and FastVideo will download just the requested subfolder. Local
|
||||
conversion via `scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py`
|
||||
is also supported.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/base",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_video/magi_human_basic/output_magi_human.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Defaults pulled from the registered preset (magi_human_base):
|
||||
# height=256, width=448, fps=25, num_inference_steps=32, seed=42.
|
||||
# Override here only if you have a specific QA scenario.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal user-runnable example for the daVinci-MagiHuman DMD-2 distilled
|
||||
text-to-AV pipeline.
|
||||
|
||||
Same arch as the base model (`basic_magi_human.py`) but with DMD-2 distilled
|
||||
weights: 8 denoising steps, no classifier-free guidance. ~4x faster than
|
||||
base at the same 256x480 resolution. Mirrors upstream
|
||||
`daVinci-MagiHuman/example/distill/run_T2V.sh`.
|
||||
|
||||
Prerequisites (one-off):
|
||||
|
||||
# 1) Accept terms on the gated HF repos with your HF_TOKEN:
|
||||
# - https://huggingface.co/google/t5gemma-9b-9b-ul2
|
||||
# - https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
# Cross-variant shared components (Wan 2.2 VAE + T5-Gemma + Stable
|
||||
# Audio VAE) are lazy-loaded from their canonical upstream HF repos
|
||||
# and shared with the base variant cache.
|
||||
# 2) Convert the distill subfolder of GAIR/daVinci-MagiHuman:
|
||||
python scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py \\
|
||||
--source GAIR/daVinci-MagiHuman \\
|
||||
--subfolder distill \\
|
||||
--output converted_weights/magi_human_distill \\
|
||||
--cast-bf16
|
||||
# `--cast-bf16` is recommended (61 GB fp32 -> 30 GB bf16); the FV pipeline
|
||||
# loads bf16 anyway, and the conversion keeps norms / RoPE bands fp32.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/distill",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_video/magi_human_basic/output_magi_human_distill.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Defaults pulled from the registered preset (magi_human_distill):
|
||||
# height=256, width=480, fps=25, num_inference_steps=8, cfg=1, seed=42.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal daVinci-MagiHuman DMD-2 distilled text+image-to-AV example."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanDistillI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/distill",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanI2VPipeline",
|
||||
pipeline_config=MagiHumanDistillI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_distill_ti2v/output_magi_human_distill_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-1080p text-to-AV in FastVideo.
|
||||
|
||||
Build the converted repo on large local storage, then symlink it into the
|
||||
workspace:
|
||||
|
||||
python scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py \
|
||||
--source GAIR/daVinci-MagiHuman \
|
||||
--subfolder base \
|
||||
--sr-source GAIR/daVinci-MagiHuman \
|
||||
--sr-subfolder 1080p_sr \
|
||||
--output /raid/william5lin_converted_weights/magi_human_sr_1080p \
|
||||
--cast-bf16
|
||||
ln -s /raid/william5lin_converted_weights/magi_human_sr_1080p \
|
||||
converted_weights/magi_human_sr_1080p
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanSR1080pConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_1080p",
|
||||
num_gpus=1,
|
||||
override_pipeline_cls_name="MagiHumanSR1080pPipeline",
|
||||
pipeline_config=MagiHumanSR1080pConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_video/magi_human_sr1080p/output_magi_human_sr1080p.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-1080p text+image-to-AV in FastVideo."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanSR1080pI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_1080p",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanSR1080pI2VPipeline",
|
||||
pipeline_config=MagiHumanSR1080pI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_sr1080p_ti2v/output_magi_human_sr1080p_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-540p text-to-AV in FastVideo.
|
||||
|
||||
The converted repo must contain both ``transformer/`` (base DiT) and
|
||||
``sr_transformer/`` (540p SR DiT). Build it with:
|
||||
|
||||
python scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py \
|
||||
--source GAIR/daVinci-MagiHuman \
|
||||
--subfolder base \
|
||||
--sr-source GAIR/daVinci-MagiHuman \
|
||||
--sr-subfolder 540p_sr \
|
||||
--output converted_weights/magi_human_sr_540p
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm afternoon scene: a person sits on a park bench reading a book, "
|
||||
"surrounded by softly swaying trees."
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_540p",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_video/magi_human_sr540p/output_magi_human_sr540p.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run daVinci-MagiHuman SR-540p text+image-to-AV in FastVideo."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanSR540pI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/sr_540p",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanSRI2VPipeline",
|
||||
pipeline_config=MagiHumanSR540pI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_sr540p_ti2v/output_magi_human_sr540p_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,34 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Minimal daVinci-MagiHuman base text+image-to-AV example."""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines.basic.magi_human.pipeline_configs import (
|
||||
MagiHumanBaseI2VConfig,
|
||||
)
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A cheerful saxophonist performs a short line with expressive facial "
|
||||
"motion, natural head movement, and synchronized audio in a small jazz club."
|
||||
)
|
||||
IMAGE_PATH = "assets/images/saxophonist.jpg"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/MagiHuman-Diffusers/base",
|
||||
num_gpus=1,
|
||||
workload_type="i2v",
|
||||
override_pipeline_cls_name="MagiHumanI2VPipeline",
|
||||
pipeline_config=MagiHumanBaseI2VConfig(),
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path="outputs_video/magi_human_ti2v/output_magi_human_ti2v.mp4",
|
||||
save_video=True,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -48,7 +48,7 @@ Prerequisites:
|
||||
and export your HF token in the shell:
|
||||
export HF_TOKEN=hf_...
|
||||
2. Install optional inference deps (one-time):
|
||||
pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
uv pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
cmake_minimum_required(VERSION 3.26 FATAL_ERROR)
|
||||
project(fastvideo-kernel LANGUAGES CXX)
|
||||
|
||||
# Prefer environment variable (used by CI or pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
# Prefer environment variable (used by CI or uv pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
|
||||
@@ -11,7 +11,7 @@ except ImportError:
|
||||
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError(
|
||||
"flash-attn is not installed. Please install it, e.g., `pip install flash-attn`."
|
||||
"flash-attn is not installed. Please install it, e.g., `uv pip install flash-attn`."
|
||||
)
|
||||
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
# `fastvideo/` — Core Package
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Inference + training framework for video DiTs. Public API entry: `from fastvideo import VideoGenerator, PipelineConfig, SamplingParam`.
|
||||
|
||||
## Public Surface (`__init__.py`)
|
||||
|
||||
```python
|
||||
VideoGenerator # entrypoints/video_generator.py — high-level inference handle
|
||||
PipelineConfig # configs/pipelines/base.py — pipeline wiring dataclass
|
||||
SamplingParam # api/sampling_param.py — runtime sampling knobs
|
||||
```
|
||||
|
||||
CLI entry: `fastvideo` script → `entrypoints/cli/main.py` (subcommands: `generate`, `serve`, `bench`).
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
fastvideo/
|
||||
├── api/ # Schema + presets for the OpenAI-compatible serving layer
|
||||
├── attention/ # Backends + selector (FlashAttn / SageAttn / SDPA / VSA / VMoBA / SLA)
|
||||
├── configs/ # Per-model arch configs + per-pipeline configs (registry-driven)
|
||||
├── dataset/ # Dataloaders (pre-commit excluded — minimal lint surface)
|
||||
├── distributed/ # SP/TP groups, device communicators, init helpers
|
||||
├── entrypoints/ # cli/, openai/, streaming/, video_generator.py
|
||||
├── hooks/ # Runtime hook system for pipelines
|
||||
├── layers/ # Tensor-parallel linears + attention wrappers (port targets)
|
||||
├── models/ # DiT / VAE / encoder / scheduler / loader (pre-commit excluded)
|
||||
├── pipelines/ # basic/<model>/, preprocess/, stages/, training/
|
||||
├── platforms/ # CUDA/ROCm capability + AttentionBackendEnum
|
||||
├── third_party/ # Vendored externals (lint excluded; do not reformat)
|
||||
├── train/ # NEW modular trainer — methods × models × callbacks
|
||||
├── training/ # LEGACY monolithic *_training/distillation_pipeline.py
|
||||
├── worker/ # Multi-process / Ray executors
|
||||
├── workflow/ # Preprocessing workflow base class
|
||||
├── registry.py # Pipeline-config + model-class lookup (canonical)
|
||||
├── envs.py # Env-var declarations
|
||||
├── fastvideo_args.py# Runtime arg dataclass passed through pipelines
|
||||
└── utils.py # FlexibleArgumentParser, qualname resolver, etc.
|
||||
```
|
||||
|
||||
## Where to Look
|
||||
|
||||
| Task | Location |
|
||||
|------|----------|
|
||||
| Add a new pipeline class | `pipelines/basic/<model>/` + `configs/pipelines/<model>.py` + register in `registry.py` |
|
||||
| Add a new model component | `models/<role>/<model>.py` + `configs/models/<role>/<model>.py` |
|
||||
| Wire an existing model into a new pipeline | `pipelines/basic/<model>/presets.py` + reuse stages from `pipelines/stages/` |
|
||||
| Add a converter | `scripts/checkpoint_conversion/<model>_to_*.py` (separate dir, separate AGENTS.md) |
|
||||
| Add an attention backend | `attention/backends/<name>.py` + register in selector |
|
||||
| Add a runtime CLI flag | `fastvideo_args.py` (avoid `argparse` ad-hoc inside stages) |
|
||||
|
||||
## Conventions Specific Here
|
||||
|
||||
- `PipelineStage` subclasses (`pipelines/stages/`) own one verb each (encode, schedule, denoise, decode). Compose, don't fork.
|
||||
- Every pipeline reads from a `PipelineConfig` subclass and a `SamplingParam`. Never read raw env vars inside a stage — go through `fastvideo.envs`.
|
||||
- Logger setup: `from fastvideo.logger import init_logger; logger = init_logger(__name__)`. Do not call `logging.getLogger` directly.
|
||||
- Imports between `train/` and `training/` are **forbidden** — they are independent stacks.
|
||||
|
||||
## Pre-Commit Exclusions (do not assume linted)
|
||||
|
||||
These dirs are listed in `.pre-commit-config.yaml` `exclude`:
|
||||
|
||||
- `fastvideo/third_party/`, `fastvideo/dataset/`, `fastvideo/models/`
|
||||
|
||||
Editing files there will NOT trigger yapf/ruff/mypy/codespell. Format manually if a sibling file shows clear style; do not introduce new violations.
|
||||
@@ -0,0 +1,58 @@
|
||||
# `fastvideo/attention/` — Attention Backends
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Backend registry + selector wrapping FlashAttn / SageAttn / SageAttn3 / SDPA / VSA / VMoBA / SLA / BSA.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
attention/
|
||||
├── __init__.py # Exports DistributedAttention, LocalAttention, get_attn_backend
|
||||
├── layer.py # DistributedAttention, DistributedAttention_VSA, LocalAttention
|
||||
├── selector.py # get_attn_backend (cached) + env-var override
|
||||
├── backends/
|
||||
│ ├── abstract.py # AttentionBackend / AttentionMetadata / AttentionMetadataBuilder
|
||||
│ ├── flash_attn.py # FA2/FA3
|
||||
│ ├── sage_attn.py # SageAttention v1
|
||||
│ ├── sage_attn3.py # SageAttention v3
|
||||
│ ├── sdpa.py # torch SDPA fallback
|
||||
│ ├── video_sparse_attn.py # VSA (paper: Video Sparse Attention)
|
||||
│ ├── vmoba.py # Video-MoBA
|
||||
│ ├── sla.py # Sliding-window (STA)
|
||||
│ └── bsa_attn.py # Block-sparse
|
||||
└── utils/
|
||||
├── flash_attn_cute.py
|
||||
└── flash_attn_no_pad.py
|
||||
```
|
||||
|
||||
## Selection Order
|
||||
|
||||
`get_attn_backend()` resolves via:
|
||||
|
||||
1. Env-var override `FASTVIDEO_ATTENTION_BACKEND` (see `STR_BACKEND_ENV_VAR` in `fastvideo/utils.py`).
|
||||
2. Per-platform default from `fastvideo/platforms/`.
|
||||
3. Heuristic fallback to SDPA.
|
||||
|
||||
The result is `@lru_cache`d. Tests that need a specific backend must use the
|
||||
`global_force_attn_backend(...)` context manager from `selector.py`, never set
|
||||
the env var mid-process.
|
||||
|
||||
## Adding a Backend
|
||||
|
||||
1. Subclass `AttentionBackend` in `backends/<name>.py`.
|
||||
2. Implement `AttentionMetadata` + `AttentionMetadataBuilder` for the new path.
|
||||
3. Register the enum value in `fastvideo/platforms/interface.py` (`AttentionBackendEnum`).
|
||||
4. Wire string → class resolution in `selector.py`.
|
||||
5. Verify the new backend works with `DistributedAttention` (sequence parallel)
|
||||
and `LocalAttention` (single-rank). If it cannot support SP, document the
|
||||
gap in the backend file's module docstring.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Calling `torch.nn.functional.scaled_dot_product_attention` directly inside a
|
||||
model's forward — go through `DistributedAttention` / `LocalAttention`.
|
||||
- Reading `os.environ[STR_BACKEND_ENV_VAR]` from arbitrary call sites. Use
|
||||
`get_env_variable_attn_backend()`.
|
||||
- Caching backend instances per-module. The selector cache is process-wide; do
|
||||
not duplicate it.
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
@@ -18,6 +17,7 @@ except ImportError:
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
|
||||
|
||||
@@ -405,7 +405,7 @@ class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
|
||||
if not SAGESLA_ENABLED:
|
||||
raise ImportError("SageSLA requires spas_sage_attn. "
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git")
|
||||
"Install with: uv pip install git+https://github.com/thu-ml/SpargeAttn.git")
|
||||
|
||||
assert head_size in [64, 128], f"SageSLA requires head_size in [64, 128], got {head_size}"
|
||||
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
# `fastvideo/configs/` — Config-Driven Model Registry
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Two layers of dataclass configs feed every pipeline: **arch configs** (what the model is) and **pipeline configs** (how to run it).
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
configs/
|
||||
├── configs.py # Dataset / loader enums (DatasetType, VideoLoaderType)
|
||||
├── utils.py # update_config_from_args, shallow_asdict helpers
|
||||
├── backend/ # Attention backend defaults
|
||||
├── models/
|
||||
│ ├── base.py # ModelConfig ABC
|
||||
│ ├── dits/ # DiTConfig per model (wanvideo, ltx2, hunyuan, ...)
|
||||
│ ├── vaes/ # VAEConfig per model
|
||||
│ ├── encoders/ # EncoderConfig (t5, clip, llama, qwen2_5, gemma, siglip, ...)
|
||||
│ ├── upsamplers/ # UpsamplerConfig (hunyuan15)
|
||||
│ └── audio/ # Audio-model configs (ltx2_audio_vae, ...)
|
||||
├── pipelines/
|
||||
│ ├── base.py # PipelineConfig ABC + (de)serialization
|
||||
│ └── <model>.py # Concrete configs (HunyuanConfig, WanT2V480PConfig, ...)
|
||||
└── *.json # Frozen reference configs for shipped models
|
||||
```
|
||||
|
||||
## How Configs Hook Into the Registry
|
||||
|
||||
`fastvideo/registry.py` imports every concrete `PipelineConfig` and exposes
|
||||
`get_pipeline_config_cls_from_name(...)`. Adding a new pipeline config requires:
|
||||
|
||||
1. Subclass `PipelineConfig` in `pipelines/<model>.py`.
|
||||
2. Reference its component arch configs (DiT / VAE / encoder / upsampler).
|
||||
3. Add the import + name mapping in `fastvideo/registry.py`.
|
||||
|
||||
Configs that do not appear in `registry.py` are unreachable from `VideoGenerator`.
|
||||
|
||||
## Arch vs Pipeline — Where Does This Field Go?
|
||||
|
||||
| Field type | Lives on |
|
||||
|-----------|----------|
|
||||
| Architecture constants (hidden dim, num heads, layer count) | `configs/models/<role>/<model>.py` |
|
||||
| Default sampling params (steps, cfg, shift, fps) | `configs/pipelines/<model>.py` |
|
||||
| Runtime overrides (precision, sp_size, tp_size, attention backend) | `configs/pipelines/base.py` defaults + CLI flags via `fastvideo_args.py` |
|
||||
| `param_names_mapping` for HF → FastVideo state-dict | Arch config (lives with the model definition) |
|
||||
|
||||
If a knob is tunable per inference call → `SamplingParam`, not `PipelineConfig`.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Hard-coding architecture constants inside model classes — always read from the arch config.
|
||||
- Using `argparse` directly here. Configs deserialize from dicts via `update_config_from_args`.
|
||||
- Importing from `fastvideo.pipelines` here. Configs are the lower layer; the dependency is one-way.
|
||||
@@ -5,6 +5,7 @@ from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.magi_human import MagiHumanVideoConfig
|
||||
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
|
||||
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
@@ -13,5 +14,5 @@ from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"StableAudioConfig"
|
||||
"MagiHumanVideoConfig", "StableAudioConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Architecture / model config for the daVinci-MagiHuman DiT.
|
||||
|
||||
The MagiHuman base DiT is a 15B-parameter single-stream transformer that
|
||||
jointly denoises video, audio, and text tokens in one flat sequence. Layout
|
||||
details verified against GAIR/daVinci-MagiHuman's base/ shards (2026-04-24).
|
||||
|
||||
This file captures only configuration. The module implementation lives in
|
||||
fastvideo/models/dits/magi_human.py and the pipeline wiring in
|
||||
fastvideo/pipelines/basic/magi_human/.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def _is_block_layer(n: str, m) -> bool:
|
||||
# Match "block.layers.<idx>" — the FSDP shard boundary for MagiHuman.
|
||||
parts = n.split(".")
|
||||
return (len(parts) >= 3 and parts[0] == "block" and parts[1] == "layers" and str.isdigit(parts[2]))
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanArchConfig(DiTArchConfig):
|
||||
"""MagiHuman base DiT architecture constants.
|
||||
|
||||
**Scope contract:** fields here must match the `transformer/config.json`
|
||||
emitted by `scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py`
|
||||
1:1, and both are sourced from the upstream Python reference
|
||||
`inference/common/config.py::ModelConfig` (the HF root `config.json`
|
||||
is empty so the Python source is canonical). Pipeline-level knobs
|
||||
(VAE stride, fps, num_inference_steps, CFG scales, flow_shift,
|
||||
t5_gemma_target_length) and data-proxy knobs (coords_style,
|
||||
frame_receptive_field, ref_audio_offset, text_offset) live on
|
||||
`MagiHumanBaseConfig`, NOT here.
|
||||
|
||||
`param_names_mapping` is intentionally empty: the FastVideo implementation
|
||||
keeps the same module tree as the reference (`adapter.*`,
|
||||
`block.layers.<i>.*`, `final_linear_{video,audio}.*`,
|
||||
`final_norm_{video,audio}.*`), so converted weights load directly.
|
||||
"""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_block_layer])
|
||||
|
||||
# No renames needed — the FastVideo module mirrors the reference names.
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
# --- transformer shape ---
|
||||
num_layers: int = 40
|
||||
hidden_size: int = 5120
|
||||
head_dim: int = 128
|
||||
num_query_groups: int = 8 # num_heads_kv (GQA)
|
||||
|
||||
# --- modality channels ---
|
||||
# video_in_channels = z_dim (48) * patch_size product (1*2*2=4), so the
|
||||
# embedder receives 192 per token. text_in_channels is T5Gemma-9B's
|
||||
# encoder hidden size.
|
||||
video_in_channels: int = 192
|
||||
audio_in_channels: int = 64
|
||||
text_in_channels: int = 3584
|
||||
|
||||
# --- block-level architecture switches ---
|
||||
# Sandwich MoE: first and last 4 layers have per-modality experts
|
||||
# (video/audio/text), middle layers share a single set of weights.
|
||||
mm_layers: tuple[int, ...] = (0, 1, 2, 3, 36, 37, 38, 39)
|
||||
local_attn_layers: tuple[int, ...] = ()
|
||||
gelu7_layers: tuple[int, ...] = (0, 1, 2, 3)
|
||||
post_norm_layers: tuple[int, ...] = ()
|
||||
enable_attn_gating: bool = True
|
||||
activation_type: str = "swiglu7"
|
||||
|
||||
# --- DiT patching (upstream `ModelConfig`-equivalent; NOT the VAE
|
||||
# stride, which is pipeline-level). ---
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
spatial_rope_interpolation: str = "extra"
|
||||
|
||||
# --- TReAD (token routing + early drop). Flattened from the upstream
|
||||
# nested `tread_config` dict so it round-trips through
|
||||
# `update_model_arch` cleanly. ---
|
||||
tread_selection_rate: float = 0.5
|
||||
tread_start_layer_idx: int = 2
|
||||
tread_end_layer_idx: int = 25
|
||||
|
||||
# --- derived fields (populated in __post_init__) ---
|
||||
num_attention_heads: int = 0 # hidden_size / head_dim
|
||||
num_heads_kv: int = 0 # == num_query_groups
|
||||
in_channels: int = 0 # mirror of video_in_channels (FastVideo contract)
|
||||
out_channels: int = 0 # mirror of video_in_channels
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.num_attention_heads = self.hidden_size // self.head_dim
|
||||
self.num_heads_kv = self.num_query_groups
|
||||
self.in_channels = self.video_in_channels
|
||||
self.out_channels = self.video_in_channels
|
||||
# num_channels_latents is the VAE latent z_dim (48 for Wan 2.2 TI2V-5B).
|
||||
# We don't declare z_dim on the arch config (it's a VAE property),
|
||||
# but we still set num_channels_latents for the BaseDiT contract.
|
||||
self.num_channels_latents = 48
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=MagiHumanArchConfig)
|
||||
|
||||
prefix: str = "magi_human"
|
||||
@@ -9,10 +9,11 @@ from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1
|
||||
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
|
||||
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
|
||||
StableAudioConditionerConfig)
|
||||
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
|
||||
|
||||
__all__ = [
|
||||
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
|
||||
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
|
||||
"StableAudioConditionerConfig"
|
||||
"StableAudioConditionerConfig", "T5GemmaEncoderConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the T5-Gemma encoder used by daVinci-MagiHuman.
|
||||
|
||||
The reference pipeline uses `transformers.models.t5gemma.T5GemmaEncoderModel`
|
||||
on `google/t5gemma-9b-9b-ul2`. That is a gated Google repository, so the
|
||||
encoder weights are not bundled inside GAIR/daVinci-MagiHuman; they are
|
||||
loaded from the T5-Gemma HF repo directly.
|
||||
|
||||
Encoder shape (verified from google/t5gemma-9b-9b-ul2/config.json):
|
||||
layers=42, hidden=3584, heads=16, kv_heads=8, head_dim=256,
|
||||
intermediate=14336, rope_theta=10000.0, max_pos=8192,
|
||||
layer_types alternate sliding_attention / full_attention.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.encoders.base import (
|
||||
TextEncoderArchConfig,
|
||||
TextEncoderConfig,
|
||||
)
|
||||
|
||||
|
||||
def _is_t5gemma_model(n: str, m) -> bool:
|
||||
return n.endswith("t5gemma_model") or n.endswith("_t5gemma_model")
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5GemmaEncoderArchConfig(TextEncoderArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["T5GemmaEncoderModel"])
|
||||
|
||||
hidden_size: int = 3584
|
||||
num_hidden_layers: int = 42
|
||||
num_attention_heads: int = 16
|
||||
num_key_value_heads: int = 8
|
||||
head_dim: int = 256
|
||||
intermediate_size: int = 14336
|
||||
max_position_embeddings: int = 8192
|
||||
rope_theta: float = 10000.0
|
||||
vocab_size: int = 256000
|
||||
|
||||
# MagiHuman fixes prompt embed length at 640 via pad_or_trim.
|
||||
text_len: int = 640
|
||||
|
||||
pad_token_id: int = 0
|
||||
eos_token_id: int = 1
|
||||
|
||||
# Path to the upstream gated repo. When set, the FastVideo loader will
|
||||
# pull the encoder directly via `T5GemmaEncoderModel.from_pretrained`.
|
||||
t5gemma_model_path: str = "google/t5gemma-9b-9b-ul2"
|
||||
t5gemma_dtype: str = "bfloat16"
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_t5gemma_model])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# WHY: upstream `t5_gemma_model.py:25` tokenizes without
|
||||
# padding/max_length, then `prompt_process.py` pad_or_trim-s the
|
||||
# encoded states. Keep only tensor return here so
|
||||
# MagiHumanLatentPreparationStage can pad/trim post-encode while
|
||||
# preserving the real original prompt length.
|
||||
self.tokenizer_kwargs.pop("truncation", None)
|
||||
self.tokenizer_kwargs.pop("max_length", None)
|
||||
self.tokenizer_kwargs.pop("padding", None)
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5GemmaEncoderConfig(TextEncoderConfig):
|
||||
arch_config: TextEncoderArchConfig = field(default_factory=T5GemmaEncoderArchConfig)
|
||||
|
||||
prefix: str = "t5gemma"
|
||||
@@ -3,6 +3,8 @@
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
from fastvideo.entrypoints.cli.router_serve import (
|
||||
cmd_init as router_serve_cmd_init, )
|
||||
from fastvideo.entrypoints.cli.serve import cmd_init as serve_cmd_init
|
||||
from fastvideo.entrypoints.cli.bench import cmd_init as bench_cmd_init
|
||||
|
||||
@@ -12,6 +14,7 @@ def cmd_init() -> list[CLISubcommand]:
|
||||
commands = []
|
||||
commands.extend(generate_cmd_init())
|
||||
commands.extend(serve_cmd_init())
|
||||
commands.extend(router_serve_cmd_init())
|
||||
commands.extend(bench_cmd_init())
|
||||
return commands
|
||||
|
||||
|
||||
@@ -0,0 +1,115 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""``fastvideo router-serve`` CLI subcommand.
|
||||
|
||||
Launches the streaming router from a YAML config. Separate from
|
||||
``fastvideo serve`` because the router is an orthogonal process: it
|
||||
fronts one or more running servers rather than hosting a generator
|
||||
itself.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.api.parser import load_raw_config
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.streaming.router.config import (
|
||||
ReplicaEndpoint,
|
||||
RouterConfig,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class RouterServeSubcommand(CLISubcommand):
|
||||
"""Start the multi-replica WebSocket router."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.name = "router-serve"
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
config = _load_router_config(args.config)
|
||||
logger.info(
|
||||
"router listening on %s:%d (%d replicas, %d primary)",
|
||||
config.host,
|
||||
config.port,
|
||||
len(config.replicas),
|
||||
sum(1 for r in config.replicas if r.primary),
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.router.main import run_router
|
||||
|
||||
run_router(config)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo router-serve requires --config PATH")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Router config file not found: {args.config}")
|
||||
|
||||
def subparser_init(
|
||||
self,
|
||||
subparsers: argparse._SubParsersAction,
|
||||
) -> FlexibleArgumentParser:
|
||||
parser = subparsers.add_parser(
|
||||
"router-serve",
|
||||
help="Start the streaming router (multi-replica load balancer)",
|
||||
usage="fastvideo router-serve --config ROUTER_CONFIG",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="Path to a YAML/JSON router config. Required.",
|
||||
)
|
||||
return cast(FlexibleArgumentParser, parser)
|
||||
|
||||
|
||||
def _load_router_config(path: str) -> RouterConfig:
|
||||
raw = load_raw_config(path)
|
||||
router_raw = raw.get("router") if isinstance(raw, dict) else None
|
||||
if not isinstance(router_raw, dict):
|
||||
raise ValueError(f"Router config {path!r} must have a top-level `router:` block")
|
||||
|
||||
replicas_raw = router_raw.get("replicas", [])
|
||||
if not isinstance(replicas_raw, list):
|
||||
raise ValueError(f"router.replicas must be a list, got {type(replicas_raw).__name__}")
|
||||
replicas = []
|
||||
for i, r in enumerate(replicas_raw):
|
||||
if not isinstance(r, dict):
|
||||
raise ValueError(f"router.replicas[{i}] must be a mapping, got {type(r).__name__}")
|
||||
url = r.get("url")
|
||||
if not url:
|
||||
raise ValueError(f"router.replicas[{i}] is missing required key 'url'")
|
||||
replicas.append(
|
||||
ReplicaEndpoint(
|
||||
url=url,
|
||||
name=r.get("name"),
|
||||
primary=bool(r.get("primary", False)),
|
||||
weight=float(r.get("weight", 1.0)),
|
||||
))
|
||||
if not replicas:
|
||||
raise ValueError("Router config must list at least one replica under `router.replicas`")
|
||||
|
||||
health_check = router_raw.get("health_check") or {}
|
||||
return RouterConfig(
|
||||
host=str(router_raw.get("host", "0.0.0.0")),
|
||||
port=int(router_raw.get("port", 9000)),
|
||||
replicas=replicas,
|
||||
health_check_path=str(health_check.get("path", "/health")),
|
||||
health_check_interval_seconds=float(health_check.get("interval_seconds", 5.0)),
|
||||
health_check_timeout_seconds=float(health_check.get("timeout_seconds", 2.0)),
|
||||
failure_threshold=int(health_check.get("failure_threshold", 3)),
|
||||
recovery_threshold=int(health_check.get("recovery_threshold", 2)),
|
||||
)
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
return [RouterServeSubcommand()]
|
||||
|
||||
|
||||
__all__ = ["RouterServeSubcommand", "cmd_init"]
|
||||
@@ -11,6 +11,28 @@ from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.gpu_pool import (
|
||||
GpuPool,
|
||||
InProcessGpuPool,
|
||||
PoolAcquireTimeout,
|
||||
SubprocessGpuPool,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.mock_server import (
|
||||
MockGenerator,
|
||||
build_mock_app,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt import (
|
||||
LLMProvider,
|
||||
PromptEnhancer,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.safety import (
|
||||
PromptSafetyFilter,
|
||||
SafetyDecision,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_logger import (
|
||||
SessionLogEvent,
|
||||
SessionLogger,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
@@ -20,12 +42,24 @@ __all__ = [
|
||||
"BlobStore",
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
"GpuPool",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"InProcessGpuPool",
|
||||
"LLMProvider",
|
||||
"MockGenerator",
|
||||
"PoolAcquireTimeout",
|
||||
"PromptEnhancer",
|
||||
"PromptSafetyFilter",
|
||||
"SafetyDecision",
|
||||
"SessionLogEvent",
|
||||
"SessionLogger",
|
||||
"build_mock_app",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionState",
|
||||
"SessionStore",
|
||||
"SubprocessGpuPool",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,542 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GPU pool manager for the streaming server.
|
||||
|
||||
Replaces the single-generator path in PR 7.5 with a typed pool
|
||||
abstraction. Three implementations ship here:
|
||||
|
||||
* :class:`InProcessGpuPool` — one in-process ``VideoGenerator``; used
|
||||
by tests and single-GPU dev deployments.
|
||||
* :class:`SubprocessGpuPool` — one ``multiprocessing.Process`` per
|
||||
GPU, each running :func:`worker_main` against a ``GeneratorConfig``.
|
||||
Jobs are dispatched via ``multiprocessing.Queue``.
|
||||
* :class:`GpuPool` (abstract) — the interface both use.
|
||||
|
||||
Session-to-GPU binding lives in the pool so continuation state stays
|
||||
on the GPU that generated the previous segment (matching the internal
|
||||
``gpu_pool.py``'s per-GPU cache behavior). Cross-GPU handoff is
|
||||
supported via :class:`SessionStore` snapshot + hydrate, which
|
||||
serializes the state before the migration and rehydrates it on the
|
||||
new worker.
|
||||
|
||||
Typed config: workers start from a :class:`GeneratorConfig` (no flat
|
||||
LTX-2 kwargs), satisfying the PR 6 + PR 7 contracts that the public
|
||||
surface doesn't reintroduce the legacy kwarg bag.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import multiprocessing as mp
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import Future
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
GeneratorConfig,
|
||||
GenerationRequest,
|
||||
GpuPoolConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.worker import worker_main
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public interface
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _GeneratorLike(Protocol):
|
||||
"""Subset the pool calls on a worker-side generator."""
|
||||
|
||||
def generate(self, request: GenerationRequest) -> Any:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class PoolAssignment:
|
||||
"""The worker a session is currently bound to."""
|
||||
|
||||
gpu_id: int
|
||||
worker_id: str
|
||||
pinned_at: float = field(default_factory=time.monotonic)
|
||||
|
||||
|
||||
class GpuPool(ABC):
|
||||
"""Abstract GPU pool.
|
||||
|
||||
``acquire`` binds a session to a worker and holds that binding
|
||||
across segments so continuation state can stay hot. ``run`` submits
|
||||
a single ``GenerationRequest`` for a bound session.
|
||||
|
||||
Acquire / release are independent of run — a session can run many
|
||||
segments on one acquired worker, and must release on disconnect.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def acquire(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float | None = None,
|
||||
) -> PoolAssignment:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def run(
|
||||
self,
|
||||
session_id: str,
|
||||
request: GenerationRequest,
|
||||
) -> Any:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def release(self, session_id: str) -> None:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def shutdown(self) -> None:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def health(self) -> PoolHealth:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class PoolHealth:
|
||||
total_workers: int
|
||||
available_workers: int
|
||||
active_sessions: int
|
||||
queued_sessions: int = 0
|
||||
|
||||
|
||||
class PoolAcquireTimeout(RuntimeError):
|
||||
"""Raised when ``acquire`` times out waiting for a free worker."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# In-process implementation (single-worker, test / dev)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class InProcessGpuPool(GpuPool):
|
||||
"""Single-process pool backed by one :class:`_GeneratorLike`.
|
||||
|
||||
This is what PR 7.5's server uses by default; PR 7.6 adds the real
|
||||
``SubprocessGpuPool`` alternative but keeps this one for tests and
|
||||
small deployments.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator: _GeneratorLike,
|
||||
*,
|
||||
gpu_id: int = 0,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> None:
|
||||
self._generator = generator
|
||||
self._gpu_id = gpu_id
|
||||
self._worker_id = f"inproc-{uuid.uuid4().hex[:6]}"
|
||||
self._session_store = session_store or InMemorySessionStore()
|
||||
self._active: dict[str, PoolAssignment] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
self._gen_lock = asyncio.Lock()
|
||||
|
||||
async def acquire(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float | None = None,
|
||||
) -> PoolAssignment:
|
||||
async with self._lock:
|
||||
existing = self._active.get(session_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
assignment = PoolAssignment(gpu_id=self._gpu_id, worker_id=self._worker_id)
|
||||
self._active[session_id] = assignment
|
||||
return assignment
|
||||
|
||||
async def run(
|
||||
self,
|
||||
session_id: str,
|
||||
request: GenerationRequest,
|
||||
) -> Any:
|
||||
if session_id not in self._active:
|
||||
raise RuntimeError(f"session {session_id!r} is not acquired on this pool")
|
||||
# Serialize generator access so one GPU runs one request at a
|
||||
# time, matching the internal gpu_pool's per-GPU lock.
|
||||
async with self._gen_lock:
|
||||
loop = asyncio.get_running_loop()
|
||||
return await loop.run_in_executor(None, self._generator.generate, request)
|
||||
|
||||
async def release(self, session_id: str) -> None:
|
||||
async with self._lock:
|
||||
self._active.pop(session_id, None)
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
self._active.clear()
|
||||
|
||||
def health(self) -> PoolHealth:
|
||||
return PoolHealth(
|
||||
total_workers=1,
|
||||
available_workers=1 if not self._active else 0,
|
||||
active_sessions=len(self._active),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Subprocess implementation (multi-worker, real deployment)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class _WorkerHandle:
|
||||
process: Any # mp.Process or compatible handle with is_alive / join / kill
|
||||
job_queue: mp.Queue
|
||||
result_queue: mp.Queue
|
||||
gpu_id: int
|
||||
worker_id: str
|
||||
ready: threading.Event
|
||||
# ``ready`` flips on either successful boot or boot failure so the
|
||||
# parent stops waiting; ``boot_ok`` is set only on a real ready
|
||||
# acknowledgement and is what gates pool admission.
|
||||
boot_ok: threading.Event
|
||||
shutdown_event: Any # mp.Event is a factory, not a type — Any keeps mypy sane
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PendingJob:
|
||||
job_id: str
|
||||
future: Future
|
||||
session_id: str
|
||||
worker_id: str
|
||||
|
||||
|
||||
class SubprocessGpuPool(GpuPool):
|
||||
"""One ``multiprocessing.Process`` per GPU.
|
||||
|
||||
Each worker boots :class:`fastvideo.VideoGenerator` from a typed
|
||||
:class:`GeneratorConfig` inside the child process (post-
|
||||
``CUDA_VISIBLE_DEVICES`` setup) and consumes jobs from an mp Queue.
|
||||
|
||||
This is the production shape: the parent process stays CPU-only, and
|
||||
GPU state never crosses process boundaries. Continuation state is
|
||||
serialized through :class:`SessionStore` for cross-GPU handoff.
|
||||
|
||||
PR 7.6 ships this as an opt-in; PR 7.5's in-process pool remains the
|
||||
default until nightly runs validate the subprocess path.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator_config: GeneratorConfig,
|
||||
*,
|
||||
pool_config: GpuPoolConfig,
|
||||
warmup_config: WarmupConfig | None = None,
|
||||
session_store: SessionStore | None = None,
|
||||
worker_factory: WorkerFactory | None = None,
|
||||
) -> None:
|
||||
self._generator_config = generator_config
|
||||
self._pool_config = pool_config
|
||||
self._warmup_config = warmup_config or WarmupConfig()
|
||||
self._session_store = session_store or InMemorySessionStore()
|
||||
self._worker_factory = worker_factory or _default_worker_factory
|
||||
self._workers: list[_WorkerHandle] = []
|
||||
self._available: asyncio.Queue[int] = asyncio.Queue()
|
||||
self._assignments: dict[str, PoolAssignment] = {}
|
||||
self._worker_by_id: dict[str, _WorkerHandle] = {}
|
||||
self._pending: dict[str, _PendingJob] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
self._result_reader_tasks: list[asyncio.Task] = []
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Spawn worker processes and wait for each to report ready."""
|
||||
num_workers = self._pool_config.num_workers or 1
|
||||
for gpu_id in range(num_workers):
|
||||
handle = self._worker_factory(
|
||||
gpu_id=gpu_id,
|
||||
generator_config=self._generator_config,
|
||||
warmup_config=self._warmup_config,
|
||||
)
|
||||
self._workers.append(handle)
|
||||
self._worker_by_id[handle.worker_id] = handle
|
||||
|
||||
# Wait for each worker's ready event in a thread to avoid
|
||||
# blocking the event loop.
|
||||
loop = asyncio.get_running_loop()
|
||||
await asyncio.gather(*[
|
||||
loop.run_in_executor(None, handle.ready.wait, self._warmup_config.timeout_seconds)
|
||||
for handle in self._workers
|
||||
])
|
||||
|
||||
# Start background result readers — one task per worker
|
||||
# drains its result queue and resolves futures in _pending.
|
||||
for handle in self._workers:
|
||||
task = asyncio.create_task(self._drain_results(handle))
|
||||
self._result_reader_tasks.append(task)
|
||||
|
||||
# Only admit workers that successfully booted. Anything that
|
||||
# failed boot (timeout, crash, error sentinel) stays out of the
|
||||
# available queue so we never assign a session to it.
|
||||
for idx, handle in enumerate(self._workers):
|
||||
if handle.boot_ok.is_set():
|
||||
await self._available.put(idx)
|
||||
else:
|
||||
logger.error(
|
||||
"pool: worker %s failed to boot; skipping",
|
||||
handle.worker_id,
|
||||
)
|
||||
|
||||
async def acquire(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float | None = None,
|
||||
) -> PoolAssignment:
|
||||
async with self._lock:
|
||||
existing = self._assignments.get(session_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
try:
|
||||
idx = await asyncio.wait_for(self._available.get(), timeout=timeout)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise PoolAcquireTimeout(f"no worker available after {timeout}s") from exc
|
||||
handle = self._workers[idx]
|
||||
assignment = PoolAssignment(gpu_id=handle.gpu_id, worker_id=handle.worker_id)
|
||||
async with self._lock:
|
||||
self._assignments[session_id] = assignment
|
||||
return assignment
|
||||
|
||||
async def run(
|
||||
self,
|
||||
session_id: str,
|
||||
request: GenerationRequest,
|
||||
) -> Any:
|
||||
assignment = self._assignments.get(session_id)
|
||||
if assignment is None:
|
||||
raise RuntimeError(f"session {session_id!r} not acquired on this pool")
|
||||
handle = self._worker_by_id[assignment.worker_id]
|
||||
job_id = uuid.uuid4().hex
|
||||
future: Future = Future()
|
||||
self._pending[job_id] = _PendingJob(
|
||||
job_id=job_id,
|
||||
future=future,
|
||||
session_id=session_id,
|
||||
worker_id=handle.worker_id,
|
||||
)
|
||||
# mp.Queue.put can block if the underlying pipe buffer is full;
|
||||
# offload to a thread so the event loop keeps serving other
|
||||
# sessions. If the put itself fails, drop the pending entry so
|
||||
# _drain_results doesn't dangle a future forever.
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await loop.run_in_executor(
|
||||
None,
|
||||
handle.job_queue.put,
|
||||
{
|
||||
"job_id": job_id,
|
||||
"request": request
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
self._pending.pop(job_id, None)
|
||||
raise
|
||||
return await asyncio.wrap_future(future)
|
||||
|
||||
async def release(self, session_id: str) -> None:
|
||||
async with self._lock:
|
||||
assignment = self._assignments.pop(session_id, None)
|
||||
if assignment is None:
|
||||
return
|
||||
idx = next((i for i, h in enumerate(self._workers) if h.worker_id == assignment.worker_id), None)
|
||||
if idx is None:
|
||||
return
|
||||
# Don't return a dead worker to the pool; otherwise the next
|
||||
# acquire will hand a session to a process that can't run jobs.
|
||||
if not self._workers[idx].process.is_alive():
|
||||
logger.warning(
|
||||
"pool: worker %s died; not returning to available queue",
|
||||
self._workers[idx].worker_id,
|
||||
)
|
||||
return
|
||||
await self._available.put(idx)
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
# Signal all workers in parallel; .put may block on a full pipe,
|
||||
# so off-load it the same way run() does.
|
||||
async def _signal(handle: _WorkerHandle) -> None:
|
||||
try:
|
||||
handle.shutdown_event.set()
|
||||
await loop.run_in_executor(None, handle.job_queue.put, None)
|
||||
except Exception: # pragma: no cover - best-effort cleanup
|
||||
pass
|
||||
|
||||
await asyncio.gather(*(_signal(h) for h in self._workers))
|
||||
# Join in parallel so total shutdown is bounded by the slowest
|
||||
# worker, not the sum of all timeouts.
|
||||
await asyncio.gather(*(loop.run_in_executor(None, handle.process.join, 5.0) for handle in self._workers))
|
||||
for handle in self._workers:
|
||||
if handle.process.is_alive():
|
||||
handle.process.kill()
|
||||
for task in self._result_reader_tasks:
|
||||
task.cancel()
|
||||
self._result_reader_tasks.clear()
|
||||
self._workers.clear()
|
||||
self._worker_by_id.clear()
|
||||
|
||||
def health(self) -> PoolHealth:
|
||||
return PoolHealth(
|
||||
total_workers=len(self._workers),
|
||||
available_workers=self._available.qsize(),
|
||||
active_sessions=len(self._assignments),
|
||||
)
|
||||
|
||||
async def _drain_results(self, handle: _WorkerHandle) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
while not handle.shutdown_event.is_set():
|
||||
try:
|
||||
msg = await loop.run_in_executor(None, _safe_queue_get, handle.result_queue, 0.5)
|
||||
except Exception:
|
||||
logger.exception("pool: worker %s result reader failed", handle.worker_id)
|
||||
return
|
||||
if msg is None:
|
||||
continue
|
||||
job_id = msg.get("job_id")
|
||||
if job_id is None:
|
||||
continue
|
||||
pending = self._pending.pop(job_id, None)
|
||||
if pending is None:
|
||||
continue
|
||||
if msg.get("kind") == "error":
|
||||
pending.future.set_exception(RuntimeError(msg["error"]))
|
||||
else:
|
||||
pending.future.set_result(msg.get("result"))
|
||||
finally:
|
||||
# If we exit for any reason — shutdown, exception, cancel —
|
||||
# surface that to any in-flight jobs on this worker so their
|
||||
# await never hangs on a future no one will resolve.
|
||||
for jid in [jid for jid, job in self._pending.items() if job.worker_id == handle.worker_id]:
|
||||
pending = self._pending.pop(jid, None)
|
||||
if pending is not None and not pending.future.done():
|
||||
pending.future.set_exception(
|
||||
RuntimeError(f"worker {handle.worker_id} result reader exited "
|
||||
"with pending jobs"))
|
||||
|
||||
|
||||
def _safe_queue_get(q: mp.Queue, timeout: float) -> Any | None:
|
||||
try:
|
||||
return q.get(timeout=timeout)
|
||||
except queue.Empty:
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Worker process
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class WorkerFactory(Protocol):
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
gpu_id: int,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> _WorkerHandle:
|
||||
...
|
||||
|
||||
|
||||
def _default_worker_factory(
|
||||
*,
|
||||
gpu_id: int,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> _WorkerHandle:
|
||||
"""Spawn a real multiprocessing worker.
|
||||
|
||||
The child process calls :func:`worker_main` which constructs a
|
||||
:class:`VideoGenerator` from ``generator_config`` and runs a
|
||||
blocking job loop. The ``ready`` event flips after the warmup
|
||||
request completes.
|
||||
"""
|
||||
ctx = mp.get_context("spawn")
|
||||
job_queue: mp.Queue = ctx.Queue()
|
||||
result_queue: mp.Queue = ctx.Queue()
|
||||
ready = threading.Event()
|
||||
boot_ok = threading.Event()
|
||||
shutdown_event = ctx.Event()
|
||||
worker_id = f"gpu{gpu_id}-{uuid.uuid4().hex[:6]}"
|
||||
process = ctx.Process(
|
||||
target=worker_main,
|
||||
kwargs={
|
||||
"gpu_id": gpu_id,
|
||||
"worker_id": worker_id,
|
||||
"generator_config": generator_config,
|
||||
"warmup_config": warmup_config,
|
||||
"job_queue": job_queue,
|
||||
"result_queue": result_queue,
|
||||
"shutdown_event": shutdown_event,
|
||||
},
|
||||
daemon=False,
|
||||
)
|
||||
process.start()
|
||||
|
||||
# Block the parent-side ``ready`` flag until the worker posts a
|
||||
# ready acknowledgement on the result queue. We drain that single
|
||||
# sentinel here; subsequent results belong to jobs. ``boot_ok``
|
||||
# only flips on a real ready; on error we set ``ready`` to unblock
|
||||
# the parent's wait but leave ``boot_ok`` clear so the pool keeps
|
||||
# the worker out of the available queue.
|
||||
def _await_ready() -> None:
|
||||
while not shutdown_event.is_set():
|
||||
try:
|
||||
msg = result_queue.get(timeout=1.0)
|
||||
except queue.Empty:
|
||||
continue
|
||||
if isinstance(msg, dict) and msg.get("kind") == "ready":
|
||||
boot_ok.set()
|
||||
ready.set()
|
||||
return
|
||||
if isinstance(msg, dict) and msg.get("kind") == "error":
|
||||
logger.error("pool: worker %s failed to boot: %s", worker_id, msg.get("error"))
|
||||
ready.set()
|
||||
return
|
||||
|
||||
threading.Thread(target=_await_ready, daemon=True).start()
|
||||
|
||||
return _WorkerHandle(
|
||||
process=process,
|
||||
job_queue=job_queue,
|
||||
result_queue=result_queue,
|
||||
gpu_id=gpu_id,
|
||||
worker_id=worker_id,
|
||||
ready=ready,
|
||||
boot_ok=boot_ok,
|
||||
shutdown_event=shutdown_event,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"GpuPool",
|
||||
"InProcessGpuPool",
|
||||
"PoolAcquireTimeout",
|
||||
"PoolAssignment",
|
||||
"PoolHealth",
|
||||
"SubprocessGpuPool",
|
||||
"WorkerFactory",
|
||||
"worker_main",
|
||||
]
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Mock streaming server — a frontend dev aid.
|
||||
|
||||
Boots the same FastAPI app the real streaming server uses, but backs
|
||||
it with :class:`InProcessGpuPool` wrapping a synthetic generator that
|
||||
emits pre-baked RGB frames. No GPU or model weights required.
|
||||
|
||||
Use cases:
|
||||
|
||||
* Frontend development without a real model loaded.
|
||||
* Integration tests that exercise the WS protocol end-to-end.
|
||||
* Reproducing protocol bugs locally.
|
||||
|
||||
Launch: ``python -m fastvideo.entrypoints.streaming.mock_server``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
StreamingConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.server import build_app
|
||||
|
||||
|
||||
@dataclass
|
||||
class MockGenerator:
|
||||
"""Generator stand-in that returns synthetic gradient frames.
|
||||
|
||||
Each call produces one segment worth of frames whose pixels vary by
|
||||
a constant derived from the request seed and segment index. Latency
|
||||
is configurable via ``sleep_ms`` so the caller can exercise slow-
|
||||
generate scenarios without spinning a GPU.
|
||||
"""
|
||||
|
||||
sleep_ms: float = 0.0
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
if self.sleep_ms:
|
||||
time.sleep(self.sleep_ms / 1000.0)
|
||||
width = max(16, request.sampling.width)
|
||||
height = max(16, request.sampling.height)
|
||||
num_frames = max(1, request.sampling.num_frames)
|
||||
frames = [_gradient_frame(height, width, idx, seed=request.sampling.seed) for idx in range(num_frames)]
|
||||
state = ContinuationState(
|
||||
kind="ltx2.v1",
|
||||
payload={
|
||||
"schema_version": 1,
|
||||
"segment_index": 0,
|
||||
"source_prompt": request.prompt,
|
||||
},
|
||||
)
|
||||
return {
|
||||
"frames": frames,
|
||||
"audio_sample_rate": 24000,
|
||||
"state": state,
|
||||
}
|
||||
|
||||
|
||||
def _gradient_frame(height: int, width: int, idx: int, *, seed: int) -> np.ndarray:
|
||||
base = (idx * 17 + seed * 3) % 256
|
||||
row = np.linspace(base, (base + 64) % 256, width, dtype=np.uint8)
|
||||
frame = np.tile(row, (height, 1))
|
||||
stacked = np.stack([frame, np.roll(frame, 8, axis=1), np.roll(frame, 16, axis=1)], axis=-1)
|
||||
return stacked.astype(np.uint8)
|
||||
|
||||
|
||||
def build_mock_app(*, sleep_ms: float = 0.0):
|
||||
"""Build a FastAPI app backed by :class:`MockGenerator`."""
|
||||
serve_config = ServeConfig(
|
||||
generator=GeneratorConfig(model_path="/models/mock"),
|
||||
streaming=StreamingConfig(
|
||||
session_timeout_seconds=120,
|
||||
generation_segment_cap=6,
|
||||
),
|
||||
)
|
||||
serve_config.default_request.sampling = SamplingConfig(
|
||||
num_frames=24,
|
||||
height=256,
|
||||
width=256,
|
||||
fps=24,
|
||||
num_inference_steps=1,
|
||||
)
|
||||
return build_app(serve_config, MockGenerator(sleep_ms=sleep_ms))
|
||||
|
||||
|
||||
def main() -> None: # pragma: no cover - CLI entry
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--host", default="127.0.0.1")
|
||||
parser.add_argument("--port", type=int, default=8000)
|
||||
parser.add_argument(
|
||||
"--sleep-ms",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Per-segment artificial latency for testing slow paths",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
import uvicorn
|
||||
|
||||
app = build_mock_app(sleep_ms=args.sleep_ms)
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MockGenerator",
|
||||
"build_mock_app",
|
||||
"main",
|
||||
]
|
||||
|
||||
if __name__ == "__main__": # pragma: no cover - CLI entry
|
||||
main()
|
||||
@@ -0,0 +1,36 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Prompt pipeline for the streaming server.
|
||||
|
||||
* :mod:`providers` — LLM backend abstraction + built-in adapters
|
||||
* :mod:`enhancer` — provider-agnostic enhance / auto-extend / rewrite
|
||||
operations on top of the provider layer
|
||||
|
||||
All of this is optional; the streaming server runs fine without it
|
||||
(PR 7.5's skeleton never invokes the enhancer). When the operator
|
||||
enables ``ServeConfig.streaming.prompt.enabled``, the server routes
|
||||
each ``session_init_v2`` curated prompt through ``enhance`` before the
|
||||
first segment.
|
||||
"""
|
||||
from fastvideo.entrypoints.streaming.prompt.enhancer import (
|
||||
PromptEnhancer,
|
||||
PromptOperation,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LLMMessage",
|
||||
"LLMProvider",
|
||||
"LLMProviderError",
|
||||
"LLMRequest",
|
||||
"LLMResponse",
|
||||
"LLMTimeoutError",
|
||||
"PromptEnhancer",
|
||||
"PromptOperation",
|
||||
]
|
||||
@@ -0,0 +1,197 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Provider-agnostic prompt orchestration for the streaming server.
|
||||
|
||||
Three operations the streaming server needs:
|
||||
|
||||
* ``enhance`` — polish a user prompt (add cinematic detail, fix syntax)
|
||||
* ``auto_extend`` — generate a follow-on prompt for loop generation
|
||||
* ``rewrite`` — rewrite a seed prompt for a user-directed rewrite flow
|
||||
|
||||
All three share the same orchestration: pick a provider in priority
|
||||
order, submit an ``LLMRequest``, fall back to the next provider on
|
||||
retryable errors, and surface a structured :class:`LLMResponse` back
|
||||
to the caller.
|
||||
|
||||
System prompts are loaded from ``system_prompt_dir`` on construction
|
||||
and can be hot-reloaded via :meth:`PromptEnhancer.reload_system_prompts`.
|
||||
The streaming server's management endpoint calls that method in
|
||||
response to a ``rewrite_seed_prompts_started`` frame.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import os
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PromptOperation(enum.Enum):
|
||||
ENHANCE = "enhance"
|
||||
AUTO_EXTEND = "auto_extend"
|
||||
REWRITE = "rewrite"
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SystemPrompts:
|
||||
enhance: str
|
||||
auto_extend: str
|
||||
rewrite: str
|
||||
|
||||
|
||||
_DEFAULT_SYSTEM_PROMPTS = _SystemPrompts(
|
||||
enhance=("You are a prompt enhancer for cinematic video generation. Given "
|
||||
"a user prompt, produce an enhanced prompt that is more vivid, "
|
||||
"specific, and concrete. Keep the subject intact; add lighting, "
|
||||
"camera, and motion detail. Reply with just the enhanced prompt."),
|
||||
auto_extend=("You are a video continuation assistant. Given the current "
|
||||
"sequence of prompts, produce one new prompt that naturally "
|
||||
"continues the sequence. Reply with just the next prompt."),
|
||||
rewrite=("You are a creative prompt rewriter. Given a seed prompt, produce "
|
||||
"a set of alternative prompts that explore different angles, "
|
||||
"styles, and moods. Reply with one prompt per line."),
|
||||
)
|
||||
|
||||
|
||||
class PromptEnhancer:
|
||||
"""Orchestrates prompt operations across a priority-ordered provider
|
||||
list with structured fallback + hot-reloadable system prompts.
|
||||
|
||||
Usage::
|
||||
|
||||
enhancer = PromptEnhancer(
|
||||
providers=[CerebrasProvider(), GroqProvider()],
|
||||
model="gpt-oss-120b",
|
||||
system_prompt_dir="/etc/fastvideo/prompts",
|
||||
)
|
||||
response = await enhancer.enhance("a fox running through snow")
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
providers: Sequence[LLMProvider],
|
||||
model: str,
|
||||
timeout_ms: int = 20000,
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int | None = 256,
|
||||
system_prompt_dir: str | None = None,
|
||||
) -> None:
|
||||
if not providers:
|
||||
raise ValueError("PromptEnhancer requires at least one LLMProvider")
|
||||
self._providers = list(providers)
|
||||
self._model = model
|
||||
self._timeout_ms = timeout_ms
|
||||
self._temperature = temperature
|
||||
self._max_tokens = max_tokens
|
||||
self._system_prompt_dir = system_prompt_dir
|
||||
self._system_prompts = self._load_system_prompts()
|
||||
|
||||
@property
|
||||
def providers(self) -> list[LLMProvider]:
|
||||
return list(self._providers)
|
||||
|
||||
def register_provider(self, provider: LLMProvider, *, priority: int = -1) -> None:
|
||||
"""Insert an additional provider. ``priority=0`` makes it primary;
|
||||
``priority=-1`` (default) appends as a fallback."""
|
||||
if priority < 0:
|
||||
self._providers.append(provider)
|
||||
else:
|
||||
self._providers.insert(priority, provider)
|
||||
|
||||
def reload_system_prompts(self) -> None:
|
||||
"""Re-read the system prompt files from ``system_prompt_dir``.
|
||||
|
||||
The streaming server exposes this via a management endpoint so
|
||||
operators can iterate on prompt templates without restarting
|
||||
workers.
|
||||
"""
|
||||
self._system_prompts = self._load_system_prompts()
|
||||
logger.info("prompt enhancer: reloaded system prompts from %s", self._system_prompt_dir or "defaults")
|
||||
|
||||
async def enhance(self, prompt: str) -> LLMResponse:
|
||||
return await self._run(
|
||||
PromptOperation.ENHANCE,
|
||||
system=self._system_prompts.enhance,
|
||||
user=prompt,
|
||||
)
|
||||
|
||||
async def auto_extend(self, prior_prompts: Sequence[str]) -> LLMResponse:
|
||||
user = "\n".join(prior_prompts)
|
||||
return await self._run(
|
||||
PromptOperation.AUTO_EXTEND,
|
||||
system=self._system_prompts.auto_extend,
|
||||
user=user,
|
||||
)
|
||||
|
||||
async def rewrite(self, seed_prompt: str) -> LLMResponse:
|
||||
return await self._run(
|
||||
PromptOperation.REWRITE,
|
||||
system=self._system_prompts.rewrite,
|
||||
user=seed_prompt,
|
||||
)
|
||||
|
||||
async def _run(
|
||||
self,
|
||||
operation: PromptOperation,
|
||||
*,
|
||||
system: str,
|
||||
user: str,
|
||||
) -> LLMResponse:
|
||||
request = LLMRequest(
|
||||
messages=[
|
||||
LLMMessage(role="system", content=system),
|
||||
LLMMessage(role="user", content=user),
|
||||
],
|
||||
model=self._model,
|
||||
max_tokens=self._max_tokens,
|
||||
temperature=self._temperature,
|
||||
timeout_ms=self._timeout_ms,
|
||||
)
|
||||
last_error: LLMProviderError | None = None
|
||||
for idx, provider in enumerate(self._providers):
|
||||
try:
|
||||
response = await provider.complete(request)
|
||||
if idx > 0:
|
||||
# Mark the fallback flag without losing any other
|
||||
# response fields the provider populated.
|
||||
response = replace(response, fallback_used=True)
|
||||
return response
|
||||
except LLMProviderError as exc:
|
||||
logger.warning("prompt %s: provider %s failed: %s; trying next", operation.value, provider.name, exc)
|
||||
last_error = exc
|
||||
if not exc.retryable:
|
||||
break
|
||||
assert last_error is not None
|
||||
raise last_error
|
||||
|
||||
def _load_system_prompts(self) -> _SystemPrompts:
|
||||
if not self._system_prompt_dir:
|
||||
return _DEFAULT_SYSTEM_PROMPTS
|
||||
return _SystemPrompts(
|
||||
enhance=_read_prompt(self._system_prompt_dir, "enhance.txt", _DEFAULT_SYSTEM_PROMPTS.enhance),
|
||||
auto_extend=_read_prompt(self._system_prompt_dir, "auto_extend.txt", _DEFAULT_SYSTEM_PROMPTS.auto_extend),
|
||||
rewrite=_read_prompt(self._system_prompt_dir, "rewrite.txt", _DEFAULT_SYSTEM_PROMPTS.rewrite),
|
||||
)
|
||||
|
||||
|
||||
def _read_prompt(dirname: str, filename: str, default: str) -> str:
|
||||
path = os.path.join(dirname, filename)
|
||||
if not os.path.exists(path):
|
||||
return default
|
||||
with open(path, encoding="utf-8") as f:
|
||||
content = f.read().strip()
|
||||
return content or default
|
||||
|
||||
|
||||
__all__ = ["PromptEnhancer", "PromptOperation"]
|
||||
@@ -0,0 +1,24 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LLM provider implementations used by the prompt enhancer."""
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.cerebras import (
|
||||
CerebrasProvider, )
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.groq import GroqProvider
|
||||
|
||||
__all__ = [
|
||||
"CerebrasProvider",
|
||||
"GroqProvider",
|
||||
"LLMMessage",
|
||||
"LLMProvider",
|
||||
"LLMProviderError",
|
||||
"LLMRequest",
|
||||
"LLMResponse",
|
||||
"LLMTimeoutError",
|
||||
]
|
||||
@@ -0,0 +1,101 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared HTTP path for OpenAI-compatible ``/chat/completions`` providers.
|
||||
|
||||
Cerebras and Groq both expose the OpenAI chat-completions schema, so
|
||||
the request shape, error mapping, and response decoding are identical
|
||||
between them. This module centralizes that logic; the per-provider
|
||||
modules stay thin (just defaults + env var wiring).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
|
||||
|
||||
async def complete_openai_compatible(
|
||||
*,
|
||||
api_key: str | None,
|
||||
api_key_hint: str,
|
||||
base_url: str,
|
||||
provider_name: str,
|
||||
request: LLMRequest,
|
||||
) -> LLMResponse:
|
||||
"""Issue a chat-completions call and decode the OpenAI response."""
|
||||
if not api_key:
|
||||
raise LLMProviderError(
|
||||
f"{provider_name} provider requires {api_key_hint} "
|
||||
"(or explicit api_key=...)",
|
||||
retryable=False,
|
||||
)
|
||||
try:
|
||||
import httpx
|
||||
except ImportError as exc: # pragma: no cover - optional dep
|
||||
raise LLMProviderError(
|
||||
f"{provider_name} provider requires httpx; install httpx",
|
||||
retryable=False,
|
||||
) from exc
|
||||
|
||||
timeout_s = (request.timeout_ms or 20000) / 1000.0
|
||||
t0 = time.perf_counter()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=timeout_s) as client:
|
||||
response = await client.post(
|
||||
f"{base_url}/chat/completions",
|
||||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
json={
|
||||
"model": request.model,
|
||||
"messages": [{
|
||||
"role": m.role,
|
||||
"content": m.content
|
||||
} for m in request.messages],
|
||||
"max_tokens": request.max_tokens,
|
||||
"temperature": request.temperature,
|
||||
},
|
||||
)
|
||||
except httpx.TimeoutException as exc:
|
||||
raise LLMTimeoutError(f"{provider_name} timed out after {timeout_s}s") from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise LLMProviderError(f"{provider_name} HTTP error: {exc}") from exc
|
||||
|
||||
if response.status_code >= 400:
|
||||
# 5xx and 429 (rate-limit) are retryable: another provider may
|
||||
# succeed. 4xx (auth, bad-request, etc.) are client errors —
|
||||
# the enhancer should stop fallback traversal.
|
||||
retryable = (response.status_code >= 500 or response.status_code == 429)
|
||||
raise LLMProviderError(
|
||||
f"{provider_name} returned {response.status_code}: "
|
||||
f"{response.text[:200]}",
|
||||
retryable=retryable,
|
||||
)
|
||||
|
||||
try:
|
||||
data = response.json()
|
||||
except Exception as exc:
|
||||
# Non-JSON body usually means a proxy / load-balancer error
|
||||
# page; leave it retryable so a fallback provider can try.
|
||||
raise LLMProviderError(f"{provider_name} returned non-JSON body: {exc}") from exc
|
||||
|
||||
choices = data.get("choices") or []
|
||||
if not choices:
|
||||
raise LLMProviderError(f"{provider_name} returned no choices")
|
||||
content = choices[0].get("message", {}).get("content") or ""
|
||||
|
||||
latency_ms = (time.perf_counter() - t0) * 1000.0
|
||||
return LLMResponse(
|
||||
content=content.strip(),
|
||||
provider=provider_name,
|
||||
model=request.model,
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["complete_openai_compatible"]
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LLM provider protocol + DTOs used by the prompt enhancer.
|
||||
|
||||
Third-party users add a new provider by implementing
|
||||
:class:`LLMProvider` and registering it with a prompt enhancer
|
||||
instance. The shipped providers live in sibling modules
|
||||
(``cerebras.py``, ``groq.py``) and each is ~100-200 LOC — the
|
||||
provider layer is intentionally thin so the enhancer stays
|
||||
provider-agnostic.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMMessage:
|
||||
role: Literal["system", "user", "assistant"]
|
||||
content: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMRequest:
|
||||
messages: list[LLMMessage]
|
||||
model: str
|
||||
max_tokens: int | None = None
|
||||
temperature: float | None = None
|
||||
timeout_ms: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMResponse:
|
||||
content: str
|
||||
provider: str
|
||||
model: str
|
||||
latency_ms: float
|
||||
fallback_used: bool = False
|
||||
|
||||
|
||||
class LLMProviderError(RuntimeError):
|
||||
"""Raised when an LLM provider fails a request.
|
||||
|
||||
``retryable`` controls whether the enhancer falls back to the next
|
||||
provider. It is settable per-instance so the same exception type
|
||||
can describe retryable transport errors (5xx, 429) and
|
||||
non-retryable client errors (4xx auth/bad-request) without forcing
|
||||
a separate subclass for every status family.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str, *, retryable: bool = True) -> None:
|
||||
super().__init__(message)
|
||||
self.retryable = retryable
|
||||
|
||||
|
||||
class LLMTimeoutError(LLMProviderError):
|
||||
"""Raised when an LLM provider times out — always retryable."""
|
||||
|
||||
def __init__(self, message: str) -> None:
|
||||
super().__init__(message, retryable=True)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class LLMProvider(Protocol):
|
||||
"""Provider interface every LLM adapter implements.
|
||||
|
||||
Providers are async-first because every built-in implementation
|
||||
talks to an HTTP API. Synchronous providers can wrap their call in
|
||||
``asyncio.to_thread`` internally.
|
||||
"""
|
||||
|
||||
name: str
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
...
|
||||
|
||||
|
||||
__all__ = [
|
||||
"LLMMessage",
|
||||
"LLMProvider",
|
||||
"LLMProviderError",
|
||||
"LLMRequest",
|
||||
"LLMResponse",
|
||||
"LLMTimeoutError",
|
||||
]
|
||||
@@ -0,0 +1,44 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cerebras LLM provider (OpenAI-compatible chat endpoint)."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers._openai_compat import (
|
||||
complete_openai_compatible, )
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
)
|
||||
|
||||
_DEFAULT_BASE_URL = "https://api.cerebras.ai/v1"
|
||||
_API_KEY_ENV = "CEREBRAS_API_KEY"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CerebrasProvider:
|
||||
"""Cerebras inference adapter.
|
||||
|
||||
``api_key`` falls back to ``CEREBRAS_API_KEY`` when unset.
|
||||
"""
|
||||
|
||||
api_key: str | None = None
|
||||
base_url: str = _DEFAULT_BASE_URL
|
||||
name: str = "cerebras"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.api_key is None:
|
||||
self.api_key = os.environ.get(_API_KEY_ENV)
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
return await complete_openai_compatible(
|
||||
api_key=self.api_key,
|
||||
api_key_hint=_API_KEY_ENV,
|
||||
base_url=self.base_url,
|
||||
provider_name=self.name,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["CerebrasProvider"]
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Groq LLM provider (OpenAI-compatible chat endpoint)."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers._openai_compat import (
|
||||
complete_openai_compatible, )
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
)
|
||||
|
||||
_DEFAULT_BASE_URL = "https://api.groq.com/openai/v1"
|
||||
_API_KEY_ENV = "GROQ_API_KEY"
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroqProvider:
|
||||
"""Groq inference adapter.
|
||||
|
||||
Identical wire format to :class:`CerebrasProvider`; both go through
|
||||
:func:`complete_openai_compatible`. The two providers differ only
|
||||
in base URL, env var, and model id conventions.
|
||||
"""
|
||||
|
||||
api_key: str | None = None
|
||||
base_url: str = _DEFAULT_BASE_URL
|
||||
name: str = "groq"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.api_key is None:
|
||||
self.api_key = os.environ.get(_API_KEY_ENV)
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
return await complete_openai_compatible(
|
||||
api_key=self.api_key,
|
||||
api_key_hint=_API_KEY_ENV,
|
||||
base_url=self.base_url,
|
||||
provider_name=self.name,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["GroqProvider"]
|
||||
@@ -0,0 +1,82 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Rewrite payload builder.
|
||||
|
||||
The UI's "rewrite seed prompts" flow asks the enhancer to produce a
|
||||
batch of alternative prompts given one seed. This module packages the
|
||||
seed + options into the payload the enhancer expects and unpacks the
|
||||
response back into a typed :class:`RewriteResult`.
|
||||
|
||||
Separating this from :mod:`enhancer` keeps the enhancer provider-
|
||||
agnostic; anything UI-specific (how many alternatives to request, how
|
||||
to split the response, temperature) lives here.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.enhancer import PromptEnhancer
|
||||
|
||||
_LEADING_MARKER_RE = re.compile(r"^(?:[-*•]\s*|\d+\s*[.)]\s*)+")
|
||||
|
||||
|
||||
@dataclass
|
||||
class RewriteOptions:
|
||||
count: int = 3
|
||||
"""Number of alternative prompts to request."""
|
||||
temperature: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RewriteResult:
|
||||
seed_prompt: str
|
||||
alternatives: list[str]
|
||||
provider: str
|
||||
model: str
|
||||
latency_ms: float
|
||||
fallback_used: bool = False
|
||||
|
||||
|
||||
async def build_rewrite(
|
||||
enhancer: PromptEnhancer,
|
||||
seed_prompt: str,
|
||||
*,
|
||||
options: RewriteOptions | None = None,
|
||||
) -> RewriteResult:
|
||||
"""Run a rewrite op through the enhancer and return a typed result."""
|
||||
if not seed_prompt.strip():
|
||||
raise ValueError("rewrite seed prompt must be non-empty")
|
||||
options = options or RewriteOptions()
|
||||
response = await enhancer.rewrite(seed_prompt)
|
||||
alternatives = _split_response(response.content, limit=options.count)
|
||||
return RewriteResult(
|
||||
seed_prompt=seed_prompt,
|
||||
alternatives=alternatives,
|
||||
provider=response.provider,
|
||||
model=response.model,
|
||||
latency_ms=response.latency_ms,
|
||||
fallback_used=response.fallback_used,
|
||||
)
|
||||
|
||||
|
||||
def _split_response(content: str, *, limit: int) -> list[str]:
|
||||
"""Split the LLM response into discrete prompt candidates.
|
||||
|
||||
The shipped system prompt instructs the model to emit one prompt
|
||||
per line; this function is forgiving about numbered lists or
|
||||
leading bullets so user-supplied system prompts don't break it.
|
||||
"""
|
||||
lines = [line.strip() for line in content.splitlines() if line.strip()]
|
||||
cleaned: list[str] = []
|
||||
for line in lines:
|
||||
stripped = _LEADING_MARKER_RE.sub("", line).strip()
|
||||
if stripped:
|
||||
cleaned.append(stripped)
|
||||
return cleaned[:max(1, limit)]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RewriteOptions",
|
||||
"RewriteResult",
|
||||
"build_rewrite",
|
||||
]
|
||||
@@ -0,0 +1,146 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Optional prompt safety filter.
|
||||
|
||||
Uses a fastText classifier to score prompts against a banned-content
|
||||
rubric. Only loaded when ``ServeConfig.streaming.safety.enabled`` is
|
||||
True and fastText is installed — users who don't need it see no
|
||||
runtime cost.
|
||||
|
||||
Install: ``pip install fastvideo[prompt-safety]`` (ships fasttext as an
|
||||
optional extra) or install fasttext directly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SafetyDecision(enum.Enum):
|
||||
ALLOW = "allow"
|
||||
BLOCK = "block"
|
||||
UNAVAILABLE = "unavailable"
|
||||
"""Returned when the classifier can't run (not configured, fastText
|
||||
missing). Safety is opt-in; the server treats ``UNAVAILABLE`` as
|
||||
``ALLOW`` but logs it so operators know the filter is off."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SafetyResult:
|
||||
prompt: str
|
||||
decision: SafetyDecision
|
||||
score: float = 0.0
|
||||
label: str | None = None
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
class PromptSafetyFilter:
|
||||
"""Minimal fastText-backed prompt safety filter.
|
||||
|
||||
Loads the classifier lazily on first use so the streaming server
|
||||
can construct the filter eagerly at startup without paying the
|
||||
model-load cost when safety is disabled.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
classifier_path: str | None,
|
||||
enabled: bool = True,
|
||||
block_threshold: float = 0.5,
|
||||
) -> None:
|
||||
self._classifier_path = classifier_path
|
||||
self._enabled = enabled
|
||||
self._block_threshold = block_threshold
|
||||
self._model: Any | None = None
|
||||
self._load_attempted = False
|
||||
self._load_lock = threading.Lock()
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._enabled and self._classifier_path is not None
|
||||
|
||||
def classify(self, prompt: str) -> SafetyResult:
|
||||
if not self.enabled:
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=SafetyDecision.UNAVAILABLE,
|
||||
reason="safety filter not enabled",
|
||||
)
|
||||
model = self._ensure_loaded()
|
||||
if model is None:
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=SafetyDecision.UNAVAILABLE,
|
||||
reason="fastText model unavailable",
|
||||
)
|
||||
try:
|
||||
labels, probs = model.predict(prompt.replace("\n", " "), k=1)
|
||||
except Exception as exc: # pragma: no cover - defensive
|
||||
logger.warning("safety: classifier failed: %s", exc)
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=SafetyDecision.UNAVAILABLE,
|
||||
reason=f"classifier error: {exc}",
|
||||
)
|
||||
label = labels[0].removeprefix("__label__") if labels else None
|
||||
score = float(probs[0]) if len(probs) else 0.0
|
||||
decision = (SafetyDecision.BLOCK if
|
||||
(label == "unsafe" and score >= self._block_threshold) else SafetyDecision.ALLOW)
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=decision,
|
||||
score=score,
|
||||
label=label,
|
||||
)
|
||||
|
||||
def _ensure_loaded(self) -> Any | None:
|
||||
if self._model is not None:
|
||||
return self._model
|
||||
if self._load_attempted:
|
||||
return None
|
||||
with self._load_lock:
|
||||
if self._model is not None:
|
||||
return self._model
|
||||
if self._load_attempted:
|
||||
return None
|
||||
self._load_attempted = True
|
||||
if self._classifier_path is None:
|
||||
return None
|
||||
try:
|
||||
import fasttext # type: ignore[import-not-found]
|
||||
except ImportError:
|
||||
logger.warning("safety: fasttext not installed; safety filter disabled. "
|
||||
"Install fastvideo[prompt-safety] to enable.")
|
||||
return None
|
||||
try:
|
||||
self._model = fasttext.load_model(self._classifier_path)
|
||||
except Exception as exc: # pragma: no cover - requires real model
|
||||
logger.warning("safety: failed to load %s: %s", self._classifier_path, exc)
|
||||
return None
|
||||
return self._model
|
||||
|
||||
|
||||
def first_blocked(
|
||||
filter_: PromptSafetyFilter,
|
||||
prompts: list[str],
|
||||
) -> SafetyResult | None:
|
||||
"""Return the first prompt the filter blocks, or ``None``."""
|
||||
for prompt in prompts:
|
||||
result = filter_.classify(prompt)
|
||||
if result.decision is SafetyDecision.BLOCK:
|
||||
return result
|
||||
return None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PromptSafetyFilter",
|
||||
"SafetyDecision",
|
||||
"SafetyResult",
|
||||
"first_blocked",
|
||||
]
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Multi-replica load balancer + WebSocket proxy for the streaming server.
|
||||
|
||||
Sits in front of one-or-more streaming-server replicas and forwards
|
||||
WebSocket sessions to a healthy primary, with failover to secondaries.
|
||||
Kept in-repo under ``fastvideo/entrypoints/streaming/router/`` per the
|
||||
PR plan's default; the alternative (separate package) is an open
|
||||
question deferred to review.
|
||||
"""
|
||||
from fastvideo.entrypoints.streaming.router.registry import (
|
||||
Replica,
|
||||
ReplicaHealth,
|
||||
ReplicaRegistry,
|
||||
ReplicaStatus,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.router.config import RouterConfig
|
||||
from fastvideo.entrypoints.streaming.router.main import build_router_app, run_router
|
||||
|
||||
__all__ = [
|
||||
"Replica",
|
||||
"ReplicaHealth",
|
||||
"ReplicaRegistry",
|
||||
"ReplicaStatus",
|
||||
"RouterConfig",
|
||||
"build_router_app",
|
||||
"run_router",
|
||||
]
|
||||
@@ -0,0 +1,88 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Typed router configuration."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReplicaEndpoint:
|
||||
"""One backend replica the router can route to."""
|
||||
|
||||
url: str
|
||||
"""HTTP base URL, e.g. ``http://host:8000``. WebSocket URL is
|
||||
derived automatically by replacing the scheme."""
|
||||
name: str | None = None
|
||||
primary: bool = False
|
||||
"""``True`` = prefer this replica over others in steady state."""
|
||||
weight: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class RouterConfig:
|
||||
"""Typed router config loaded from a YAML file.
|
||||
|
||||
Example::
|
||||
|
||||
router:
|
||||
host: 0.0.0.0
|
||||
port: 9000
|
||||
replicas:
|
||||
- url: http://streamer-a:8000
|
||||
primary: true
|
||||
- url: http://streamer-b:8000
|
||||
health_check:
|
||||
path: /health
|
||||
interval_seconds: 5
|
||||
failure_threshold: 3
|
||||
|
||||
Validation runs in ``__post_init__``: empty replicas, non-positive
|
||||
intervals/timeouts, thresholds < 1, non-http(s) URLs, and more than
|
||||
one primary all raise ``ValueError`` so misconfigurations surface at
|
||||
load time rather than as confusing runtime failures.
|
||||
"""
|
||||
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 9000
|
||||
replicas: list[ReplicaEndpoint] = field(default_factory=list)
|
||||
health_check_path: str = "/health"
|
||||
health_check_interval_seconds: float = 5.0
|
||||
health_check_timeout_seconds: float = 2.0
|
||||
failure_threshold: int = 3
|
||||
recovery_threshold: int = 2
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.replicas:
|
||||
raise ValueError("RouterConfig.replicas must list at least one replica")
|
||||
if self.health_check_interval_seconds <= 0:
|
||||
raise ValueError(f"health_check_interval_seconds must be > 0, got {self.health_check_interval_seconds}")
|
||||
if self.health_check_timeout_seconds <= 0:
|
||||
raise ValueError(f"health_check_timeout_seconds must be > 0, got {self.health_check_timeout_seconds}")
|
||||
if self.failure_threshold < 1:
|
||||
raise ValueError(f"failure_threshold must be >= 1, got {self.failure_threshold}")
|
||||
if self.recovery_threshold < 1:
|
||||
raise ValueError(f"recovery_threshold must be >= 1, got {self.recovery_threshold}")
|
||||
seen_urls: set[str] = set()
|
||||
for replica in self.replicas:
|
||||
if not replica.url.startswith(("http://", "https://")):
|
||||
raise ValueError(f"ReplicaEndpoint.url must start with http:// or https://, got {replica.url!r}")
|
||||
parsed = urlparse(replica.url)
|
||||
if parsed.path not in ("", "/"):
|
||||
raise ValueError(f"ReplicaEndpoint.url must be a base host[:port] URL without a path; "
|
||||
f"got {replica.url!r} with path {parsed.path!r}. The router appends "
|
||||
"`/health` and `/v1/stream` itself.")
|
||||
if parsed.query or parsed.fragment:
|
||||
raise ValueError(f"ReplicaEndpoint.url must not include query/fragment; got {replica.url!r}")
|
||||
if replica.url in seen_urls:
|
||||
raise ValueError(f"Duplicate ReplicaEndpoint.url {replica.url!r}; "
|
||||
"router selection keys by URL so duplicates would silently collapse")
|
||||
seen_urls.add(replica.url)
|
||||
primaries = sum(1 for r in self.replicas if r.primary)
|
||||
if primaries > 1:
|
||||
raise ValueError(f"RouterConfig allows at most one primary replica; got {primaries}. "
|
||||
"Multi-primary load distribution is deferred — promote one replica to "
|
||||
"primary and treat the rest as secondaries.")
|
||||
|
||||
|
||||
__all__ = ["ReplicaEndpoint", "RouterConfig"]
|
||||
@@ -0,0 +1,218 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Router FastAPI entry point.
|
||||
|
||||
Exposes the same ``/v1/stream`` WebSocket path the backend servers do,
|
||||
accepts a client, picks a healthy replica from the registry, and
|
||||
proxies frames bidirectionally.
|
||||
|
||||
PR 7.9 ships the minimum-viable shape: explicit replica list, single
|
||||
primary, JSON + binary passthrough in both directions, and a
|
||||
``/status`` endpoint for operators. Sticky-session routing (so a
|
||||
reconnect lands on the same backend) is left for a follow-up.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.config import RouterConfig
|
||||
from fastvideo.entrypoints.streaming.router.registry import (
|
||||
ReplicaRegistry,
|
||||
run_health_check_loop,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RouterState:
|
||||
config: RouterConfig
|
||||
registry: ReplicaRegistry
|
||||
stop_event: asyncio.Event
|
||||
health_task: asyncio.Task | None = None
|
||||
|
||||
|
||||
def build_router_app(
|
||||
config: RouterConfig,
|
||||
*,
|
||||
registry: ReplicaRegistry | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the router FastAPI app.
|
||||
|
||||
``registry`` can be injected for tests; defaults to one built from
|
||||
``config.replicas``.
|
||||
"""
|
||||
registry = registry or ReplicaRegistry(config.replicas)
|
||||
state = _RouterState(
|
||||
config=config,
|
||||
registry=registry,
|
||||
stop_event=asyncio.Event(),
|
||||
)
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _lifespan(_app: FastAPI):
|
||||
state.health_task = asyncio.create_task(
|
||||
run_health_check_loop(
|
||||
registry=state.registry,
|
||||
config=state.config,
|
||||
stop_event=state.stop_event,
|
||||
))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
state.stop_event.set()
|
||||
if state.health_task is not None:
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await state.health_task
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming Router", lifespan=_lifespan)
|
||||
|
||||
@app.get("/status")
|
||||
async def _status() -> JSONResponse:
|
||||
return JSONResponse({
|
||||
"replicas": [{
|
||||
"url": r.url,
|
||||
"primary": r.primary,
|
||||
"status": r.health.status.value,
|
||||
"last_ok_at": r.health.last_ok_at,
|
||||
"last_latency_ms": r.health.last_latency_ms,
|
||||
"consecutive_failures": r.health.consecutive_failures,
|
||||
} for r in state.registry.all()],
|
||||
})
|
||||
|
||||
@app.websocket("/v1/stream")
|
||||
async def _proxy(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
replica = state.registry.select()
|
||||
if replica is None:
|
||||
await websocket.send_json({
|
||||
"type": "error",
|
||||
"code": "gpu_unavailable",
|
||||
"message": "router: no healthy replica available",
|
||||
"retryable": True,
|
||||
})
|
||||
await websocket.close(code=1013, reason="no_healthy_replica")
|
||||
return
|
||||
|
||||
ws_url = _websocket_url_for(replica.url)
|
||||
try:
|
||||
await _bridge_session(websocket, ws_url)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("router: client disconnected")
|
||||
except Exception as exc:
|
||||
logger.exception("router: bridge failed: %s", exc)
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await websocket.send_json({
|
||||
"type": "error",
|
||||
"code": "worker_failed",
|
||||
"message": f"router bridge failed: {exc}",
|
||||
"retryable": True,
|
||||
})
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await websocket.close(code=1011)
|
||||
|
||||
app.state.router_state = state
|
||||
return app
|
||||
|
||||
|
||||
def run_router(config: RouterConfig) -> None: # pragma: no cover - CLI
|
||||
import uvicorn
|
||||
|
||||
app = build_router_app(config)
|
||||
uvicorn.run(app, host=config.host, port=config.port)
|
||||
|
||||
|
||||
async def _bridge_session(
|
||||
client_ws: WebSocket,
|
||||
backend_ws_url: str,
|
||||
) -> None:
|
||||
"""Connect to backend and shuttle messages in both directions.
|
||||
|
||||
Uses ``websockets`` for the backend side; imported lazily to keep
|
||||
the router's import graph small for users who only want the server.
|
||||
|
||||
Cancellation: when either direction completes (client disconnect,
|
||||
backend close, exception), the other is cancelled explicitly and
|
||||
both are drained before returning. Unexpected exceptions from the
|
||||
direction that completed first are re-raised; normal disconnect
|
||||
paths (``WebSocketDisconnect``, ``ConnectionClosed``,
|
||||
``CancelledError``) are swallowed.
|
||||
"""
|
||||
try:
|
||||
import websockets
|
||||
except ImportError as exc: # pragma: no cover - optional extra
|
||||
raise RuntimeError("router requires the `websockets` package for backend proxying") from exc
|
||||
|
||||
async with websockets.connect(backend_ws_url + "/v1/stream") as backend_ws:
|
||||
c2b = asyncio.create_task(_forward_client_to_backend(client_ws, backend_ws))
|
||||
b2c = asyncio.create_task(_forward_backend_to_client(backend_ws, client_ws))
|
||||
try:
|
||||
done, _pending = await asyncio.wait(
|
||||
{c2b, b2c},
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
finally:
|
||||
for task in (c2b, b2c):
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(c2b, b2c, return_exceptions=True)
|
||||
for task in done:
|
||||
task_exc = task.exception()
|
||||
if task_exc is not None and not _is_normal_disconnect(task_exc):
|
||||
raise task_exc
|
||||
|
||||
|
||||
def _is_normal_disconnect(exc: BaseException) -> bool:
|
||||
"""Whether ``exc`` is a routine WebSocket teardown vs a real bridge fault."""
|
||||
if isinstance(exc, asyncio.CancelledError | WebSocketDisconnect):
|
||||
return True
|
||||
name = type(exc).__name__
|
||||
# websockets.exceptions.ConnectionClosed{,OK,Error} all subclass
|
||||
# WebSocketException; check by name to avoid the lazy-import dance.
|
||||
return name.startswith("ConnectionClosed")
|
||||
|
||||
|
||||
async def _forward_client_to_backend(client_ws: WebSocket, backend_ws) -> None:
|
||||
try:
|
||||
while True:
|
||||
msg = await client_ws.receive()
|
||||
if msg.get("type") == "websocket.disconnect":
|
||||
break
|
||||
if "text" in msg and msg["text"] is not None:
|
||||
await backend_ws.send(msg["text"])
|
||||
elif "bytes" in msg and msg["bytes"] is not None:
|
||||
await backend_ws.send(msg["bytes"])
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
await backend_ws.close()
|
||||
|
||||
|
||||
async def _forward_backend_to_client(backend_ws, client_ws: WebSocket) -> None:
|
||||
try:
|
||||
async for frame in backend_ws:
|
||||
if isinstance(frame, bytes):
|
||||
await client_ws.send_bytes(frame)
|
||||
else:
|
||||
await client_ws.send_text(frame)
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
await client_ws.close()
|
||||
|
||||
|
||||
def _websocket_url_for(http_url: str) -> str:
|
||||
if http_url.startswith("https://"):
|
||||
return "wss://" + http_url[len("https://"):]
|
||||
if http_url.startswith("http://"):
|
||||
return "ws://" + http_url[len("http://"):]
|
||||
return http_url
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_router_app",
|
||||
"run_router",
|
||||
]
|
||||
@@ -0,0 +1,268 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Replica registry + health-check loop.
|
||||
|
||||
The registry tracks the set of known backend replicas and their live
|
||||
health. The router consults it for "pick a backend for this session"
|
||||
decisions and a background task updates it from periodic HTTP probes.
|
||||
|
||||
State machine per replica::
|
||||
|
||||
HEALTHY ──(N consecutive failures)──▶ UNHEALTHY
|
||||
▲ │
|
||||
└──────(M consecutive successes)──────┘
|
||||
|
||||
Where N = :attr:`RouterConfig.failure_threshold` and
|
||||
M = :attr:`RouterConfig.recovery_threshold`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import enum
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.config import (
|
||||
ReplicaEndpoint,
|
||||
RouterConfig,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
HttpProbe = Any
|
||||
"""Structural alias for health-probe callables. Concrete signature is
|
||||
``async def __call__(url: str, *, timeout: float) -> tuple[float,
|
||||
str | None]``; typing.Callable cannot express keyword-only parameters,
|
||||
so duck-typing is the pragmatic compromise."""
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ReplicaStatus(enum.Enum):
|
||||
UNKNOWN = "unknown"
|
||||
HEALTHY = "healthy"
|
||||
UNHEALTHY = "unhealthy"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReplicaHealth:
|
||||
status: ReplicaStatus = ReplicaStatus.UNKNOWN
|
||||
last_ok_at: float | None = None
|
||||
last_failure_at: float | None = None
|
||||
consecutive_failures: int = 0
|
||||
consecutive_successes: int = 0
|
||||
last_latency_ms: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Replica:
|
||||
endpoint: ReplicaEndpoint
|
||||
health: ReplicaHealth = field(default_factory=ReplicaHealth)
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
return self.endpoint.url
|
||||
|
||||
@property
|
||||
def primary(self) -> bool:
|
||||
return self.endpoint.primary
|
||||
|
||||
@property
|
||||
def is_healthy(self) -> bool:
|
||||
return self.health.status is ReplicaStatus.HEALTHY
|
||||
|
||||
|
||||
class ReplicaRegistry:
|
||||
"""Stateful map of replica URL → :class:`Replica`.
|
||||
|
||||
Selection favors primary replicas when healthy; otherwise the first
|
||||
healthy non-primary is returned. When none are healthy, the
|
||||
registry returns ``None`` so the router can reject incoming
|
||||
sessions with ``gpu_unavailable``.
|
||||
"""
|
||||
|
||||
def __init__(self, replicas: list[ReplicaEndpoint]) -> None:
|
||||
if not replicas:
|
||||
raise ValueError("ReplicaRegistry requires at least one replica")
|
||||
self._replicas: dict[str, Replica] = {endpoint.url: Replica(endpoint=endpoint) for endpoint in replicas}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
def all(self) -> list[Replica]:
|
||||
return list(self._replicas.values())
|
||||
|
||||
def get(self, url: str) -> Replica | None:
|
||||
return self._replicas.get(url)
|
||||
|
||||
def primaries(self) -> list[Replica]:
|
||||
return [r for r in self._replicas.values() if r.primary]
|
||||
|
||||
def select(self) -> Replica | None:
|
||||
"""Pick the best healthy replica.
|
||||
|
||||
Priority order:
|
||||
|
||||
1. The first healthy primary (insertion order).
|
||||
2. The first healthy non-primary (insertion order).
|
||||
3. ``None`` when nothing is healthy.
|
||||
|
||||
This MVP picks the first match within each tier; it does NOT
|
||||
load-balance across multiple healthy replicas of the same tier.
|
||||
Round-robin and weighted distribution are deferred until a real
|
||||
N-way active deployment exists.
|
||||
"""
|
||||
healthy_primaries = [r for r in self._replicas.values() if r.primary and r.is_healthy]
|
||||
if healthy_primaries:
|
||||
return healthy_primaries[0]
|
||||
healthy = [r for r in self._replicas.values() if r.is_healthy]
|
||||
if healthy:
|
||||
return healthy[0]
|
||||
return None
|
||||
|
||||
async def record_success(
|
||||
self,
|
||||
replica: Replica,
|
||||
*,
|
||||
recovery_threshold: int,
|
||||
latency_ms: float,
|
||||
) -> None:
|
||||
async with self._lock:
|
||||
h = replica.health
|
||||
h.last_ok_at = time.time()
|
||||
h.last_latency_ms = latency_ms
|
||||
h.consecutive_failures = 0
|
||||
h.consecutive_successes += 1
|
||||
# State machine: UNKNOWN -> HEALTHY is immediate; only the
|
||||
# UNHEALTHY -> HEALTHY transition is gated by recovery_threshold.
|
||||
if h.status is ReplicaStatus.UNKNOWN:
|
||||
logger.info("router: replica %s initial probe ok, marking HEALTHY", replica.url)
|
||||
h.status = ReplicaStatus.HEALTHY
|
||||
h.consecutive_successes = 0
|
||||
elif (h.status is ReplicaStatus.UNHEALTHY and h.consecutive_successes >= recovery_threshold):
|
||||
logger.info("router: replica %s recovered to HEALTHY after %d successes", replica.url,
|
||||
h.consecutive_successes)
|
||||
h.status = ReplicaStatus.HEALTHY
|
||||
h.consecutive_successes = 0
|
||||
|
||||
async def record_failure(
|
||||
self,
|
||||
replica: Replica,
|
||||
*,
|
||||
failure_threshold: int,
|
||||
reason: str,
|
||||
) -> None:
|
||||
async with self._lock:
|
||||
h = replica.health
|
||||
h.last_failure_at = time.time()
|
||||
h.consecutive_successes = 0
|
||||
h.consecutive_failures += 1
|
||||
if (h.status is not ReplicaStatus.UNHEALTHY and h.consecutive_failures >= failure_threshold):
|
||||
logger.warning("router: replica %s marked UNHEALTHY after %d failures: %s", replica.url,
|
||||
h.consecutive_failures, reason)
|
||||
h.status = ReplicaStatus.UNHEALTHY
|
||||
|
||||
|
||||
async def run_health_check_loop(
|
||||
registry: ReplicaRegistry,
|
||||
config: RouterConfig,
|
||||
*,
|
||||
stop_event: asyncio.Event,
|
||||
http_get: HttpProbe | None = None,
|
||||
) -> None:
|
||||
"""Poll all replicas' health endpoints in parallel on a fixed interval.
|
||||
|
||||
``http_get`` is pluggable so unit tests can inject a deterministic
|
||||
probe without hitting the network. The default builds a single
|
||||
``httpx.AsyncClient`` shared across the loop's lifetime so the
|
||||
common case (steady polling against a stable replica set) reuses
|
||||
TCP/TLS connections instead of paying handshake cost per probe.
|
||||
|
||||
Probes within one polling cycle run concurrently via ``asyncio.gather``
|
||||
so a slow replica doesn't push the cycle past
|
||||
``health_check_interval_seconds``.
|
||||
"""
|
||||
if http_get is not None:
|
||||
await _run_loop(registry, config, stop_event, http_get)
|
||||
return
|
||||
async with _build_default_probe(config) as probe:
|
||||
await _run_loop(registry, config, stop_event, probe)
|
||||
|
||||
|
||||
async def _run_loop(
|
||||
registry: ReplicaRegistry,
|
||||
config: RouterConfig,
|
||||
stop_event: asyncio.Event,
|
||||
http_get: Callable[..., Awaitable[tuple[float, str | None]]],
|
||||
) -> None:
|
||||
while not stop_event.is_set():
|
||||
replicas = registry.all()
|
||||
results = await asyncio.gather(
|
||||
*[
|
||||
http_get(replica.url + config.health_check_path, timeout=config.health_check_timeout_seconds)
|
||||
for replica in replicas
|
||||
],
|
||||
return_exceptions=True,
|
||||
)
|
||||
for replica, result in zip(replicas, results, strict=True):
|
||||
if isinstance(result, BaseException):
|
||||
await registry.record_failure(
|
||||
replica,
|
||||
failure_threshold=config.failure_threshold,
|
||||
reason=f"{type(result).__name__}: {result}",
|
||||
)
|
||||
continue
|
||||
status_ms, error = result
|
||||
if error is None:
|
||||
await registry.record_success(
|
||||
replica,
|
||||
recovery_threshold=config.recovery_threshold,
|
||||
latency_ms=status_ms,
|
||||
)
|
||||
else:
|
||||
await registry.record_failure(
|
||||
replica,
|
||||
failure_threshold=config.failure_threshold,
|
||||
reason=error,
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
stop_event.wait(),
|
||||
timeout=config.health_check_interval_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _build_default_probe(
|
||||
config: RouterConfig, ) -> AsyncIterator[Callable[..., Awaitable[tuple[float, str | None]]]]:
|
||||
try:
|
||||
import httpx
|
||||
except ImportError as exc: # pragma: no cover - optional extra
|
||||
raise RuntimeError("router health checks require httpx; install with "
|
||||
"`pip install fastvideo[streaming]` or `pip install httpx`") from exc
|
||||
|
||||
async with httpx.AsyncClient(timeout=config.health_check_timeout_seconds) as client:
|
||||
|
||||
async def probe(url: str, *, timeout: float) -> tuple[float, str | None]:
|
||||
start = time.perf_counter()
|
||||
try:
|
||||
response = await client.get(url, timeout=timeout)
|
||||
except Exception as exc:
|
||||
return 0.0, f"{type(exc).__name__}: {exc}"
|
||||
latency_ms = (time.perf_counter() - start) * 1000.0
|
||||
if response.status_code >= 400:
|
||||
return latency_ms, f"HTTP {response.status_code}"
|
||||
return latency_ms, None
|
||||
|
||||
yield probe
|
||||
|
||||
|
||||
__all__ = [
|
||||
"HttpProbe",
|
||||
"Replica",
|
||||
"ReplicaHealth",
|
||||
"ReplicaRegistry",
|
||||
"ReplicaStatus",
|
||||
"run_health_check_loop",
|
||||
]
|
||||
@@ -51,6 +51,11 @@ from fastvideo.entrypoints.streaming.session import (
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image, )
|
||||
from fastvideo.entrypoints.streaming.gpu_pool import (
|
||||
GpuPool,
|
||||
InProcessGpuPool,
|
||||
PoolAcquireTimeout,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
@@ -75,25 +80,37 @@ class _GeneratorProto(Protocol):
|
||||
@dataclass
|
||||
class ServerState:
|
||||
serve_config: ServeConfig
|
||||
generator: _GeneratorProto
|
||||
pool: GpuPool
|
||||
sessions: SessionManager
|
||||
session_store: SessionStore
|
||||
|
||||
|
||||
def build_app(
|
||||
serve_config: ServeConfig,
|
||||
generator: _GeneratorProto,
|
||||
generator: _GeneratorProto | None = None,
|
||||
*,
|
||||
pool: GpuPool | None = None,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI app used by :func:`run_server`.
|
||||
|
||||
Exposed so tests can drive the WebSocket endpoint in-process via
|
||||
``starlette.testclient.TestClient(app).websocket_connect(...)``.
|
||||
|
||||
Exactly one of ``generator`` (backed by :class:`InProcessGpuPool`)
|
||||
or ``pool`` (for the subprocess-backed production shape) must be
|
||||
given.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
|
||||
"server; got None. Add a `streaming:` block to your serve config.")
|
||||
if (generator is None) == (pool is None):
|
||||
raise ValueError("build_app requires exactly one of `generator` or `pool`")
|
||||
|
||||
store = session_store or InMemorySessionStore()
|
||||
if pool is None:
|
||||
assert generator is not None
|
||||
pool = InProcessGpuPool(generator, session_store=store)
|
||||
|
||||
sessions = SessionManager(
|
||||
segment_cap=serve_config.streaming.generation_segment_cap,
|
||||
@@ -101,9 +118,9 @@ def build_app(
|
||||
)
|
||||
state = ServerState(
|
||||
serve_config=serve_config,
|
||||
generator=generator,
|
||||
pool=pool,
|
||||
sessions=sessions,
|
||||
session_store=session_store or InMemorySessionStore(),
|
||||
session_store=store,
|
||||
)
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming")
|
||||
@@ -135,6 +152,8 @@ def build_app(
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
await state.pool.release(session.id)
|
||||
_cleanup_session(session, state)
|
||||
|
||||
app.state.server_state = state
|
||||
@@ -178,10 +197,22 @@ async def _handle_session(
|
||||
await _apply_session_init(session, init, state)
|
||||
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
|
||||
session.transition(SessionState.GPU_BINDING)
|
||||
await _send_json(websocket, GpuAssigned(
|
||||
gpu_id=0,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
try:
|
||||
assignment = await state.pool.acquire(
|
||||
session.id,
|
||||
timeout=float(state.sessions.session_timeout_seconds),
|
||||
)
|
||||
except PoolAcquireTimeout as exc:
|
||||
await _send_error(websocket, "gpu_unavailable", str(exc), retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return
|
||||
session.gpu_id = assignment.gpu_id
|
||||
await _send_json(websocket,
|
||||
GpuAssigned(
|
||||
gpu_id=assignment.gpu_id,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
session.transition(SessionState.ACTIVE)
|
||||
await _send_json(websocket, _build_stream_start(session, state))
|
||||
|
||||
@@ -327,15 +358,13 @@ async def _run_segment(
|
||||
))
|
||||
|
||||
start = time.perf_counter()
|
||||
loop = asyncio.get_running_loop()
|
||||
# TODO: executor-wrapped generate() cannot be cancelled, so a
|
||||
# client disconnect mid-segment leaves the GPU work running to
|
||||
# completion. Real cancellation needs the generate_async API.
|
||||
# TODO: pool.run() runs to completion even if the client disconnects
|
||||
# mid-segment. Real cancellation needs the generate_async API.
|
||||
try:
|
||||
result = await loop.run_in_executor(None, state.generator.generate, request)
|
||||
result = await state.pool.run(session.id, request)
|
||||
except Exception as exc:
|
||||
logger.exception("session %s: generator failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
|
||||
logger.exception("session %s: pool.run failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"pool.run failed: {exc}", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-session JSONL event logger.
|
||||
|
||||
Each session gets its own JSONL file under the configured log root so
|
||||
post-hoc analytics (enhancer latency, GPU assignment, segment timings)
|
||||
can be recovered without a tracing backend. The internal UI uses this
|
||||
format; keeping the same shape makes log tooling portable.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, TextIO
|
||||
|
||||
_FILENAME_SANITIZE_RE = re.compile(r"[^A-Za-z0-9._-]")
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionLogEvent:
|
||||
"""One line in the session JSONL file."""
|
||||
|
||||
session_id: str
|
||||
event: str
|
||||
payload: dict[str, Any] = field(default_factory=dict)
|
||||
ts: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
class SessionLogger:
|
||||
"""Append-only JSONL logger keyed by session id.
|
||||
|
||||
Thread-safe; the server may be writing from multiple asyncio tasks
|
||||
(fMP4 encoder thread + control-frame handler) for the same session.
|
||||
"""
|
||||
|
||||
def __init__(self, log_dir: str | None) -> None:
|
||||
self._log_dir = log_dir
|
||||
self._files: dict[str, TextIO] = {}
|
||||
self._locks: dict[str, threading.Lock] = {}
|
||||
self._registry_lock = threading.Lock()
|
||||
self._ensure_dir()
|
||||
|
||||
def log(self, event: SessionLogEvent) -> None:
|
||||
if self._log_dir is None:
|
||||
return
|
||||
opened = self._get_file(event.session_id)
|
||||
if opened is None:
|
||||
return
|
||||
handle, lock = opened
|
||||
line = json.dumps({
|
||||
"session_id": event.session_id,
|
||||
"event": event.event,
|
||||
"ts": event.ts,
|
||||
"payload": event.payload,
|
||||
})
|
||||
with lock, contextlib.suppress(ValueError):
|
||||
handle.write(line + "\n")
|
||||
handle.flush()
|
||||
|
||||
def close(self, session_id: str) -> None:
|
||||
with self._registry_lock:
|
||||
handle = self._files.pop(session_id, None)
|
||||
lock = self._locks.pop(session_id, None)
|
||||
if handle is None or lock is None:
|
||||
return
|
||||
with lock, contextlib.suppress(Exception):
|
||||
handle.close()
|
||||
|
||||
def close_all(self) -> None:
|
||||
with self._registry_lock:
|
||||
sids = list(self._files)
|
||||
for sid in sids:
|
||||
self.close(sid)
|
||||
|
||||
def _ensure_dir(self) -> None:
|
||||
if self._log_dir is None:
|
||||
return
|
||||
os.makedirs(self._log_dir, exist_ok=True)
|
||||
|
||||
def _get_file(self, session_id: str) -> tuple[TextIO, threading.Lock] | None:
|
||||
if self._log_dir is None:
|
||||
return None
|
||||
with self._registry_lock:
|
||||
handle = self._files.get(session_id)
|
||||
lock = self._locks.get(session_id)
|
||||
if handle is not None and lock is not None:
|
||||
return handle, lock
|
||||
# Defense-in-depth: session_id is server-generated UUID today,
|
||||
# but sanitize against path traversal in case future code paths
|
||||
# allow client-supplied ids.
|
||||
safe_id = _FILENAME_SANITIZE_RE.sub("_", session_id) or "unknown"
|
||||
path = os.path.join(
|
||||
self._log_dir,
|
||||
f"session-{safe_id}.jsonl",
|
||||
)
|
||||
try:
|
||||
handle = open(path, "a", encoding="utf-8") # noqa: SIM115
|
||||
except OSError:
|
||||
return None
|
||||
lock = threading.Lock()
|
||||
self._files[session_id] = handle
|
||||
self._locks[session_id] = lock
|
||||
return handle, lock
|
||||
|
||||
|
||||
__all__ = [
|
||||
"SessionLogEvent",
|
||||
"SessionLogger",
|
||||
]
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-GPU worker subprocess entry for :class:`SubprocessGpuPool`.
|
||||
|
||||
The pool manages binding, lifecycle, and message dispatch in the parent
|
||||
process. The worker constructs its :class:`VideoGenerator` from a typed
|
||||
:class:`GeneratorConfig`, runs the two-segment warmup so both
|
||||
initial-segment and continuation-branch compile graphs are hot, and
|
||||
then loops on the job queue.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import multiprocessing as mp
|
||||
import queue
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
GeneratorConfig,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Synthetic warmup dimensions: small enough to keep boot fast, big enough
|
||||
# to exercise the real shape-dependent compile paths. Keep in sync with
|
||||
# WarmupConfig if these become user-tunable.
|
||||
_WARMUP_NUM_FRAMES = 8
|
||||
_WARMUP_HEIGHT = 256
|
||||
_WARMUP_WIDTH = 256
|
||||
_WARMUP_NUM_INFERENCE_STEPS = 1
|
||||
|
||||
|
||||
def worker_main(
|
||||
*,
|
||||
gpu_id: int,
|
||||
worker_id: str,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
job_queue: mp.Queue,
|
||||
result_queue: mp.Queue,
|
||||
shutdown_event: Any,
|
||||
) -> None: # pragma: no cover - exercised via integration only
|
||||
"""Per-worker subprocess entry.
|
||||
|
||||
Runs inside the child spawned by ``SubprocessGpuPool``. Blocking
|
||||
``VideoGenerator`` construction + generation happens here, not in
|
||||
the parent's event loop.
|
||||
"""
|
||||
import os
|
||||
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id)
|
||||
try:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
generator = VideoGenerator.from_pretrained(config=generator_config)
|
||||
if warmup_config.enabled:
|
||||
_warmup_worker(generator, warmup_config)
|
||||
result_queue.put({"kind": "ready", "worker_id": worker_id})
|
||||
except Exception as exc:
|
||||
result_queue.put({"kind": "error", "error": repr(exc)})
|
||||
return
|
||||
|
||||
while not shutdown_event.is_set():
|
||||
try:
|
||||
item = job_queue.get(timeout=0.5)
|
||||
except queue.Empty:
|
||||
continue
|
||||
if item is None:
|
||||
break
|
||||
job_id = item["job_id"]
|
||||
request = item["request"]
|
||||
try:
|
||||
result = generator.generate(request)
|
||||
result_queue.put({
|
||||
"kind": "result",
|
||||
"job_id": job_id,
|
||||
"result": result,
|
||||
})
|
||||
except Exception as exc:
|
||||
result_queue.put({
|
||||
"kind": "error",
|
||||
"job_id": job_id,
|
||||
"error": repr(exc),
|
||||
})
|
||||
|
||||
|
||||
def _warmup_worker(
|
||||
generator: Any,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> None:
|
||||
"""Run two synthetic generations so both compile branches are primed.
|
||||
|
||||
Segment 1 is a fresh start (no continuation state) and exercises
|
||||
the initial-segment graph. Segment 2 feeds segment 1's continuation
|
||||
state back in so the conditioning branch is also compiled before
|
||||
the first user request lands.
|
||||
"""
|
||||
sampling = SamplingConfig(
|
||||
num_frames=_WARMUP_NUM_FRAMES,
|
||||
height=_WARMUP_HEIGHT,
|
||||
width=_WARMUP_WIDTH,
|
||||
num_inference_steps=_WARMUP_NUM_INFERENCE_STEPS,
|
||||
)
|
||||
seg1 = GenerationRequest(
|
||||
prompt=warmup_config.prompt,
|
||||
sampling=sampling,
|
||||
inputs=InputConfig(),
|
||||
output=OutputConfig(save_video=False, return_frames=False, return_state=True),
|
||||
)
|
||||
seg1_result = generator.generate(seg1)
|
||||
|
||||
seg2 = GenerationRequest(
|
||||
prompt=warmup_config.prompt,
|
||||
sampling=sampling,
|
||||
inputs=InputConfig(),
|
||||
output=OutputConfig(save_video=False, return_frames=False),
|
||||
state=_extract_continuation_state(seg1_result),
|
||||
)
|
||||
generator.generate(seg2)
|
||||
|
||||
|
||||
def _extract_continuation_state(result: Any) -> Any:
|
||||
state = getattr(result, "state", None)
|
||||
if state is None and isinstance(result, dict):
|
||||
state = result.get("state")
|
||||
return state
|
||||
|
||||
|
||||
__all__ = ["worker_main"]
|
||||
@@ -65,6 +65,7 @@ _FROM_PRETRAINED_CONVENIENCE_KWARGS = frozenset({
|
||||
"pin_cpu_memory",
|
||||
"enable_torch_compile",
|
||||
"torch_compile_kwargs",
|
||||
"output_type",
|
||||
})
|
||||
|
||||
|
||||
@@ -601,10 +602,20 @@ class VideoGenerator:
|
||||
thread = threading.Thread(target=execute_forward_thread)
|
||||
thread.start()
|
||||
latent_batch_size = _infer_latent_batch_size(batch)
|
||||
samples = torch.empty(
|
||||
(latent_batch_size, 3, sampling_param.num_frames, sampling_param.height, sampling_param.width),
|
||||
device='cpu',
|
||||
pin_memory=fastvideo_args.pin_cpu_memory)
|
||||
# When ``output_type == "latent"`` the forward output has latent
|
||||
# shape (e.g. ``[B, C_latent, T_latent, H_latent, W_latent]``)
|
||||
# rather than the pre-allocation's pixel shape. Skip the pinned
|
||||
# ~50 MB buffer entirely; we always fall through to the
|
||||
# ``samples = output_batch.output.cpu()`` branch below in that
|
||||
# mode. ``skip_pixel_prealloc`` also gates the slow-path warning.
|
||||
skip_pixel_prealloc = fastvideo_args.output_type == "latent"
|
||||
if skip_pixel_prealloc:
|
||||
samples = torch.empty(0, device='cpu')
|
||||
else:
|
||||
samples = torch.empty(
|
||||
(latent_batch_size, 3, sampling_param.num_frames, sampling_param.height, sampling_param.width),
|
||||
device='cpu',
|
||||
pin_memory=fastvideo_args.pin_cpu_memory)
|
||||
thread.join()
|
||||
|
||||
if thread_error["error"] is not None:
|
||||
@@ -619,29 +630,44 @@ class VideoGenerator:
|
||||
if output_batch.output.shape == samples.shape:
|
||||
samples.copy_(output_batch.output)
|
||||
else:
|
||||
logger.warning("Output shape %s does not match expected shape %s; use slow path", output_batch.output.shape,
|
||||
samples.shape)
|
||||
if not skip_pixel_prealloc:
|
||||
logger.warning("Output shape %s does not match expected shape %s; use slow path",
|
||||
output_batch.output.shape, samples.shape)
|
||||
samples = output_batch.output.cpu()
|
||||
logging_info = output_batch.logging_info
|
||||
|
||||
gen_time = time.perf_counter() - start_time
|
||||
logger.info("Generated successfully in %.2f seconds", gen_time)
|
||||
|
||||
# Process outputs (skip the make_grid loop for audio-only, where
|
||||
# `samples` is a 1×3×1×8×8 placeholder no caller will use).
|
||||
# Three mutually-exclusive output modes determine whether (a) we
|
||||
# build an RGB frame buffer and (b) what file we write to disk:
|
||||
#
|
||||
# 1. `output_type == "latent"` — VAE is bypassed in DecodingStage
|
||||
# and `samples` holds raw latents (arbitrary channel count).
|
||||
# The RGB grid / uint8 / mp4 / png pipeline below cannot
|
||||
# consume those, so we skip it entirely and let callers work
|
||||
# with the latent tensor directly via `result["samples"]`.
|
||||
# 2. Audio-only workload — `samples` is a 1×3×1×8×8 placeholder
|
||||
# no caller will use; skip the grid loop and save a `.wav`.
|
||||
# 3. Pixel video / image — the historical happy path.
|
||||
is_latent_output = fastvideo_args.output_type == "latent"
|
||||
audio_only = bool(output_batch.extra.get("audio_only"))
|
||||
frames: list[np.ndarray] = []
|
||||
if not audio_only:
|
||||
|
||||
frames: list[np.ndarray] | None
|
||||
if is_latent_output or audio_only:
|
||||
frames = None if is_latent_output else []
|
||||
else:
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.permute(1, 2, 0).squeeze(-1)
|
||||
x = (x * 255).to(torch.uint8)
|
||||
frames.append(x.cpu().numpy())
|
||||
|
||||
# Save output if requested
|
||||
if batch.save_video:
|
||||
if output_batch.extra.get("audio_only"):
|
||||
save_to_disk = batch.save_video and not is_latent_output
|
||||
if save_to_disk:
|
||||
if audio_only:
|
||||
# Audio-only workload: write a standalone .wav rather than
|
||||
# muxing the audio into a placeholder mp4 (which forces
|
||||
# ffmpeg to round 8x8 placeholder frames up to 16x16).
|
||||
@@ -654,9 +680,11 @@ class VideoGenerator:
|
||||
logger.info("Saved audio to %s", output_path)
|
||||
elif self._is_image_workload():
|
||||
# Image workloads (t2i, i2i, …): save the first frame as PNG.
|
||||
assert frames is not None # implied by save_to_disk and not audio_only
|
||||
imageio.imwrite(output_path, frames[0])
|
||||
logger.info("Saved image to %s", output_path)
|
||||
else:
|
||||
assert frames is not None # implied by save_to_disk and not audio_only
|
||||
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
audio = output_batch.extra.get("audio")
|
||||
@@ -680,7 +708,7 @@ class VideoGenerator:
|
||||
"trajectory": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
"video_path": output_path if batch.save_video else None,
|
||||
"video_path": output_path if save_to_disk else None,
|
||||
"peak_memory_mb": output_batch.extra.get("peak_memory_mb"),
|
||||
}
|
||||
|
||||
@@ -759,7 +787,7 @@ class VideoGenerator:
|
||||
import av
|
||||
except ImportError:
|
||||
logger.warning("PyAV not installed; cannot mux audio. "
|
||||
"Install with: pip install av")
|
||||
"Install with: uv pip install av")
|
||||
return False
|
||||
|
||||
try:
|
||||
|
||||
@@ -35,6 +35,11 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS: int = 1
|
||||
FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS: int = 2
|
||||
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
|
||||
FASTVIDEO_TRACE_ACTIVATIONS: bool = False
|
||||
FASTVIDEO_TRACE_LAYERS: str = ""
|
||||
FASTVIDEO_TRACE_STATS: str = "abs_mean,sum"
|
||||
FASTVIDEO_TRACE_OUTPUT: str = "/tmp/fv_trace_<pid>.jsonl"
|
||||
FASTVIDEO_TRACE_STEPS: str = ""
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
FASTVIDEO_STAGE_LOGGING: bool = False
|
||||
FASTVIDEO_HOST_IP: str = ""
|
||||
@@ -252,6 +257,22 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_TORCH_PROFILE_REGIONS":
|
||||
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
|
||||
|
||||
# Enable activation trace hooks if set.
|
||||
"FASTVIDEO_TRACE_ACTIVATIONS":
|
||||
lambda: bool(os.getenv("FASTVIDEO_TRACE_ACTIVATIONS", "0") != "0"),
|
||||
# Regex filter for traced module names. Empty means all modules.
|
||||
"FASTVIDEO_TRACE_LAYERS":
|
||||
lambda: os.getenv("FASTVIDEO_TRACE_LAYERS", ""),
|
||||
# Comma-separated activation stats to dump for each output tensor.
|
||||
"FASTVIDEO_TRACE_STATS":
|
||||
lambda: os.getenv("FASTVIDEO_TRACE_STATS", "abs_mean,sum"),
|
||||
# JSONL sink path. The literal <pid> is replaced at runtime.
|
||||
"FASTVIDEO_TRACE_OUTPUT":
|
||||
lambda: os.getenv("FASTVIDEO_TRACE_OUTPUT", "/tmp/fv_trace_<pid>.jsonl"),
|
||||
# Comma-separated denoise step indices. Empty means all steps.
|
||||
"FASTVIDEO_TRACE_STEPS":
|
||||
lambda: os.getenv("FASTVIDEO_TRACE_STEPS", ""),
|
||||
|
||||
# If set, fastvideo will run in development mode, which will enable
|
||||
# some additional endpoints for developing and debugging,
|
||||
# e.g. `/reset_prefix_cache`
|
||||
|
||||
@@ -0,0 +1,221 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Zero-overhead-when-off activation trace mode for FastVideo pipelines.
|
||||
|
||||
Enable by setting FASTVIDEO_TRACE_ACTIVATIONS=1. When off, this module
|
||||
adds zero overhead — no hooks are registered, no branches exist in the
|
||||
production forward path. When on, registers forward hooks on modules
|
||||
whose name matches FASTVIDEO_TRACE_LAYERS, computes the requested stats
|
||||
(FASTVIDEO_TRACE_STATS) on each output tensor, and writes JSONL records
|
||||
to FASTVIDEO_TRACE_OUTPUT.
|
||||
|
||||
Useful for parity debugging across model ports — log on both the
|
||||
FastVideo path and the upstream reference, diff the two JSONL files
|
||||
to find the first divergent layer.
|
||||
|
||||
Example:
|
||||
|
||||
FASTVIDEO_TRACE_ACTIVATIONS=1 \
|
||||
FASTVIDEO_TRACE_LAYERS="^block\\.layers\\.[0-9]+$" \
|
||||
FASTVIDEO_TRACE_STATS="abs_mean,max,shape" \
|
||||
FASTVIDEO_TRACE_OUTPUT="/tmp/fv_trace.jsonl" \
|
||||
python examples/inference/basic/basic_magi_human.py
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from collections.abc import Callable, Iterator
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_TRACE_STATE = threading.local()
|
||||
|
||||
|
||||
def current_step_idx() -> int | None:
|
||||
return getattr(_TRACE_STATE, "step_idx", None)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def trace_step(step_idx: int) -> Iterator[None]:
|
||||
"""Context manager that sets the current denoise step for trace records."""
|
||||
prev = getattr(_TRACE_STATE, "step_idx", None)
|
||||
_TRACE_STATE.step_idx = step_idx
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_TRACE_STATE.step_idx = prev
|
||||
|
||||
|
||||
_STAT_FNS: dict[str, Callable[[torch.Tensor], Any]] = {
|
||||
"abs_mean": lambda t: float(t.detach().float().abs().mean().item()),
|
||||
"sum": lambda t: float(t.detach().float().sum().item()),
|
||||
"min": lambda t: float(t.detach().float().min().item()),
|
||||
"max": lambda t: float(t.detach().float().max().item()),
|
||||
"mean": lambda t: float(t.detach().float().mean().item()),
|
||||
"std": lambda t: float(t.detach().float().std().item()),
|
||||
"shape": lambda t: list(t.shape),
|
||||
"dtype": lambda t: str(t.dtype),
|
||||
}
|
||||
|
||||
|
||||
def _resolve_stats(spec: str) -> list[tuple[str, Callable[[torch.Tensor], Any]]]:
|
||||
stats = []
|
||||
for name in [s.strip() for s in spec.split(",") if s.strip()]:
|
||||
stat_fn = _STAT_FNS.get(name)
|
||||
if stat_fn is None:
|
||||
logger.warning(
|
||||
"FASTVIDEO_TRACE_STATS contains unknown stat %r; valid: %s",
|
||||
name,
|
||||
sorted(_STAT_FNS),
|
||||
)
|
||||
continue
|
||||
stats.append((name, stat_fn))
|
||||
return stats
|
||||
|
||||
|
||||
def _resolve_output_path(template: str) -> Path:
|
||||
return Path(template.replace("<pid>", str(os.getpid())))
|
||||
|
||||
|
||||
def _parse_step_filter(spec: str) -> set[int] | None:
|
||||
if not spec.strip():
|
||||
return None
|
||||
return {int(s.strip()) for s in spec.split(",") if s.strip()}
|
||||
|
||||
|
||||
class JsonlSink:
|
||||
"""Buffered JSONL writer with thread-safe append."""
|
||||
|
||||
def __init__(self, path: Path) -> None:
|
||||
self.path = path
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._fh = open(self.path, "a", buffering=1) # noqa: SIM115
|
||||
self._lock = threading.Lock()
|
||||
logger.info("Activation trace JSONL sink: %s", self.path)
|
||||
|
||||
def write(self, record: dict[str, Any]) -> None:
|
||||
line = json.dumps(record, default=str) + "\n"
|
||||
with self._lock:
|
||||
self._fh.write(line)
|
||||
|
||||
def close(self) -> None:
|
||||
with self._lock:
|
||||
if not self._fh.closed:
|
||||
self._fh.close()
|
||||
|
||||
|
||||
class ActivationStatHook(ForwardHook):
|
||||
"""Forward hook that emits per-tensor stats to a JSONL sink."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
module_name: str,
|
||||
stats: list[tuple[str, Callable[[torch.Tensor], Any]]],
|
||||
sink: JsonlSink,
|
||||
step_filter: set[int] | None,
|
||||
) -> None:
|
||||
self.module_name = module_name
|
||||
self.stats = stats
|
||||
self.sink = sink
|
||||
self.step_filter = step_filter
|
||||
|
||||
def name(self) -> str:
|
||||
return "ActivationStatHook"
|
||||
|
||||
def post_forward(self, module: nn.Module, output: Any) -> Any:
|
||||
step_idx = current_step_idx()
|
||||
if self.step_filter is not None and step_idx not in self.step_filter:
|
||||
return output
|
||||
for tensor_label, tensor in _flatten_tensors(output):
|
||||
record: dict[str, Any] = {
|
||||
"module": self.module_name,
|
||||
"tensor": tensor_label,
|
||||
"step": step_idx,
|
||||
}
|
||||
for stat_name, stat_fn in self.stats:
|
||||
try:
|
||||
record[stat_name] = stat_fn(tensor)
|
||||
except Exception as exc: # pragma: no cover - defensive logging
|
||||
record[stat_name] = f"<error: {exc!r}>"
|
||||
self.sink.write(record)
|
||||
return output
|
||||
|
||||
|
||||
def _flatten_tensors(obj: Any, prefix: str = "out") -> list[tuple[str, torch.Tensor]]:
|
||||
"""Yield (label, tensor) pairs from arbitrarily-nested forward outputs."""
|
||||
if isinstance(obj, torch.Tensor):
|
||||
return [(prefix, obj)]
|
||||
if isinstance(obj, tuple | list):
|
||||
out = []
|
||||
for idx, item in enumerate(obj):
|
||||
out.extend(_flatten_tensors(item, f"{prefix}[{idx}]"))
|
||||
return out
|
||||
if isinstance(obj, dict):
|
||||
out = []
|
||||
for key, value in obj.items():
|
||||
out.extend(_flatten_tensors(value, f"{prefix}.{key}"))
|
||||
return out
|
||||
return []
|
||||
|
||||
|
||||
class ActivationTraceManager:
|
||||
|
||||
def __init__(self, managers: list[ModuleHookManager], sink: JsonlSink) -> None:
|
||||
self.managers = managers
|
||||
self.sink = sink
|
||||
|
||||
def remove_from_manager(self) -> None:
|
||||
for manager in self.managers:
|
||||
if manager.get_forward_hook("ActivationStatHook") is not None:
|
||||
manager.remove_forward_hook("ActivationStatHook")
|
||||
if not manager.forward_hooks:
|
||||
ModuleHookManager.remove_from_manager(manager.module)
|
||||
self.sink.close()
|
||||
|
||||
|
||||
def attach_activation_trace(model: nn.Module | None) -> ActivationTraceManager | None:
|
||||
"""Attach activation-stat hooks to model. Returns None if trace is off."""
|
||||
if not envs.FASTVIDEO_TRACE_ACTIVATIONS or model is None:
|
||||
return None
|
||||
|
||||
pattern_spec = envs.FASTVIDEO_TRACE_LAYERS
|
||||
pattern = re.compile(pattern_spec) if pattern_spec else re.compile(".*")
|
||||
stats = _resolve_stats(envs.FASTVIDEO_TRACE_STATS)
|
||||
if not stats:
|
||||
logger.warning("FASTVIDEO_TRACE_STATS yielded no valid stats; trace disabled.")
|
||||
return None
|
||||
|
||||
sink = JsonlSink(_resolve_output_path(envs.FASTVIDEO_TRACE_OUTPUT))
|
||||
step_filter = _parse_step_filter(envs.FASTVIDEO_TRACE_STEPS)
|
||||
managers = []
|
||||
for name, module in model.named_modules():
|
||||
if not name or not pattern.search(name):
|
||||
continue
|
||||
manager = ModuleHookManager.get_from_or_default(module)
|
||||
manager.append_forward_hook(ActivationStatHook(name, stats, sink, step_filter))
|
||||
managers.append(manager)
|
||||
|
||||
logger.info(
|
||||
"Activation trace attached to %d modules (pattern=%r, stats=%s)",
|
||||
len(managers),
|
||||
pattern_spec,
|
||||
[stat_name for stat_name, _ in stats],
|
||||
)
|
||||
return ActivationTraceManager(managers, sink)
|
||||
|
||||
|
||||
def detach_activation_trace(mgr: ActivationTraceManager | None) -> None:
|
||||
if mgr is not None:
|
||||
mgr.remove_from_manager()
|
||||
@@ -0,0 +1,53 @@
|
||||
# Layer Guidance For Model Ports
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Use this file when adding FastVideo-native model components. Keep it generic:
|
||||
model-specific parameter mappings belong in `scripts/checkpoint_conversion/`, not
|
||||
in this directory.
|
||||
|
||||
## Linear Layers
|
||||
|
||||
- Use `ReplicatedLinear` for DiT and VAE hot paths when the layer is not tensor
|
||||
parallel and should expose a normal `weight`/`bias` state-dict surface.
|
||||
- Use `QKVParallelLinear` for LLM-style fused query/key/value projections when
|
||||
the existing encoder pattern already expects tensor parallel loading.
|
||||
- Use `MergedColumnParallelLinear` for fused MLP gate/up projections that are
|
||||
loaded as packed column shards.
|
||||
- Use `ColumnParallelLinear` and `RowParallelLinear` for tensor-parallel encoder
|
||||
blocks that follow existing `t5.py`, `clip.py`, `llama.py`, or `qwen2_5.py`
|
||||
patterns.
|
||||
- Do not replace a simple official layer with a fused FastVideo layer unless the
|
||||
conversion script explicitly handles the resulting key and tensor layout.
|
||||
|
||||
## Attention Layers
|
||||
|
||||
- Use `DistributedAttention` for standard DiT full-sequence attention when the
|
||||
model should participate in sequence parallel execution.
|
||||
- Use `LocalAttention` for local/window attention or narrow single-GPU parity
|
||||
paths that match existing component style.
|
||||
- Raw `torch.nn.functional.scaled_dot_product_attention` is acceptable for
|
||||
unusual cross-modality flat streams when no FastVideo distributed primitive
|
||||
matches yet. Document the sequence-parallel gap in the owning model file.
|
||||
|
||||
## State-Dict Surface
|
||||
|
||||
- Prototype the native component before writing conversion mappings. The
|
||||
prototype's `state_dict()` is the source of truth for FastVideo target keys and
|
||||
shapes.
|
||||
- Conversion scripts should map official keys into the native state-dict surface;
|
||||
production model code should not be contorted to match checkpoint naming.
|
||||
- Fused and packed FastVideo layers may require tensor split/fuse logic in the
|
||||
converter, especially QKV/KV projections and gated MLP projections.
|
||||
- Record intentional skipped keys in the conversion script with a reason, such
|
||||
as training-only EMA/logvar/optimizer state or dynamically computed buffers.
|
||||
|
||||
## Porting Discipline
|
||||
|
||||
- Match the official layer definition and the official instantiation arguments.
|
||||
A reusable class with different constructor args is not reused.
|
||||
- Keep architecture constants on the component arch config. Runtime sampling,
|
||||
guidance, precision, and pipeline defaults belong on pipeline config or
|
||||
presets.
|
||||
- Prefer small, direct implementations until parity passes. Add helpers only
|
||||
when they serve multiple call sites or make the mapping clearer.
|
||||
@@ -0,0 +1,55 @@
|
||||
# `fastvideo/models/` — Model Implementations
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
DiT / VAE / encoder / scheduler / upsampler / audio model classes. **Pre-commit excludes this directory** — yapf/ruff/mypy do not run on commits here. Match neighboring file style manually.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
models/
|
||||
├── dits/
|
||||
│ ├── <model>.py # Single-file DiT (wanvideo, ltx2, hunyuanvideo, cosmos, ...)
|
||||
│ ├── hyworld/ # Multi-file DiT family
|
||||
│ ├── lingbotworld/ # ditto
|
||||
│ └── matrixgame/ # ditto
|
||||
├── vaes/ # AutoencoderKL variants per model family
|
||||
├── encoders/ # T5, CLIP, Llama, Qwen2.5, Gemma, SigLIP, Reason1, audio conditioner
|
||||
├── schedulers/ # FlowMatch / EulerDiscrete / DPM custom schedulers
|
||||
├── upsamplers/ # Hunyuan15 super-resolution
|
||||
├── audio/ # Audio-VAE/decoder modules (LTX-2 audio, Stable Audio)
|
||||
├── camera/ # Camera-conditioning modules (Gen3C)
|
||||
└── loader/ # component_loader.py, fsdp_load.py, weight_utils.py
|
||||
```
|
||||
|
||||
`loader/component_loader.py` is the central entry point that the pipeline uses
|
||||
to instantiate model components from a HF directory. New components plug in
|
||||
through `register_*` calls or by extending the `ComponentLoader` mappings.
|
||||
|
||||
## Adding a Model Component (DiT / VAE / Encoder)
|
||||
|
||||
1. Read `fastvideo/layers/AGENTS.md` first — it defines which tensor-parallel
|
||||
linear / attention layer to use. Do not freelance.
|
||||
2. Define the arch in `models/<role>/<model>.py`. Mirror the official reference's
|
||||
constructor args; do not "improve" the layer choices.
|
||||
3. Add the matching arch config in `configs/models/<role>/<model>.py`.
|
||||
4. Expose `param_names_mapping` on the config — it is the **source of truth** for
|
||||
the converter under `scripts/checkpoint_conversion/`.
|
||||
5. Use `init_logger(__name__)`, not stdlib logging.
|
||||
|
||||
## State-Dict Discipline
|
||||
|
||||
- The native component's `state_dict()` defines target keys + shapes.
|
||||
- Conversion scripts (`scripts/checkpoint_conversion/`) bend to the model, not
|
||||
the other way around.
|
||||
- Fused QKV / packed MLP layouts must be documented in the config or the model
|
||||
module — converters need to split/fuse accordingly.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Importing `transformers` / `diffusers` model classes at runtime inside the
|
||||
forward path — these belong in the loader, not the architecture file.
|
||||
- Adding training-only state (EMA buffers, optimizer state) to the inference
|
||||
state-dict surface.
|
||||
- Calling `torch.distributed` directly. Go through `fastvideo.distributed`.
|
||||
- Treating this directory as lint-clean. It isn't (see pre-commit excludes).
|
||||
@@ -0,0 +1,867 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""daVinci-MagiHuman DiT (base variant).
|
||||
|
||||
Ported from https://github.com/GAIR-NLP/daVinci-MagiHuman
|
||||
(inference/model/dit/dit_module.py, ~950 lines in the reference).
|
||||
|
||||
Architecture summary (verified against GAIR/daVinci-MagiHuman/base/ weights):
|
||||
|
||||
- 40 transformer layers, hidden 5120, head_dim 128.
|
||||
- GQA with 40 query heads and 8 KV heads.
|
||||
- Multi-modality "sandwich": layers 0..3 and 36..39 use 3-way modality
|
||||
experts (video/audio/text) packed inside each linear as
|
||||
weight[..., out * 3, in]. Middle layers share a single expert.
|
||||
- Per-head attention gating: the QKV projection emits an extra
|
||||
num_heads_q channels that are sigmoid-gated onto the attention output.
|
||||
- Activation is GELU7 on layers 0..3 (non-gated, intermediate=4*hidden)
|
||||
and SwiGLU7 elsewhere (gated, intermediate=int(hidden*4*2/3)//4*4).
|
||||
- Position encoding is an element-wise Fourier embedding over 9-column
|
||||
coords (t,h,w + original TxHxW + reference TxHxW), not a standard
|
||||
1D/3D RoPE.
|
||||
- Forward takes a flat concatenated token stream (video first, then
|
||||
audio, then text) plus a modality map; the internal ModalityDispatcher
|
||||
permutes by modality before each linear so per-expert chunks line up.
|
||||
|
||||
Deviations from the "use fastvideo.layers primitives everywhere" guideline
|
||||
in the add-model skill:
|
||||
|
||||
- The packed-expert linears store weight as [out * num_experts, in].
|
||||
FastVideo's ReplicatedLinear does not model this layout; we use raw
|
||||
nn.Parameter with a small wrapper below. This is deliberate and scoped
|
||||
to this DiT: ReplicatedLinear still handles the adapter.* embedders
|
||||
and final_linear_{video,audio} (single-expert) here.
|
||||
- Self-attention is full-sequence and crosses modalities inside the flat
|
||||
concat stream; DistributedAttention assumes a clean spatial-sequence
|
||||
layout, so for the first port we use torch SDPA. Multi-GPU sequence
|
||||
parallelism is a follow-up.
|
||||
- torch.compile via magi_compiler is replaced with a plain nn.Module.
|
||||
|
||||
For the full history and shape-by-shape verification notes, see
|
||||
.claude/skills/add-model/SKILL.md and the scaffold PR description.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.configs.models.dits.magi_human import (
|
||||
MagiHumanArchConfig,
|
||||
MagiHumanVideoConfig,
|
||||
)
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Enums
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class Modality(IntEnum):
|
||||
VIDEO = 0
|
||||
AUDIO = 1
|
||||
TEXT = 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Activations
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def swiglu7(x: torch.Tensor, alpha: float = 1.702, limit: float = 7.0) -> torch.Tensor:
|
||||
"""Gated swish-GLU with OpenAI-OSS-style limits and +1 linear bias."""
|
||||
in_dtype = x.dtype
|
||||
x = x.to(torch.float32)
|
||||
x_glu, x_linear = x[..., ::2], x[..., 1::2]
|
||||
x_glu = x_glu.clamp(max=limit)
|
||||
x_linear = x_linear.clamp(min=-limit, max=limit)
|
||||
out_glu = x_glu * torch.sigmoid(alpha * x_glu)
|
||||
return (out_glu * (x_linear + 1)).to(in_dtype)
|
||||
|
||||
|
||||
def gelu7(x: torch.Tensor, alpha: float = 1.702, limit: float = 7.0) -> torch.Tensor:
|
||||
in_dtype = x.dtype
|
||||
x = x.to(torch.float32).clamp(max=limit)
|
||||
return (x * torch.sigmoid(alpha * x)).to(in_dtype)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Modality dispatcher
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ModalityDispatcher:
|
||||
"""Permute a flat token stream so same-modality tokens are contiguous.
|
||||
|
||||
The DiT's multi-expert linears apply a different weight chunk per modality.
|
||||
Instead of carrying a branch inside each Linear, we pre-permute tokens so
|
||||
each chunk sees a contiguous slice, then un-permute before computing
|
||||
RoPE/attention across the full sequence.
|
||||
"""
|
||||
|
||||
def __init__(self, modality_mapping: torch.Tensor, num_modalities: int):
|
||||
self.modality_mapping = modality_mapping
|
||||
self.num_modalities = num_modalities
|
||||
self.permute_mapping = torch.argsort(modality_mapping)
|
||||
self.inv_permute_mapping = torch.argsort(self.permute_mapping)
|
||||
permuted = modality_mapping[self.permute_mapping]
|
||||
self.group_size = torch.bincount(permuted, minlength=num_modalities).to(torch.int32)
|
||||
self.group_size_cpu: list[int] = [int(x) for x in self.group_size.cpu().tolist()]
|
||||
|
||||
def dispatch(self, x: torch.Tensor) -> list[torch.Tensor]:
|
||||
return list(torch.split(x, self.group_size_cpu, dim=0))
|
||||
|
||||
def undispatch(self, *chunks: torch.Tensor) -> torch.Tensor:
|
||||
return torch.cat(chunks, dim=0)
|
||||
|
||||
@staticmethod
|
||||
def permute(x: torch.Tensor, permute_mapping: torch.Tensor) -> torch.Tensor:
|
||||
return x[permute_mapping]
|
||||
|
||||
@staticmethod
|
||||
def inv_permute(x: torch.Tensor, inv_permute_mapping: torch.Tensor) -> torch.Tensor:
|
||||
return x[inv_permute_mapping]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Norms, rotary embed
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MultiModalityRMSNorm(nn.Module):
|
||||
"""RMSNorm with optional per-modality scale.
|
||||
|
||||
When num_modality == 1, behaves identically to a standard RMSNorm with
|
||||
weight initialized to zero (effective weight is 1 + weight, hence the
|
||||
learnable +1 offset baked into the forward path). When num_modality > 1,
|
||||
the weight tensor packs per-modality scales along its flat axis and the
|
||||
dispatcher selects the right chunk per modality.
|
||||
"""
|
||||
|
||||
def __init__(self, dim: int, eps: float = 1e-6, num_modality: int = 1):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.num_modality = num_modality
|
||||
# Always stored in fp32; matches the reference initialization.
|
||||
self.weight = nn.Parameter(torch.zeros(dim * num_modality, dtype=torch.float32))
|
||||
|
||||
def _rms(self, x: torch.Tensor) -> torch.Tensor:
|
||||
t = x.float()
|
||||
return t * torch.rsqrt(t.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
modality_dispatcher: Optional[ModalityDispatcher] = None,
|
||||
) -> torch.Tensor:
|
||||
original_dtype = x.dtype
|
||||
t = self._rms(x)
|
||||
if self.num_modality == 1:
|
||||
return (t * (self.weight + 1)).to(original_dtype)
|
||||
assert modality_dispatcher is not None, (
|
||||
"MultiModalityRMSNorm with num_modality>1 requires a dispatcher"
|
||||
)
|
||||
weight_chunks = self.weight.chunk(self.num_modality, dim=0)
|
||||
parts = modality_dispatcher.dispatch(t)
|
||||
for i in range(self.num_modality):
|
||||
parts[i] = parts[i] * (weight_chunks[i] + 1)
|
||||
return modality_dispatcher.undispatch(*parts).to(original_dtype)
|
||||
|
||||
|
||||
def _freq_bands(num_bands: int, temperature: float = 10000.0) -> torch.Tensor:
|
||||
exp = torch.arange(0, num_bands, 1, dtype=torch.int64).float() / num_bands
|
||||
return 1.0 / (temperature ** exp)
|
||||
|
||||
|
||||
class ElementWiseFourierEmbed(nn.Module):
|
||||
"""Element-wise Fourier embedding over 9-column coords (t, h, w, T, H, W,
|
||||
ref_T, ref_H, ref_W). Produces a per-token positional embedding that
|
||||
acts as the RoPE angle input for attention.
|
||||
|
||||
Weight: `bands` of shape `[dim // 8]` (fixed at init via freq_bands).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
temperature: float = 10000.0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
bands = _freq_bands(dim // 8, temperature=temperature).to(dtype)
|
||||
# `register_buffer` so state_dict keeps it, matching upstream naming.
|
||||
self.register_buffer("bands", bands)
|
||||
|
||||
def forward(self, coords: torch.Tensor) -> torch.Tensor:
|
||||
# coords: [L, 9] = (t, h, w, T, H, W, ref_T, ref_H, ref_W)
|
||||
coords_xyz = coords[:, :3]
|
||||
sizes = coords[:, 3:6]
|
||||
refs = coords[:, 6:9]
|
||||
|
||||
scales = (refs - 1) / (sizes - 1)
|
||||
scales[(refs == 1) & (sizes == 1)] = 1
|
||||
# Center H and W (leave time uncentered).
|
||||
centers = (sizes - 1) / 2
|
||||
centers[:, 0] = 0
|
||||
coords_xyz = coords_xyz - centers
|
||||
|
||||
proj = coords_xyz.unsqueeze(-1) * scales.unsqueeze(-1) * self.bands # [L, 3, B]
|
||||
sin_proj = proj.sin()
|
||||
cos_proj = proj.cos()
|
||||
return torch.cat((sin_proj, cos_proj), dim=1).flatten(1)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Packed-expert linear
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PackedExpertLinear(nn.Module):
|
||||
"""Linear where the weight is packed per-modality along the output axis.
|
||||
|
||||
Shapes:
|
||||
weight: [out_features * num_experts, in_features]
|
||||
bias: [out_features * num_experts] (optional)
|
||||
|
||||
When `num_experts == 1`, behaves exactly like `nn.Linear`. When
|
||||
`num_experts > 1`, `forward` dispatches the input via the supplied
|
||||
`ModalityDispatcher`, applies the per-modality weight/bias chunk, and
|
||||
gathers the outputs in original order.
|
||||
|
||||
Why not use `ReplicatedLinear`? Because the packed-expert layout is not
|
||||
what ReplicatedLinear (or any other fastvideo.layers.linear) is wired
|
||||
for. Using raw `nn.Parameter` keeps weight loading trivial (names map
|
||||
1:1 to the upstream checkpoint) and avoids quantization-path assumptions
|
||||
that don't match this layout.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_features: int,
|
||||
out_features: int,
|
||||
num_experts: int = 1,
|
||||
bias: bool = False,
|
||||
dtype: torch.dtype = torch.bfloat16,
|
||||
):
|
||||
super().__init__()
|
||||
self.in_features = in_features
|
||||
self.out_features = out_features
|
||||
self.num_experts = num_experts
|
||||
self.use_bias = bias
|
||||
self.weight = nn.Parameter(
|
||||
torch.empty(out_features * num_experts, in_features, dtype=dtype)
|
||||
)
|
||||
if bias:
|
||||
self.bias = nn.Parameter(
|
||||
torch.empty(out_features * num_experts, dtype=dtype)
|
||||
)
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
modality_dispatcher: Optional[ModalityDispatcher] = None,
|
||||
) -> torch.Tensor:
|
||||
if self.num_experts == 1:
|
||||
return F.linear(x, self.weight, self.bias)
|
||||
assert modality_dispatcher is not None, (
|
||||
"PackedExpertLinear with num_experts>1 requires a dispatcher"
|
||||
)
|
||||
parts = modality_dispatcher.dispatch(x)
|
||||
w_chunks = self.weight.chunk(self.num_experts, dim=0)
|
||||
b_chunks = (
|
||||
self.bias.chunk(self.num_experts, dim=0)
|
||||
if self.bias is not None else [None] * self.num_experts
|
||||
)
|
||||
for i in range(self.num_experts):
|
||||
parts[i] = F.linear(parts[i], w_chunks[i], b_chunks[i])
|
||||
return modality_dispatcher.undispatch(*parts)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Attention & MLP
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class AttentionSubConfig:
|
||||
hidden_size: int
|
||||
num_heads_q: int
|
||||
num_heads_kv: int
|
||||
head_dim: int
|
||||
num_modality: int
|
||||
enable_attn_gating: bool
|
||||
use_local_attn: bool = False
|
||||
frame_receptive_field: int = 11
|
||||
|
||||
|
||||
class MagiAttention(nn.Module):
|
||||
"""Self-attention with GQA + optional per-head sigmoid gating."""
|
||||
|
||||
def __init__(self, cfg: AttentionSubConfig):
|
||||
super().__init__()
|
||||
self.cfg = cfg
|
||||
self.gating_size = cfg.num_heads_q if cfg.enable_attn_gating else 0
|
||||
qkv_out = (
|
||||
cfg.num_heads_q * cfg.head_dim
|
||||
+ 2 * cfg.num_heads_kv * cfg.head_dim
|
||||
+ self.gating_size
|
||||
)
|
||||
self.pre_norm = MultiModalityRMSNorm(cfg.hidden_size, num_modality=cfg.num_modality)
|
||||
self.linear_qkv = PackedExpertLinear(
|
||||
cfg.hidden_size, qkv_out, num_experts=cfg.num_modality, bias=False,
|
||||
)
|
||||
self.linear_proj = PackedExpertLinear(
|
||||
cfg.num_heads_q * cfg.head_dim, cfg.hidden_size,
|
||||
num_experts=cfg.num_modality, bias=False,
|
||||
)
|
||||
self.q_norm = MultiModalityRMSNorm(cfg.head_dim, num_modality=cfg.num_modality)
|
||||
self.k_norm = MultiModalityRMSNorm(cfg.head_dim, num_modality=cfg.num_modality)
|
||||
|
||||
self.q_size = cfg.num_heads_q * cfg.head_dim
|
||||
self.kv_size = cfg.num_heads_kv * cfg.head_dim
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=cfg.num_heads_q,
|
||||
head_size=cfg.head_dim,
|
||||
num_kv_heads=cfg.num_heads_kv,
|
||||
causal=False,
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
),
|
||||
)
|
||||
|
||||
def configure_local_attention(
|
||||
self,
|
||||
*,
|
||||
enabled: bool,
|
||||
frame_receptive_field: int = 11,
|
||||
) -> None:
|
||||
self.cfg.use_local_attn = enabled
|
||||
self.cfg.frame_receptive_field = frame_receptive_field
|
||||
|
||||
def _sdpa(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
|
||||
"""Run SDPA on [L, H, D] tensors and return [L, Hq, D]."""
|
||||
if q.numel() == 0:
|
||||
return q.new_empty(q.shape)
|
||||
out = F.scaled_dot_product_attention(
|
||||
q.transpose(0, 1).unsqueeze(0).contiguous(),
|
||||
k.transpose(0, 1).unsqueeze(0).contiguous(),
|
||||
v.transpose(0, 1).unsqueeze(0).contiguous(),
|
||||
enable_gqa=self.cfg.num_heads_q != self.cfg.num_heads_kv,
|
||||
)
|
||||
return out.squeeze(0).transpose(0, 1).contiguous()
|
||||
|
||||
def _local_window_attention(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
*,
|
||||
num_video_tokens: int,
|
||||
num_frames: int,
|
||||
) -> torch.Tensor:
|
||||
"""Approximate upstream FFAHandler block accumulation with SDPA.
|
||||
|
||||
SR-1080p's reference kernel sums three independently-normalized
|
||||
attention contributions:
|
||||
|
||||
* video frame queries -> local-window video keys;
|
||||
* all video queries -> all audio+text keys;
|
||||
* all audio+text queries -> full sequence keys.
|
||||
|
||||
This method mirrors that accumulator semantics with ordinary SDPA
|
||||
slices. It is intentionally scoped to single-process inference; layers
|
||||
without ``use_local_attn`` keep the existing full LocalAttention path.
|
||||
"""
|
||||
if num_frames <= 0 or num_video_tokens <= 0:
|
||||
return self._sdpa(q, k, v)
|
||||
if num_video_tokens % num_frames != 0:
|
||||
raise ValueError(
|
||||
f"MagiHuman local attention expects video tokens divisible by "
|
||||
f"frames, got {num_video_tokens=} and {num_frames=}."
|
||||
)
|
||||
|
||||
token_per_frame = num_video_tokens // num_frames
|
||||
out = torch.zeros(
|
||||
q.shape[0],
|
||||
self.cfg.num_heads_q,
|
||||
self.cfg.head_dim,
|
||||
device=q.device,
|
||||
dtype=q.dtype,
|
||||
)
|
||||
rf = int(self.cfg.frame_receptive_field)
|
||||
|
||||
q_video = q[:num_video_tokens]
|
||||
k_video = k[:num_video_tokens]
|
||||
v_video = v[:num_video_tokens]
|
||||
for frame_idx in range(num_frames):
|
||||
q_start = frame_idx * token_per_frame
|
||||
q_end = q_start + token_per_frame
|
||||
k_start = max(0, (frame_idx - rf) * token_per_frame)
|
||||
k_end = min(num_video_tokens, (frame_idx + rf + 1) * token_per_frame)
|
||||
out[q_start:q_end] = self._sdpa(
|
||||
q_video[q_start:q_end],
|
||||
k_video[k_start:k_end],
|
||||
v_video[k_start:k_end],
|
||||
)
|
||||
|
||||
if num_video_tokens < q.shape[0]:
|
||||
k_at = k[num_video_tokens:]
|
||||
v_at = v[num_video_tokens:]
|
||||
out[:num_video_tokens] = out[:num_video_tokens] + self._sdpa(
|
||||
q[:num_video_tokens],
|
||||
k_at,
|
||||
v_at,
|
||||
)
|
||||
out[num_video_tokens:] = self._sdpa(
|
||||
q[num_video_tokens:],
|
||||
k,
|
||||
v,
|
||||
)
|
||||
return out
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rope: torch.Tensor,
|
||||
permute_mapping: torch.Tensor,
|
||||
inv_permute_mapping: torch.Tensor,
|
||||
modality_dispatcher: ModalityDispatcher,
|
||||
num_video_tokens: int | None = None,
|
||||
num_frames: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
orig_dtype = self.linear_qkv.weight.dtype
|
||||
h = self.pre_norm(hidden_states, modality_dispatcher=modality_dispatcher).to(orig_dtype)
|
||||
qkv = self.linear_qkv(h, modality_dispatcher=modality_dispatcher).float()
|
||||
q, k, v, g = torch.split(
|
||||
qkv, [self.q_size, self.kv_size, self.kv_size, self.gating_size], dim=-1,
|
||||
)
|
||||
q = q.view(-1, self.cfg.num_heads_q, self.cfg.head_dim)
|
||||
k = k.view(-1, self.cfg.num_heads_kv, self.cfg.head_dim)
|
||||
v = v.view(-1, self.cfg.num_heads_kv, self.cfg.head_dim)
|
||||
g = g.view(-1, self.cfg.num_heads_q, 1) if self.gating_size else None
|
||||
|
||||
q = self.q_norm(q, modality_dispatcher=modality_dispatcher)
|
||||
k = self.k_norm(k, modality_dispatcher=modality_dispatcher)
|
||||
|
||||
# Un-permute before RoPE + attention so positional order reflects
|
||||
# the original (video, audio, text) concat — matches reference.
|
||||
q = ModalityDispatcher.inv_permute(q, inv_permute_mapping)
|
||||
k = ModalityDispatcher.inv_permute(k, inv_permute_mapping)
|
||||
v = ModalityDispatcher.inv_permute(v, inv_permute_mapping)
|
||||
if g is not None:
|
||||
g = ModalityDispatcher.inv_permute(g, inv_permute_mapping)
|
||||
|
||||
# Element-wise Fourier embed packs sin/cos of 3 axes into a single
|
||||
# `rope` tensor. Match reference's split:
|
||||
# sin_emb, cos_emb = rope.tensor_split(2, -1)
|
||||
# Reference passes (cos_emb, sin_emb) but splits sin first — replicated
|
||||
# exactly so weight parity holds. Partial RoPE: rope dim is
|
||||
# 6 * (head_dim // 8) = 96 < head_dim (128), so the trailing 32
|
||||
# head_dim positions stay unrotated, matching the reference.
|
||||
sin_emb, cos_emb = rope.tensor_split(2, -1)
|
||||
rot_dim = cos_emb.shape[-1] * 2
|
||||
q_rot = _apply_rotary_emb(q[..., :rot_dim], cos_emb, sin_emb, is_neox_style=True)
|
||||
k_rot = _apply_rotary_emb(k[..., :rot_dim], cos_emb, sin_emb, is_neox_style=True)
|
||||
if rot_dim < q.shape[-1]:
|
||||
q = torch.cat([q_rot, q[..., rot_dim:]], dim=-1)
|
||||
k = torch.cat([k_rot, k[..., rot_dim:]], dim=-1)
|
||||
else:
|
||||
q, k = q_rot, k_rot
|
||||
|
||||
# Run SDPA via FastVideo's LocalAttention so the backend selection
|
||||
# (SDPA / FlashAttn / SLA / SageAttn) flows through the standard
|
||||
# configurable path. GQA is handled inside the SDPA backend via
|
||||
# `enable_gqa=True` when num_heads_q != num_heads_kv, so we no
|
||||
# longer need the manual `repeat_interleave` here.
|
||||
# Attention math runs at orig_dtype (bf16 in production and in the
|
||||
# parity test, since PackedExpertLinear's default is bf16, matching
|
||||
# upstream BaseLinear at dit_module.py:330). The gating multiply
|
||||
# promotes back to fp32 implicitly via PyTorch's type-promotion
|
||||
# rules: bf16_attn_out * sigmoid(fp32_g) -> fp32, mirroring upstream
|
||||
# dit_module.py:649.
|
||||
q = q.to(orig_dtype)
|
||||
k = k.to(orig_dtype)
|
||||
v = v.to(orig_dtype)
|
||||
if self.cfg.use_local_attn:
|
||||
if num_video_tokens is None or num_frames is None:
|
||||
raise ValueError("MagiHuman local attention requires video token/frame metadata.")
|
||||
out = self._local_window_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
num_video_tokens=num_video_tokens,
|
||||
num_frames=num_frames,
|
||||
)
|
||||
else:
|
||||
out = self.attn(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0)).squeeze(0)
|
||||
|
||||
out = ModalityDispatcher.permute(out, permute_mapping)
|
||||
if g is not None:
|
||||
g = ModalityDispatcher.permute(g, permute_mapping)
|
||||
out = out * torch.sigmoid(g)
|
||||
out = out.reshape(-1, self.cfg.num_heads_q * self.cfg.head_dim).to(orig_dtype)
|
||||
return self.linear_proj(out, modality_dispatcher=modality_dispatcher)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MLPSubConfig:
|
||||
hidden_size: int
|
||||
intermediate_size: int
|
||||
activation: str # "swiglu7" or "gelu7"
|
||||
num_modality: int
|
||||
gated: bool
|
||||
|
||||
|
||||
class MagiMLP(nn.Module):
|
||||
def __init__(self, cfg: MLPSubConfig):
|
||||
super().__init__()
|
||||
self.cfg = cfg
|
||||
self.pre_norm = MultiModalityRMSNorm(cfg.hidden_size, num_modality=cfg.num_modality)
|
||||
up_out = cfg.intermediate_size * 2 if cfg.gated else cfg.intermediate_size
|
||||
self.up_gate_proj = PackedExpertLinear(
|
||||
cfg.hidden_size, up_out, num_experts=cfg.num_modality, bias=False,
|
||||
)
|
||||
self.down_proj = PackedExpertLinear(
|
||||
cfg.intermediate_size, cfg.hidden_size,
|
||||
num_experts=cfg.num_modality, bias=False,
|
||||
)
|
||||
self._act = swiglu7 if cfg.activation == "swiglu7" else gelu7
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
modality_dispatcher: ModalityDispatcher,
|
||||
) -> torch.Tensor:
|
||||
orig_dtype = self.up_gate_proj.weight.dtype
|
||||
x = self.pre_norm(x, modality_dispatcher=modality_dispatcher).to(orig_dtype)
|
||||
x = self.up_gate_proj(x, modality_dispatcher=modality_dispatcher).float()
|
||||
x = self._act(x).to(orig_dtype)
|
||||
x = self.down_proj(x, modality_dispatcher=modality_dispatcher).float()
|
||||
return x
|
||||
|
||||
|
||||
class MagiTransformerLayer(nn.Module):
|
||||
def __init__(self, arch: MagiHumanArchConfig, layer_idx: int):
|
||||
super().__init__()
|
||||
num_modality = 3 if layer_idx in arch.mm_layers else 1
|
||||
self.post_norm = layer_idx in arch.post_norm_layers
|
||||
self.layer_idx = layer_idx
|
||||
|
||||
self.attention = MagiAttention(AttentionSubConfig(
|
||||
hidden_size=arch.hidden_size,
|
||||
num_heads_q=arch.num_attention_heads,
|
||||
num_heads_kv=arch.num_heads_kv,
|
||||
head_dim=arch.head_dim,
|
||||
num_modality=num_modality,
|
||||
enable_attn_gating=arch.enable_attn_gating,
|
||||
use_local_attn=layer_idx in arch.local_attn_layers,
|
||||
))
|
||||
|
||||
is_gelu7 = layer_idx in arch.gelu7_layers
|
||||
if is_gelu7:
|
||||
intermediate = arch.hidden_size * 4
|
||||
gated = False
|
||||
activation = "gelu7"
|
||||
else:
|
||||
intermediate = (arch.hidden_size * 4 * 2 // 3) // 4 * 4
|
||||
gated = True
|
||||
activation = "swiglu7"
|
||||
|
||||
self.mlp = MagiMLP(MLPSubConfig(
|
||||
hidden_size=arch.hidden_size,
|
||||
intermediate_size=intermediate,
|
||||
activation=activation,
|
||||
num_modality=num_modality,
|
||||
gated=gated,
|
||||
))
|
||||
|
||||
if self.post_norm:
|
||||
self.attn_post_norm = MultiModalityRMSNorm(arch.hidden_size, num_modality=num_modality)
|
||||
self.mlp_post_norm = MultiModalityRMSNorm(arch.hidden_size, num_modality=num_modality)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rope: torch.Tensor,
|
||||
permute_mapping: torch.Tensor,
|
||||
inv_permute_mapping: torch.Tensor,
|
||||
modality_dispatcher: ModalityDispatcher,
|
||||
num_video_tokens: int | None = None,
|
||||
num_frames: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
attn_out = self.attention(
|
||||
hidden_states, rope, permute_mapping, inv_permute_mapping, modality_dispatcher,
|
||||
num_video_tokens=num_video_tokens,
|
||||
num_frames=num_frames,
|
||||
)
|
||||
if self.post_norm:
|
||||
attn_out = self.attn_post_norm(attn_out, modality_dispatcher=modality_dispatcher)
|
||||
hidden_states = hidden_states + attn_out
|
||||
|
||||
mlp_out = self.mlp(hidden_states, modality_dispatcher=modality_dispatcher)
|
||||
if self.post_norm:
|
||||
mlp_out = self.mlp_post_norm(mlp_out, modality_dispatcher=modality_dispatcher)
|
||||
return hidden_states + mlp_out
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Adapter (per-modality embedders + Fourier RoPE producer)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class MagiAdapter(nn.Module):
|
||||
def __init__(self, arch: MagiHumanArchConfig):
|
||||
super().__init__()
|
||||
# Embedders stay in fp32 to match the reference dtype exactly.
|
||||
self.video_embedder = nn.Linear(
|
||||
arch.video_in_channels, arch.hidden_size, bias=True, dtype=torch.float32,
|
||||
)
|
||||
self.text_embedder = nn.Linear(
|
||||
arch.text_in_channels, arch.hidden_size, bias=True, dtype=torch.float32,
|
||||
)
|
||||
self.audio_embedder = nn.Linear(
|
||||
arch.audio_in_channels, arch.hidden_size, bias=True, dtype=torch.float32,
|
||||
)
|
||||
self.rope = ElementWiseFourierEmbed(arch.head_dim)
|
||||
# RoPE cache: coords_mapping is the same tensor object across timesteps
|
||||
# in the denoising loop, so data_ptr()+shape+dtype+device is a fast,
|
||||
# collision-free key that avoids recomputing the Fourier embed each step.
|
||||
self._cached_rope: Optional[torch.Tensor] = None
|
||||
self._cached_rope_key: Optional[tuple] = None
|
||||
|
||||
def _rope_cache_key(self, t: torch.Tensor) -> tuple:
|
||||
return (t.data_ptr(), t.shape, t.dtype, t.device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
coords_mapping: torch.Tensor,
|
||||
video_mask: torch.Tensor,
|
||||
audio_mask: torch.Tensor,
|
||||
text_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
key = self._rope_cache_key(coords_mapping)
|
||||
if key != self._cached_rope_key:
|
||||
self._cached_rope = self.rope(coords_mapping)
|
||||
self._cached_rope_key = key
|
||||
rope = self._cached_rope
|
||||
# Embedder dtypes may differ from x's dtype when FastVideo's FSDP
|
||||
# loader casts all weights to `pipeline_config.precision` (bf16).
|
||||
# Match the weight dtype per modality.
|
||||
v_w = self.video_embedder.weight
|
||||
a_w = self.audio_embedder.weight
|
||||
t_w = self.text_embedder.weight
|
||||
out = torch.zeros(
|
||||
x.shape[0], self.video_embedder.out_features,
|
||||
device=x.device, dtype=v_w.dtype,
|
||||
)
|
||||
out[text_mask] = self.text_embedder(
|
||||
x[text_mask, : self.text_embedder.in_features].to(t_w.dtype)
|
||||
).to(out.dtype)
|
||||
out[audio_mask] = self.audio_embedder(
|
||||
x[audio_mask, : self.audio_embedder.in_features].to(a_w.dtype)
|
||||
).to(out.dtype)
|
||||
out[video_mask] = self.video_embedder(
|
||||
x[video_mask, : self.video_embedder.in_features].to(v_w.dtype)
|
||||
).to(out.dtype)
|
||||
return out, rope
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Top-level DiT
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _TransformerBlock(nn.Module):
|
||||
"""Thin ModuleList wrapper to keep the 'block.layers.<i>' state_dict
|
||||
naming identical to the upstream checkpoint (which uses a magi_compile
|
||||
decorator producing `block.layers.<i>.*`)."""
|
||||
|
||||
def __init__(self, arch: MagiHumanArchConfig):
|
||||
super().__init__()
|
||||
self.layers = nn.ModuleList([
|
||||
MagiTransformerLayer(arch, i) for i in range(arch.num_layers)
|
||||
])
|
||||
|
||||
def configure_local_attention(
|
||||
self,
|
||||
local_attn_layers: tuple[int, ...],
|
||||
frame_receptive_field: int = 11,
|
||||
) -> None:
|
||||
enabled_layers = set(local_attn_layers)
|
||||
for idx, layer in enumerate(self.layers):
|
||||
layer.attention.configure_local_attention(
|
||||
enabled=idx in enabled_layers,
|
||||
frame_receptive_field=frame_receptive_field,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
rope: torch.Tensor,
|
||||
permute_mapping: torch.Tensor,
|
||||
inv_permute_mapping: torch.Tensor,
|
||||
modality_dispatcher: ModalityDispatcher,
|
||||
num_video_tokens: int | None = None,
|
||||
num_frames: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
for layer in self.layers:
|
||||
x = layer(
|
||||
x,
|
||||
rope,
|
||||
permute_mapping,
|
||||
inv_permute_mapping,
|
||||
modality_dispatcher,
|
||||
num_video_tokens=num_video_tokens,
|
||||
num_frames=num_frames,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
_CFG = MagiHumanVideoConfig()
|
||||
|
||||
|
||||
class MagiHumanDiT(BaseDiT):
|
||||
"""Top-level DiT for daVinci-MagiHuman (base).
|
||||
|
||||
Forward signature mirrors the reference `DiTModel.forward`: it takes a
|
||||
flat token stream, its per-token coords and modality mapping, and
|
||||
returns per-modality outputs packed into a max-channel-width tensor.
|
||||
|
||||
This scaffold is single-GPU only; the `ulysses_scheduler().dispatch(...)`
|
||||
sequence-parallel wrapping in the reference has no equivalent here yet.
|
||||
"""
|
||||
|
||||
# BaseDiT requires these class attrs. Source them from the config so
|
||||
# they stay in sync with MagiHumanVideoConfig edits.
|
||||
_fsdp_shard_conditions = _CFG._fsdp_shard_conditions
|
||||
_compile_conditions = _CFG._compile_conditions
|
||||
_supported_attention_backends = _CFG._supported_attention_backends
|
||||
param_names_mapping = _CFG.param_names_mapping
|
||||
reverse_param_names_mapping = _CFG.reverse_param_names_mapping
|
||||
lora_param_names_mapping = _CFG.lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: MagiHumanVideoConfig, hf_config: dict | None = None, **kwargs):
|
||||
super().__init__(config=config, hf_config=hf_config or {})
|
||||
arch: MagiHumanArchConfig = getattr(config, "arch_config", config)
|
||||
self.arch = arch
|
||||
|
||||
# BaseDiT contract instance vars.
|
||||
self.hidden_size = arch.hidden_size
|
||||
self.num_attention_heads = arch.num_attention_heads
|
||||
self.num_channels_latents = arch.num_channels_latents
|
||||
|
||||
self.adapter = MagiAdapter(arch)
|
||||
self.block = _TransformerBlock(arch)
|
||||
self.final_norm_video = MultiModalityRMSNorm(arch.hidden_size)
|
||||
self.final_norm_audio = MultiModalityRMSNorm(arch.hidden_size)
|
||||
self.final_linear_video = nn.Linear(
|
||||
arch.hidden_size, arch.video_in_channels, bias=False, dtype=torch.float32,
|
||||
)
|
||||
self.final_linear_audio = nn.Linear(
|
||||
arch.hidden_size, arch.audio_in_channels, bias=False, dtype=torch.float32,
|
||||
)
|
||||
# Dispatcher + mask cache: modality_mapping is the same tensor object
|
||||
# across all timesteps in the denoising loop; data_ptr()+shape+dtype+device
|
||||
# is a fast, collision-free key that avoids rebuilding ModalityDispatcher
|
||||
# (which calls argsort + bincount) on every forward call.
|
||||
self._cached_dispatcher: Optional[ModalityDispatcher] = None
|
||||
self._cached_video_mask: Optional[torch.Tensor] = None
|
||||
self._cached_audio_mask: Optional[torch.Tensor] = None
|
||||
self._cached_text_mask: Optional[torch.Tensor] = None
|
||||
self._cached_modality_key: Optional[tuple] = None
|
||||
|
||||
def configure_local_attention(
|
||||
self,
|
||||
local_attn_layers: tuple[int, ...] | list[int],
|
||||
frame_receptive_field: int = 11,
|
||||
) -> None:
|
||||
layers = tuple(int(layer) for layer in local_attn_layers)
|
||||
self.arch.local_attn_layers = layers
|
||||
self.block.configure_local_attention(layers, frame_receptive_field)
|
||||
|
||||
def _modality_cache_key(self, t: torch.Tensor) -> tuple:
|
||||
return (t.data_ptr(), t.shape, t.dtype, t.device)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
coords_mapping: torch.Tensor,
|
||||
modality_mapping: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
x: [L, max(V_ch, A_ch, T_ch)]
|
||||
coords_mapping: [L, 9]
|
||||
modality_mapping: [L] (int in {VIDEO, AUDIO, TEXT})
|
||||
Returns:
|
||||
out: [L, max(V_ch, A_ch)] with video channels in video slots and
|
||||
audio channels in audio slots; text slots are zero.
|
||||
"""
|
||||
key = self._modality_cache_key(modality_mapping)
|
||||
if key != self._cached_modality_key:
|
||||
self._cached_dispatcher = ModalityDispatcher(modality_mapping, num_modalities=3)
|
||||
self._cached_video_mask = modality_mapping == Modality.VIDEO
|
||||
self._cached_audio_mask = modality_mapping == Modality.AUDIO
|
||||
self._cached_text_mask = modality_mapping == Modality.TEXT
|
||||
self._cached_modality_key = key
|
||||
dispatcher = self._cached_dispatcher
|
||||
video_mask = self._cached_video_mask
|
||||
audio_mask = self._cached_audio_mask
|
||||
text_mask = self._cached_text_mask
|
||||
num_video_tokens = int(video_mask.sum().item())
|
||||
if num_video_tokens:
|
||||
num_frames = int(coords_mapping[:num_video_tokens, 0].max().item()) + 1
|
||||
else:
|
||||
num_frames = 0
|
||||
|
||||
x, rope = self.adapter(x, coords_mapping, video_mask, audio_mask, text_mask)
|
||||
# Keep the residual stream in adapter dtype (fp32) entering the block.
|
||||
# Upstream daVinci-MagiHuman dit_module.py:923 casts to params_dtype,
|
||||
# which is fp32 by default; each layer's pre_norm.to(bf16) handles
|
||||
# the bf16 internal-compute boundary, and linear_proj outputs bf16
|
||||
# which gets promoted back to fp32 by the residual addition. Casting
|
||||
# the residual to bf16 here degrades the cross-layer accumulator and
|
||||
# compounds visibly over 40 layers in pipeline parity.
|
||||
x = ModalityDispatcher.permute(x, dispatcher.permute_mapping)
|
||||
|
||||
x = self.block(
|
||||
x, rope,
|
||||
permute_mapping=dispatcher.permute_mapping,
|
||||
inv_permute_mapping=dispatcher.inv_permute_mapping,
|
||||
modality_dispatcher=dispatcher,
|
||||
num_video_tokens=num_video_tokens,
|
||||
num_frames=num_frames,
|
||||
)
|
||||
x = ModalityDispatcher.inv_permute(x, dispatcher.inv_permute_mapping)
|
||||
|
||||
x_video = x[video_mask].to(self.final_norm_video.weight.dtype)
|
||||
x_video = self.final_norm_video(x_video)
|
||||
x_video = self.final_linear_video(x_video)
|
||||
|
||||
x_audio = x[audio_mask].to(self.final_norm_audio.weight.dtype)
|
||||
x_audio = self.final_norm_audio(x_audio)
|
||||
x_audio = self.final_linear_audio(x_audio)
|
||||
|
||||
max_ch = max(self.arch.video_in_channels, self.arch.audio_in_channels)
|
||||
out = torch.zeros(x.shape[0], max_ch, device=x.device, dtype=x.dtype)
|
||||
out[video_mask, : self.arch.video_in_channels] = x_video.to(out.dtype)
|
||||
out[audio_mask, : self.arch.audio_in_channels] = x_audio.to(out.dtype)
|
||||
return out
|
||||
|
||||
|
||||
EntryClass = MagiHumanDiT
|
||||
@@ -0,0 +1,127 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""T5-Gemma encoder wrapper for daVinci-MagiHuman.
|
||||
|
||||
MagiHuman uses `transformers.models.t5gemma.T5GemmaEncoderModel` on
|
||||
`google/t5gemma-9b-9b-ul2` (a gated Google repo). This wrapper follows the
|
||||
same lazy-loading pattern as `fastvideo/models/encoders/gemma.py`: we keep
|
||||
the HF module under `self._t5gemma_model` and exclude it from
|
||||
`named_parameters` so FastVideo's weight loader does not try to load
|
||||
encoder shards from the converted repo directory.
|
||||
|
||||
For the base MagiHuman T2V port there are no additional connector layers
|
||||
on top — the pipeline prompt-preprocessing stage handles pad-or-trim to
|
||||
`text_len` and exposes both the padded embedding and the original length.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class T5GemmaEncoderModel(TextEncoder):
|
||||
"""Thin wrapper over HuggingFace's `T5GemmaEncoderModel`.
|
||||
|
||||
On first `forward`, the wrapper lazily instantiates the upstream encoder
|
||||
from `t5gemma_model_path` (defaulting to `google/t5gemma-9b-9b-ul2`).
|
||||
Afterwards, forward returns a `BaseEncoderOutput` with
|
||||
`last_hidden_state = [B, L, 3584]` matching MagiHuman's
|
||||
`context.half()` output.
|
||||
"""
|
||||
|
||||
_supported_attention_backends = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__(config)
|
||||
arch = config.arch_config
|
||||
self.t5gemma_model_path: str = arch.t5gemma_model_path
|
||||
self.t5gemma_dtype: str = arch.t5gemma_dtype
|
||||
self._t5gemma_model = None
|
||||
|
||||
def named_parameters(self, prefix: str = "", recurse: bool = True):
|
||||
# The upstream encoder is loaded lazily and its parameters are
|
||||
# managed by HF, not FastVideo's loader. Hide them from the parent
|
||||
# module-tree traversal so Diffusers-repo weight loading does not
|
||||
# try to match them.
|
||||
for name, param in super().named_parameters(prefix=prefix, recurse=recurse):
|
||||
if name.startswith("_t5gemma_model.") or name == "_t5gemma_model":
|
||||
continue
|
||||
yield name, param
|
||||
|
||||
def _build_t5gemma_model(self, device: torch.device | None = None):
|
||||
from transformers.models.t5gemma import T5GemmaEncoderModel as HFEncoder
|
||||
|
||||
path = self.t5gemma_model_path
|
||||
if not path:
|
||||
raise ValueError(
|
||||
"t5gemma_model_path must be set. Expected "
|
||||
"`google/t5gemma-9b-9b-ul2` or a local path to an "
|
||||
"equivalent T5-Gemma encoder."
|
||||
)
|
||||
dtype = getattr(torch, self.t5gemma_dtype, torch.bfloat16)
|
||||
model = HFEncoder.from_pretrained(
|
||||
path,
|
||||
is_encoder_decoder=False,
|
||||
dtype=dtype,
|
||||
)
|
||||
if os.getenv("FASTVIDEO_ATTENTION_BACKEND") == "TORCH_SDPA":
|
||||
if hasattr(model.config, "attn_implementation"):
|
||||
model.config.attn_implementation = "sdpa"
|
||||
if hasattr(model.config, "_attn_implementation"):
|
||||
model.config._attn_implementation = "sdpa"
|
||||
if device is not None:
|
||||
model = model.to(device=device)
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
@property
|
||||
def t5gemma_model(self):
|
||||
if self._t5gemma_model is None:
|
||||
# Lazy-load on CPU if no device is known yet; `forward` will
|
||||
# move the model to the input's device on first call.
|
||||
self._t5gemma_model = self._build_t5gemma_model()
|
||||
return self._t5gemma_model
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None = 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:
|
||||
# Ensure the lazy-loaded encoder lives on the same device as the
|
||||
# input; lazy-loading leaves it on CPU until the first forward.
|
||||
ref = input_ids if input_ids is not None else inputs_embeds
|
||||
target_device = ref.device if ref is not None else None
|
||||
model = self.t5gemma_model
|
||||
if target_device is not None:
|
||||
first_param = next(model.parameters(), None)
|
||||
if first_param is not None and first_param.device != target_device:
|
||||
model = model.to(device=target_device)
|
||||
self._t5gemma_model = model
|
||||
outputs = model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
output_hidden_states=bool(output_hidden_states),
|
||||
)
|
||||
# MagiHuman casts to fp16 at this point; keep the raw dtype here and
|
||||
# leave precision management to the pipeline's postprocess stage.
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=outputs["last_hidden_state"],
|
||||
hidden_states=getattr(outputs, "hidden_states", None),
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
|
||||
EntryClass = T5GemmaEncoderModel
|
||||
@@ -78,6 +78,7 @@ class ComponentLoader(ABC):
|
||||
module_loaders = {
|
||||
"scheduler": (SchedulerLoader, "diffusers"),
|
||||
"transformer": (TransformerLoader, "diffusers"),
|
||||
"sr_transformer": (TransformerLoader, "diffusers"),
|
||||
"transformer_2": (TransformerLoader, "diffusers"),
|
||||
"transformer_3": (TransformerLoader, "diffusers"),
|
||||
"vae": (VAELoader, "diffusers"),
|
||||
|
||||
@@ -183,7 +183,7 @@ def _load_video_with_ffmpeg(
|
||||
except AttributeError as e:
|
||||
raise AttributeError(
|
||||
"Unable to find an ffmpeg installation on your machine. "
|
||||
"Please install via `pip install imageio-ffmpeg`") from e
|
||||
"Please install via `uv pip install imageio-ffmpeg`") from e
|
||||
|
||||
pil_images = []
|
||||
original_fps = None
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# `fastvideo/pipelines/` — Pipeline Composition
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Diffusion pipelines are **compositions of `PipelineStage` objects**. Each stage owns one verb (validate / encode / schedule / denoise / decode). Adding a model means assembling stages, not subclassing a megapipeline.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
pipelines/
|
||||
├── pipeline_batch_info.py # ForwardBatch — the dict passed between stages
|
||||
├── lora_pipeline.py # LoRA-aware base
|
||||
├── composed_pipeline_base.py # Base for stage-composed pipelines
|
||||
├── stages/ # Reusable stage implementations (~30 files)
|
||||
│ ├── base.py # PipelineStage ABC + StageVerificationError
|
||||
│ ├── input_validation.py # Validates ForwardBatch shape/keys
|
||||
│ ├── text_encoding.py # Generic prompt encoder stage
|
||||
│ ├── image_encoding.py # Image conditioning
|
||||
│ ├── latent_preparation.py # Init noise + scheduler
|
||||
│ ├── conditioning.py # CFG / negative prompt fan-out
|
||||
│ ├── denoising.py # Standard diffusion loop
|
||||
│ ├── sd35_conditioning.py # Per-model overrides (named by family)
|
||||
│ ├── longcat_*.py # LongCat I2V/V2V/refine variants
|
||||
│ ├── gen3c_stages.py # Gen3C-specific stages
|
||||
│ ├── gamecraft_denoising.py # GameCraft-specific
|
||||
│ └── matrixgame_denoising.py # MatrixGame-specific
|
||||
├── basic/ # Per-model end-to-end pipelines
|
||||
│ ├── hunyuan/, hunyuan15/, hyworld/, gamecraft/, gen3c/, cosmos/
|
||||
│ ├── wan/, longcat/, ltx2/, lingbotworld/, magi_human/, matrixgame/
|
||||
│ ├── sd35/, stable_audio/, turbodiffusion/
|
||||
│ └── <model>/{<model>_pipeline.py, presets.py, __init__.py}
|
||||
├── preprocess/ # Data preprocessing pipelines (ltx2, wan, matrixgame)
|
||||
└── training/ # Training-time pipeline glue
|
||||
```
|
||||
|
||||
## Stage Authoring Rules
|
||||
|
||||
- Subclass `PipelineStage` from `stages/base.py`. Implement `forward(batch, args) -> ForwardBatch`.
|
||||
- Implement `verify_input` / `verify_output` — both return `VerificationResult`. Failures raise `StageVerificationError`.
|
||||
- Mutate `ForwardBatch` only by reassigning fields you declared in `pipeline_batch_info.py`. New keys → add to the dataclass first.
|
||||
- Stages must be **deterministic given the same `ForwardBatch + FastVideoArgs`**. Side effects (logging, profiling) only.
|
||||
- Read all knobs from the passed-in `FastVideoArgs` / `PipelineConfig`. Never `os.getenv` directly.
|
||||
|
||||
## Per-Model Pipeline Pattern (`basic/<model>/`)
|
||||
|
||||
Every model directory has the same skeleton:
|
||||
|
||||
```
|
||||
basic/<model>/
|
||||
├── __init__.py
|
||||
├── <model>_pipeline.py # Composes stages list
|
||||
├── presets.py # Default PipelineConfig + SamplingParam combos
|
||||
└── (optional) stage_overrides.py, continuation.py, ...
|
||||
```
|
||||
|
||||
`presets.py` is the entry point that `registry.py` imports — it must export the named preset constants used elsewhere in the codebase.
|
||||
|
||||
## Forking vs Reusing a Stage
|
||||
|
||||
Reuse `stages/text_encoding.py` if your model takes text → embeddings via a standard encoder. Fork only when:
|
||||
|
||||
- The model needs a **different ForwardBatch shape** (extra inputs, different output keys).
|
||||
- The denoising loop has structural differences (causal, refine-then-denoise, multi-stream).
|
||||
|
||||
When forking, keep the file name model-prefixed (`longcat_*`, `gamecraft_*`) so the registry stays grep-able.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Putting a full pipeline in a single file under `basic/<model>/` instead of composing stages.
|
||||
- Reading config from globals or env vars inside a stage.
|
||||
- Adding cross-stage state via module-level dicts. Use `ForwardBatch`.
|
||||
@@ -37,7 +37,7 @@ def load_moge_model(
|
||||
from moge.model.v1 import MoGeModel
|
||||
except ImportError as exc:
|
||||
raise ImportError("MoGe is required for GEN3C 3D cache conditioning. "
|
||||
"Install it with: pip install git+https://github.com/microsoft/MoGe.git. "
|
||||
"Install it with: uv pip install git+https://github.com/microsoft/MoGe.git. "
|
||||
"If import fails with libGL.so.1, install system deps: "
|
||||
"sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1") from exc
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,417 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MagiHuman base text-to-AV pipeline.
|
||||
|
||||
Top-level composition for the daVinci-MagiHuman base model. Wires:
|
||||
|
||||
InputValidationStage -> TextEncodingStage (T5-Gemma)
|
||||
-> MagiHumanLatentPreparationStage
|
||||
-> MagiHumanDenoisingStage
|
||||
-> DecodingStage (Wan 2.2 TI2V-5B VAE decode for video)
|
||||
-> MagiHumanAudioDecodingStage (Stable Audio Open 1.0 VAE decode)
|
||||
|
||||
The base checkpoint is a joint audio-visual generator; both the video
|
||||
and audio paths run in the denoising loop and both are decoded.
|
||||
|
||||
`load_modules` is overridden so the four cross-variant shared components
|
||||
(text_encoder, tokenizer, audio_vae, video vae) lazy-load from their
|
||||
canonical upstream HF repos at first build time instead of being
|
||||
bundled inside every converted MagiHuman variant. This keeps each
|
||||
variant's converted repo at ~5-30 GB (transformer + scheduler +
|
||||
model_index.json) instead of ~30-55 GB, and lets all variants share
|
||||
the same ~25 GB of cached upstream weights.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
|
||||
from fastvideo.configs.models.vaes import OobleckVAEConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.encoders.t5gemma import T5GemmaEncoderModel
|
||||
from fastvideo.models.vaes.sa_audio import SAAudioVAEModel
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler, )
|
||||
from fastvideo.pipelines.basic.magi_human.stages import (
|
||||
MagiHumanAudioDecodingStage,
|
||||
MagiHumanDenoisingStage,
|
||||
MagiHumanLatentPreparationStage,
|
||||
MagiHumanReferenceImageStage,
|
||||
MagiHumanSRDenoisingStage,
|
||||
MagiHumanSRLatentPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (
|
||||
DecodingStage,
|
||||
InputValidationStage,
|
||||
TextEncodingStage,
|
||||
)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_T5GEMMA_HF_ID = "google/t5gemma-9b-9b-ul2"
|
||||
_SA_AUDIO_HF_ID = "stabilityai/stable-audio-open-1.0"
|
||||
_WAN_VAE_HF_ID = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
|
||||
|
||||
def _ensure_hf_token_env() -> str | None:
|
||||
"""Surface any of the three common HF token env vars as `HF_TOKEN`.
|
||||
|
||||
FastVideo workers spawn child processes that inherit env; both
|
||||
`huggingface_hub` and `transformers.AutoTokenizer.from_pretrained`
|
||||
look at `HF_TOKEN` / `HUGGINGFACE_HUB_TOKEN` by default but not
|
||||
`HF_API_KEY`. If only the latter is set, gated downloads fail with
|
||||
401. Aliasing at pipeline-load time is the minimum-disruption fix.
|
||||
"""
|
||||
for src in ("HF_TOKEN", "HUGGINGFACE_HUB_TOKEN", "HF_API_KEY"):
|
||||
value = os.environ.get(src)
|
||||
if value:
|
||||
os.environ.setdefault("HF_TOKEN", value)
|
||||
os.environ.setdefault("HUGGINGFACE_HUB_TOKEN", value)
|
||||
return value
|
||||
return None
|
||||
|
||||
|
||||
class MagiHumanPipeline(ComposedPipelineBase):
|
||||
"""Base MagiHuman text-to-AV pipeline (no LoRA, no distill, no SR)."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
"audio_vae",
|
||||
]
|
||||
|
||||
def load_modules(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
loaded_modules: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Load the variant-specific transformer + scheduler from the
|
||||
converted MagiHuman repo and lazy-load the four cross-variant
|
||||
shared components from their canonical upstream HF repos:
|
||||
|
||||
* text_encoder, tokenizer -> ``google/t5gemma-9b-9b-ul2``
|
||||
(gated, requires HF token with accepted terms of use)
|
||||
* audio_vae -> ``stabilityai/stable-audio-open-1.0`` (gated)
|
||||
* vae -> ``Wan-AI/Wan2.2-TI2V-5B-Diffusers``
|
||||
|
||||
Backwards-compatible with bundled converted repos: if any of
|
||||
these subfolders is present locally and listed in
|
||||
``model_index.json``, the standard component loader picks it up
|
||||
via super(). Otherwise the loader is told to skip the entry and
|
||||
we lazy-load it here.
|
||||
"""
|
||||
# T5-Gemma is gated: expose `HF_API_KEY` as `HF_TOKEN` if needed.
|
||||
_ensure_hf_token_env()
|
||||
|
||||
# Resolve to a local cache path so we can inspect
|
||||
# model_index.json before invoking super(). `maybe_download_model`
|
||||
# is idempotent for local paths; super() repeats the call cheaply
|
||||
# via `_load_config`.
|
||||
local_path = maybe_download_model(self.model_path)
|
||||
|
||||
# Identify which cross-variant shared keys are bundled in the
|
||||
# converted repo (declared in model_index.json with a non-null
|
||||
# spec) versus absent (the umbrella scheme). Bundled keys stay
|
||||
# in `required_config_modules` and are loaded normally by super()
|
||||
# from `<model_path>/<key>/`. Absent keys are temporarily
|
||||
# dropped so super() does not fail the "every required entry
|
||||
# must appear in model_index.json" check, then lazy-loaded
|
||||
# below.
|
||||
model_index: dict[str, Any] = {}
|
||||
try:
|
||||
with open(Path(local_path) / "model_index.json") as f:
|
||||
model_index = json.load(f)
|
||||
except (FileNotFoundError, json.JSONDecodeError):
|
||||
pass
|
||||
|
||||
def _is_bundled(key: str) -> bool:
|
||||
spec = model_index.get(key)
|
||||
return (isinstance(spec, list | tuple) and len(spec) >= 1 and spec[0] is not None)
|
||||
|
||||
deferred = []
|
||||
for key in ("text_encoder", "tokenizer", "audio_vae", "vae"):
|
||||
if key in self.required_config_modules and not _is_bundled(key):
|
||||
self.required_config_modules.remove(key)
|
||||
deferred.append(key)
|
||||
|
||||
try:
|
||||
modules = super().load_modules(fastvideo_args, loaded_modules)
|
||||
finally:
|
||||
for key in deferred:
|
||||
if key not in self.required_config_modules:
|
||||
self.required_config_modules.append(key)
|
||||
|
||||
# For each lazy-load key, prefer whatever super() already loaded
|
||||
# (a bundled subfolder, or a caller-provided override merged in
|
||||
# via `loaded_modules`). Fall back to the caller-provided
|
||||
# `loaded_modules` entry for keys absent from model_index.json
|
||||
# (super() never iterates those). Otherwise lazy-load from the
|
||||
# canonical upstream HF repo.
|
||||
def _resolve(key: str) -> bool:
|
||||
"""Return True if `modules[key]` is already populated."""
|
||||
if modules.get(key) is not None:
|
||||
return True
|
||||
if loaded_modules and key in loaded_modules:
|
||||
modules[key] = loaded_modules[key]
|
||||
return True
|
||||
return False
|
||||
|
||||
if not _resolve("text_encoder"):
|
||||
logger.info("Building T5-Gemma text encoder (lazy-load from %s)", _T5GEMMA_HF_ID)
|
||||
enc_config = T5GemmaEncoderConfig()
|
||||
enc_config.arch_config.t5gemma_model_path = _T5GEMMA_HF_ID
|
||||
modules["text_encoder"] = T5GemmaEncoderModel(enc_config)
|
||||
|
||||
if not _resolve("tokenizer"):
|
||||
logger.info("Loading T5-Gemma tokenizer from %s", _T5GEMMA_HF_ID)
|
||||
modules["tokenizer"] = AutoTokenizer.from_pretrained(_T5GEMMA_HF_ID)
|
||||
|
||||
if not _resolve("audio_vae"):
|
||||
logger.info(
|
||||
"Building Stable Audio Open 1.0 VAE (lazy-load from %s) — "
|
||||
"requires HF terms accepted for gated repo",
|
||||
_SA_AUDIO_HF_ID,
|
||||
)
|
||||
audio_config = OobleckVAEConfig()
|
||||
audio_config.pretrained_path = _SA_AUDIO_HF_ID
|
||||
modules["audio_vae"] = SAAudioVAEModel(audio_config)
|
||||
|
||||
if not _resolve("vae"):
|
||||
modules["vae"] = self._load_video_vae(fastvideo_args)
|
||||
|
||||
return modules
|
||||
|
||||
def _load_video_vae(self, fastvideo_args: FastVideoArgs) -> Any:
|
||||
"""Resolve the video VAE: prefer a bundled ``vae/`` subfolder in
|
||||
the converted repo (legacy), fall back to lazy-downloading the
|
||||
Wan 2.2 TI2V-5B VAE shards from upstream.
|
||||
|
||||
Either way the load goes through FastVideo's standard
|
||||
``VAELoader`` so the result is the same FV ``AutoencoderKLWan``
|
||||
nn.Module that the bundled path produces.
|
||||
"""
|
||||
from fastvideo.models.loader.component_loader import VAELoader
|
||||
|
||||
bundled = Path(self.model_path) / "vae"
|
||||
if bundled.is_dir() and (bundled / "config.json").is_file():
|
||||
logger.info("Loading bundled video VAE from %s", bundled)
|
||||
return VAELoader().load(str(bundled), fastvideo_args)
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
logger.info(
|
||||
"Bundled vae/ not found at %s; lazy-loading Wan 2.2 TI2V-5B VAE from %s",
|
||||
self.model_path,
|
||||
_WAN_VAE_HF_ID,
|
||||
)
|
||||
snapshot = snapshot_download(
|
||||
repo_id=_WAN_VAE_HF_ID,
|
||||
allow_patterns=["vae/*"],
|
||||
)
|
||||
vae_dir = os.path.join(snapshot, "vae")
|
||||
if not os.path.isdir(vae_dir):
|
||||
raise RuntimeError(
|
||||
f"snapshot_download returned {snapshot} but no vae/ "
|
||||
f"subfolder was found inside it. Check that {_WAN_VAE_HF_ID} "
|
||||
"still exposes a Diffusers-format vae/ folder.", )
|
||||
return VAELoader().load(vae_dir, fastvideo_args)
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
# MagiHuman applies `flow_shift` during timestep setup; keep the
|
||||
# scheduler constructor at its default no-op shift.
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler()
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self._add_input_and_conditioning_stages(fastvideo_args)
|
||||
self._add_base_latent_and_denoising_stages(fastvideo_args)
|
||||
self._add_decode_stages()
|
||||
|
||||
def _add_input_and_conditioning_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
stage=InputValidationStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
)
|
||||
|
||||
self._add_reference_image_stage(fastvideo_args)
|
||||
|
||||
def _add_base_latent_and_denoising_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
dit_arch = pc.dit_config.arch_config
|
||||
|
||||
# Data-proxy + eval knobs come from the PipelineConfig (`pc`).
|
||||
# Only DiT-architecture fields live on `dit_arch` now.
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=MagiHumanLatentPreparationStage(
|
||||
vae_stride=tuple(pc.vae_stride),
|
||||
z_dim=pc.z_dim,
|
||||
patch_size=tuple(dit_arch.patch_size),
|
||||
fps=pc.fps,
|
||||
t5_gemma_target_length=pc.t5_gemma_target_length,
|
||||
coords_style=pc.coords_style,
|
||||
text_offset=pc.text_offset,
|
||||
audio_in_channels=dit_arch.audio_in_channels,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=MagiHumanDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
patch_size=tuple(dit_arch.patch_size),
|
||||
video_in_channels=dit_arch.video_in_channels,
|
||||
audio_in_channels=dit_arch.audio_in_channels,
|
||||
video_txt_guidance_scale=pc.video_txt_guidance_scale,
|
||||
audio_txt_guidance_scale=pc.audio_txt_guidance_scale,
|
||||
cfg_number=pc.cfg_number,
|
||||
coords_style=pc.coords_style,
|
||||
video_guidance_high_t_threshold=pc.video_guidance_high_t_threshold,
|
||||
video_guidance_low_t_value=pc.video_guidance_low_t_value,
|
||||
),
|
||||
)
|
||||
|
||||
def _add_decode_stages(self) -> None:
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae"), pipeline=self),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="audio_decoding_stage",
|
||||
stage=MagiHumanAudioDecodingStage(audio_vae=self.get_module("audio_vae"), ),
|
||||
)
|
||||
|
||||
def _add_reference_image_stage(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
return
|
||||
|
||||
|
||||
class MagiHumanI2VPipeline(MagiHumanPipeline):
|
||||
"""MagiHuman text+image-to-AV pipeline using the T2V DiT weights."""
|
||||
|
||||
def _add_reference_image_stage(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
self.add_stage(
|
||||
stage_name="reference_image_stage",
|
||||
stage=MagiHumanReferenceImageStage(
|
||||
vae=self.get_module("vae"),
|
||||
vae_scale_factor=pc.vae_stride[1],
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class MagiHumanSRPipeline(MagiHumanPipeline):
|
||||
"""Two-stage MagiHuman base + SR-540p text-to-AV pipeline."""
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"sr_transformer",
|
||||
"scheduler",
|
||||
"audio_vae",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self._add_input_and_conditioning_stages(fastvideo_args)
|
||||
self._add_base_latent_and_denoising_stages(fastvideo_args)
|
||||
self._add_sr_latent_and_denoising_stages(fastvideo_args)
|
||||
self._add_decode_stages()
|
||||
|
||||
def _add_sr_latent_and_denoising_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
dit_arch = pc.dit_config.arch_config
|
||||
sr_transformer = self.get_module("sr_transformer")
|
||||
sr_local_attn_layers = tuple(getattr(pc, "sr_local_attn_layers", ()))
|
||||
if sr_local_attn_layers and hasattr(sr_transformer, "configure_local_attention"):
|
||||
sr_transformer.configure_local_attention(
|
||||
sr_local_attn_layers,
|
||||
frame_receptive_field=pc.frame_receptive_field,
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="sr_latent_preparation_stage",
|
||||
stage=MagiHumanSRLatentPreparationStage(
|
||||
vae=self.get_module("vae"),
|
||||
vae_stride=tuple(pc.vae_stride),
|
||||
patch_size=tuple(dit_arch.patch_size),
|
||||
noise_value=pc.noise_value,
|
||||
sr_audio_noise_scale=pc.sr_audio_noise_scale,
|
||||
sr_height=pc.sr_height,
|
||||
sr_width=pc.sr_width,
|
||||
vae_scale_factor=pc.vae_stride[1],
|
||||
),
|
||||
)
|
||||
self.add_stage(
|
||||
stage_name="sr_denoising_stage",
|
||||
stage=MagiHumanSRDenoisingStage(
|
||||
transformer=sr_transformer,
|
||||
scheduler=self.get_module("scheduler"),
|
||||
patch_size=tuple(dit_arch.patch_size),
|
||||
video_in_channels=dit_arch.video_in_channels,
|
||||
audio_in_channels=dit_arch.audio_in_channels,
|
||||
sr_num_inference_steps=pc.sr_num_inference_steps,
|
||||
sr_video_txt_guidance_scale=pc.sr_video_txt_guidance_scale,
|
||||
use_cfg_trick=pc.use_cfg_trick,
|
||||
cfg_trick_start_frame=pc.cfg_trick_start_frame,
|
||||
cfg_trick_value=pc.cfg_trick_value,
|
||||
cfg_number=pc.cfg_number,
|
||||
coords_style="v1",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class MagiHumanSRI2VPipeline(MagiHumanSRPipeline):
|
||||
"""Two-stage MagiHuman base + SR-540p text+image-to-AV pipeline."""
|
||||
|
||||
def _add_reference_image_stage(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
pc = fastvideo_args.pipeline_config
|
||||
self.add_stage(
|
||||
stage_name="reference_image_stage",
|
||||
stage=MagiHumanReferenceImageStage(
|
||||
vae=self.get_module("vae"),
|
||||
vae_scale_factor=pc.vae_stride[1],
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class MagiHumanSR1080pPipeline(MagiHumanSRPipeline):
|
||||
"""Two-stage MagiHuman base + SR-1080p text-to-AV pipeline.
|
||||
|
||||
The stage chain is identical to SR-540p. The paired pipeline config enables
|
||||
block-sparse local-window attention on 32 SR-DiT layers and requests the
|
||||
1080p latent target.
|
||||
"""
|
||||
|
||||
|
||||
class MagiHumanSR1080pI2VPipeline(MagiHumanSRI2VPipeline):
|
||||
"""Two-stage MagiHuman base + SR-1080p text+image-to-AV pipeline."""
|
||||
|
||||
|
||||
EntryClass = [
|
||||
MagiHumanPipeline,
|
||||
MagiHumanI2VPipeline,
|
||||
MagiHumanSRPipeline,
|
||||
MagiHumanSRI2VPipeline,
|
||||
MagiHumanSR1080pPipeline,
|
||||
MagiHumanSR1080pI2VPipeline,
|
||||
]
|
||||
@@ -0,0 +1,236 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""PipelineConfig for the daVinci-MagiHuman base text-to-AV pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import MagiHumanVideoConfig
|
||||
from fastvideo.configs.models.encoders import (
|
||||
BaseEncoderOutput,
|
||||
T5GemmaEncoderConfig,
|
||||
)
|
||||
from fastvideo.configs.models.vaes import OobleckVAEConfig, WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def t5gemma_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""Return per-prompt last_hidden_state as a batched [B, L, D] tensor.
|
||||
|
||||
MagiHuman pads/trims the embedding to a fixed length in its own
|
||||
`pad_or_trim` helper at pipeline time. Here we simply hand through
|
||||
whatever the tokenizer produced — the latent-prep stage is responsible
|
||||
for pad/trim so that the original context length can be preserved.
|
||||
"""
|
||||
hidden = outputs.last_hidden_state
|
||||
assert torch.isnan(hidden).sum() == 0
|
||||
# Keep the shape the tokenizer emitted; the pipeline stage handles
|
||||
# pad-or-trim to t5_gemma_target_length=640.
|
||||
return hidden
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanBaseConfig(PipelineConfig):
|
||||
"""Base MagiHuman text-to-AV pipeline config (prompt → video + audio).
|
||||
|
||||
MagiHuman's base model is a joint audio-visual generator. This config
|
||||
wires up both the video VAE (Wan 2.2 TI2V-5B) and the audio VAE
|
||||
(Stable Audio Open 1.0); the pipeline produces an mp4 with a muxed
|
||||
audio track. The framework's `WorkloadType` enum has no `T2AV`
|
||||
variant yet, so the registry entry uses `WorkloadType.T2V` as a
|
||||
placeholder.
|
||||
"""
|
||||
|
||||
# DiT
|
||||
dit_config: DiTConfig = field(default_factory=MagiHumanVideoConfig)
|
||||
# VAE — Wan 2.2 TI2V-5B. Diffusers `vae/config.json` drives arch_config
|
||||
# at load time, including z_dim=48 and scale_factor_temporal=4 /
|
||||
# scale_factor_spatial=16.
|
||||
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Audio VAE — Stable Audio Open 1.0 (Oobleck), shared with the
|
||||
# standalone Stable Audio pipeline. Lazy-loaded from
|
||||
# `stabilityai/stable-audio-open-1.0` (HF gated, Apache 2.0).
|
||||
audio_vae_config: VAEConfig = field(default_factory=OobleckVAEConfig)
|
||||
|
||||
# Denoising (flow-matching UniPC).
|
||||
flow_shift: float | None = 5.0
|
||||
|
||||
# Text encoding
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5GemmaEncoderConfig(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (t5gemma_postprocess_text, ))
|
||||
|
||||
# Precisions — the DiT runs bf16 internally, the text encoder is
|
||||
# bf16-native, and the VAE decode path benefits from fp32 for long
|
||||
# sequences.
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
|
||||
# MagiHuman-specific defaults surfaced for the pipeline stages. These
|
||||
# are pipeline-level knobs sourced from the upstream
|
||||
# `EvaluationConfig` / `DataProxyConfig` (not `ModelConfig`), so they
|
||||
# belong here, NOT on `MagiHumanArchConfig`.
|
||||
t5_gemma_target_length: int = 640
|
||||
fps: int = 25
|
||||
num_inference_steps: int = 32
|
||||
video_txt_guidance_scale: float = 5.0
|
||||
audio_txt_guidance_scale: float = 5.0
|
||||
cfg_number: int = 2
|
||||
|
||||
# VAE / data-proxy knobs (were on ArchConfig before; moved here).
|
||||
vae_stride: tuple[int, int, int] = (4, 16, 16)
|
||||
z_dim: int = 48
|
||||
frame_receptive_field: int = 11
|
||||
coords_style: str = "v2"
|
||||
ref_audio_offset: int = 1000
|
||||
text_offset: int = 0
|
||||
|
||||
# Video CFG step-dependent guidance: low-t steps use a relaxed scale.
|
||||
# Upstream daVinci-MagiHuman/inference/pipeline/video_generate.py:426
|
||||
# uses 5.0 for high-t and 2.0 for low-t with cutoff at t=500.
|
||||
video_guidance_high_t_threshold: int = 500
|
||||
video_guidance_low_t_value: float = 2.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# Base text-to-AV does not need the VAE encoder (no reference-image
|
||||
# conditioning). Keep decoder only to save memory.
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanBaseI2VConfig(MagiHumanBaseConfig):
|
||||
"""Base MagiHuman text+image-to-AV pipeline config.
|
||||
|
||||
TI2V reuses the T2V DiT weights; the only pipeline-side difference is
|
||||
that a reference image is encoded with the Wan VAE and reinserted into
|
||||
the first video-latent frame before every denoise step.
|
||||
"""
|
||||
|
||||
image_conditioning: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanDistillConfig(MagiHumanBaseConfig):
|
||||
"""DMD-2 distilled MagiHuman text-to-AV pipeline config.
|
||||
|
||||
Same arch as base (identical 331 keys, same shapes, same module tree),
|
||||
but trained via DMD-2 for 8-step inference without classifier-free
|
||||
guidance. Weights are stored in fp32 upstream; the conversion script's
|
||||
`--cast-bf16` flag reduces the checkpoint to ~30 GB on disk.
|
||||
"""
|
||||
|
||||
num_inference_steps: int = 8
|
||||
cfg_number: int = 1 # DMD distilled models skip CFG.
|
||||
# Lower flow_shift matches the distilled DMD schedule; if parity later
|
||||
# shows drift, measure against `scheduler_config.json` generated by the
|
||||
# conversion script for the distill subfolder.
|
||||
flow_shift: float | None = 5.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanDistillI2VConfig(MagiHumanDistillConfig):
|
||||
"""DMD-2 distilled MagiHuman text+image-to-AV pipeline config."""
|
||||
|
||||
image_conditioning: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanSR540pConfig(MagiHumanBaseConfig):
|
||||
"""Two-stage MagiHuman base + SR-540p text-to-AV pipeline config."""
|
||||
|
||||
noise_value: int = 220
|
||||
sr_audio_noise_scale: float = 0.7
|
||||
sr_num_inference_steps: int = 5
|
||||
sr_video_txt_guidance_scale: float = 3.5
|
||||
use_cfg_trick: bool = True
|
||||
cfg_trick_start_frame: int = 13
|
||||
cfg_trick_value: float = 2.0
|
||||
# Upstream example/sr_540p uses sr_height=512, sr_width=896. Despite the
|
||||
# marketing name, these are the VAE/patch-aligned dimensions actually run.
|
||||
sr_height: int = 512
|
||||
sr_width: int = 896
|
||||
sr_local_attn_layers: tuple[int, ...] = ()
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanSR540pI2VConfig(MagiHumanSR540pConfig):
|
||||
"""Two-stage MagiHuman base + SR-540p text+image-to-AV config."""
|
||||
|
||||
image_conditioning: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
_SR_1080P_LOCAL_ATTN_LAYERS: tuple[int, ...] = (
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
4,
|
||||
5,
|
||||
6,
|
||||
8,
|
||||
9,
|
||||
10,
|
||||
12,
|
||||
13,
|
||||
14,
|
||||
16,
|
||||
17,
|
||||
18,
|
||||
20,
|
||||
21,
|
||||
22,
|
||||
24,
|
||||
25,
|
||||
26,
|
||||
28,
|
||||
29,
|
||||
30,
|
||||
32,
|
||||
33,
|
||||
34,
|
||||
35,
|
||||
36,
|
||||
37,
|
||||
38,
|
||||
39,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanSR1080pConfig(MagiHumanSR540pConfig):
|
||||
"""Two-stage MagiHuman base + SR-1080p text-to-AV pipeline config."""
|
||||
|
||||
sr_height: int = 1080
|
||||
sr_width: int = 1920
|
||||
sr_local_attn_layers: tuple[int, ...] = _SR_1080P_LOCAL_ATTN_LAYERS
|
||||
|
||||
|
||||
@dataclass
|
||||
class MagiHumanSR1080pI2VConfig(MagiHumanSR1080pConfig):
|
||||
"""Two-stage MagiHuman base + SR-1080p text+image-to-AV config."""
|
||||
|
||||
image_conditioning: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -0,0 +1,225 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Presets for the daVinci-MagiHuman pipelines."""
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
# Keep this in sync with upstream MagiEvaluator.negative_prompt
|
||||
# (daVinci-MagiHuman/inference/pipeline/video_generate.py:222-224): the
|
||||
# video, audio-quality, and speech-delivery blocks all condition CFG.
|
||||
_MAGI_HUMAN_NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, style, works, "
|
||||
"paintings, images, static, overall gray, worst quality, low quality, JPEG "
|
||||
"compression residue, ugly, incomplete, extra fingers, poorly drawn hands, "
|
||||
"poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, "
|
||||
"still picture, messy background, three legs, many people in the background, "
|
||||
"walking backwards, low quality, worst quality, poor quality, noise, background "
|
||||
"noise, hiss, hum, buzz, crackle, static, compression artifacts, MP3 artifacts, "
|
||||
"digital clipping, distortion, muffled, muddy, unclear, echo, reverb, room echo, "
|
||||
"over-reverberated, hollow sound, distant, washed out, harsh, shrill, piercing, "
|
||||
"grating, tinny, thin sound, boomy, bass-heavy, flat EQ, over-compressed, "
|
||||
"abrupt cut, jarring transition, sudden silence, looping artifact, music, "
|
||||
"instrumental, sirens, alarms, crowd noise, unrelated sound effects, chaotic, "
|
||||
"disorganized, messy, cheap sound, emotionless, flat delivery, deadpan, lifeless, "
|
||||
"apathetic, robotic, mechanical, monotone, flat intonation, undynamic, boring, "
|
||||
"reading from a script, AI voice, synthetic, text-to-speech, TTS, insincere, "
|
||||
"fake emotion, exaggerated, overly dramatic, melodramatic, cheesy, cringey, "
|
||||
"hesitant, unconfident, tired, weak voice, stuttering, stammering, mumbling, "
|
||||
"slurred speech, mispronounced, bad articulation, lisp, vocal fry, creaky voice, "
|
||||
"mouth clicks, lip smacks, wet mouth sounds, heavy breathing, audible inhales, "
|
||||
"plosives, p-pops, coughing, clearing throat, sneezing, speaking too fast, rushed, "
|
||||
"speaking too slow, dragged out, unnatural pauses, awkward silence, choppy, "
|
||||
"disjointed, multiple speakers, two voices, background talking, out of tune, "
|
||||
"off-key, autotune artifacts")
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Joint video+audio UniPC flow-matching denoise pass.",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
MAGI_HUMAN_BASE = InferencePreset(
|
||||
name="magi_human_base",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman base text-to-AV at 256x480, 4s @ 25 fps. "
|
||||
"Produces an mp4 with muxed audio + video. workload_type "
|
||||
"is `t2v` because the framework enum has no `t2av` variant yet."),
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
# Upstream pipeline.py:61-64 defaults br_width=480, br_height=272,
|
||||
# and video_generate.py:254-261 snaps height to 256 while width stays
|
||||
# 480, so the rendered default is 256x480.
|
||||
"width": 480,
|
||||
# num_frames is derived by the pipeline as `seconds*fps + 1`; we
|
||||
# surface it here for APIs that expect a concrete default.
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 5.0, # used as video_txt_guidance_scale
|
||||
"num_inference_steps": 32,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_DISTILL = InferencePreset(
|
||||
name="magi_human_distill",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman DMD-2 distilled text-to-AV at 256x480, 4s @ "
|
||||
"25 fps. 8-step inference, no classifier-free guidance. Produces "
|
||||
"an mp4 with muxed audio + video."),
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
# DMD: cfg=1 at the pipeline level. guidance_scale is kept at 1.0
|
||||
# for interop; the DenoisingStage ignores it when cfg_number=1.
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 8,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_BASE_TI2V = InferencePreset(
|
||||
name="magi_human_base_ti2v",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman base text+image-to-AV at 256x480, 4s @ 25 fps. "
|
||||
"The reference image is VAE-encoded and pinned to the first "
|
||||
"video latent frame at each denoise step."),
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_DISTILL_TI2V = InferencePreset(
|
||||
name="magi_human_distill_ti2v",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman DMD-2 distilled text+image-to-AV at 256x480, "
|
||||
"4s @ 25 fps. 8-step inference, no CFG."),
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 8,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_SR_540P = InferencePreset(
|
||||
name="magi_human_sr_540p",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman two-stage base + SR-540p text-to-AV. "
|
||||
"Base pass runs at 256x480; SR pass refines to upstream's "
|
||||
"aligned 512x896 output with muxed audio."),
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_SR_540P_TI2V = InferencePreset(
|
||||
name="magi_human_sr_540p_ti2v",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman two-stage base + SR-540p text+image-to-AV. "
|
||||
"The reference image is encoded at base resolution and then "
|
||||
"re-encoded at SR resolution before the SR denoise pass."),
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_SR_1080P = InferencePreset(
|
||||
name="magi_human_sr_1080p",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman two-stage base + SR-1080p text-to-AV. "
|
||||
"The SR DiT uses upstream local-window attention in 32 of "
|
||||
"40 layers and refines to 1080p-class output."),
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
MAGI_HUMAN_SR_1080P_TI2V = InferencePreset(
|
||||
name="magi_human_sr_1080p_ti2v",
|
||||
version=1,
|
||||
model_family="magi_human",
|
||||
description=("daVinci-MagiHuman two-stage base + SR-1080p text+image-to-AV. "
|
||||
"The SR DiT uses upstream local-window attention in 32 of "
|
||||
"40 layers; the reference image is re-encoded at SR resolution."),
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"seed": 42,
|
||||
"height": 256,
|
||||
"width": 480,
|
||||
"num_frames": 101,
|
||||
"fps": 25,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"negative_prompt": _MAGI_HUMAN_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (
|
||||
MAGI_HUMAN_BASE,
|
||||
MAGI_HUMAN_DISTILL,
|
||||
MAGI_HUMAN_BASE_TI2V,
|
||||
MAGI_HUMAN_DISTILL_TI2V,
|
||||
MAGI_HUMAN_SR_540P,
|
||||
MAGI_HUMAN_SR_540P_TI2V,
|
||||
MAGI_HUMAN_SR_1080P,
|
||||
MAGI_HUMAN_SR_1080P_TI2V,
|
||||
)
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.pipelines.basic.magi_human.stages.audio_decoding import MagiHumanAudioDecodingStage
|
||||
from fastvideo.pipelines.basic.magi_human.stages.denoising import MagiHumanDenoisingStage
|
||||
from fastvideo.pipelines.basic.magi_human.stages.latent_preparation import MagiHumanLatentPreparationStage
|
||||
from fastvideo.pipelines.basic.magi_human.stages.reference_image import MagiHumanReferenceImageStage
|
||||
from fastvideo.pipelines.basic.magi_human.stages.sr_denoising import MagiHumanSRDenoisingStage
|
||||
from fastvideo.pipelines.basic.magi_human.stages.sr_latent_preparation import MagiHumanSRLatentPreparationStage
|
||||
|
||||
__all__ = [
|
||||
"MagiHumanAudioDecodingStage",
|
||||
"MagiHumanDenoisingStage",
|
||||
"MagiHumanLatentPreparationStage",
|
||||
"MagiHumanReferenceImageStage",
|
||||
"MagiHumanSRDenoisingStage",
|
||||
"MagiHumanSRLatentPreparationStage",
|
||||
]
|
||||
@@ -0,0 +1,111 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Audio decoding stage for daVinci-MagiHuman.
|
||||
|
||||
Takes the denoised audio latent that `MagiHumanDenoisingStage` leaves
|
||||
on `batch.audio_latents` and decodes it to a waveform using the
|
||||
Stable Audio Open 1.0 VAE. Mirrors the upstream post-process path
|
||||
(see `MagiEvaluator.post_process` in
|
||||
daVinci-MagiHuman/inference/pipeline/video_generate.py:503):
|
||||
|
||||
latent_audio.squeeze(0) # (L, C_latent)
|
||||
audio = self.audio_vae.decode(latent_audio.T) # (1, audio_ch, samples)
|
||||
audio = audio.squeeze(0).T.cpu().numpy() # (samples, audio_ch)
|
||||
audio = resample_audio_sinc(audio, _UPSTREAM_AUDIO_TIME_STRETCH)
|
||||
|
||||
The stage stores the resampled waveform on `batch.extra["audio"]`
|
||||
(shape `[samples, audio_channels]`) and the sample rate on
|
||||
`batch.extra["audio_sample_rate"]`. FastVideo's `VideoGenerator._mux_audio`
|
||||
then reads those, writes a temp wav, and muxes it into the output mp4
|
||||
via PyAV — same plumbing LTX-2 and Stable Audio use.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from scipy.signal import resample as _scipy_resample
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
# 441/512 is daVinci-MagiHuman's audio time-stretch ratio that aligns
|
||||
# the 44.1 kHz Stable-Audio output with the 25-fps video frame rate.
|
||||
# See daVinci-MagiHuman/inference/pipeline/video_generate.py:516.
|
||||
_UPSTREAM_AUDIO_TIME_STRETCH = 441.0 / 512.0
|
||||
|
||||
# Stable Audio Open 1.0 native sample rate (per stabilityai/stable-audio-open-1.0
|
||||
# model card and fastvideo/configs/models/vaes/oobleck.py::OobleckVAEArchConfig.sampling_rate).
|
||||
_SA_AUDIO_OPEN_SAMPLE_RATE = 44100
|
||||
|
||||
|
||||
def _resample_sinc(audio: np.ndarray, time_stretching: float) -> np.ndarray:
|
||||
"""Resample the audio to ``new_length = int(L * time_stretching)`` samples.
|
||||
|
||||
Mirrors upstream ``video_process.resample_audio_sinc`` which calls
|
||||
``scipy.signal.resample`` (FFT-based polyphase resampling that
|
||||
approximates ideal sinc interpolation). This avoids the
|
||||
high-frequency aliasing and roll-off that ``F.interpolate(mode='linear')``
|
||||
would introduce on a 25 fps × ~5 s wav (`scipy` is already a direct
|
||||
fastvideo dep, so this is dependency-free relative to the previous
|
||||
implementation).
|
||||
"""
|
||||
if time_stretching == 1.0:
|
||||
return audio
|
||||
new_length = int(audio.shape[0] * time_stretching)
|
||||
resampled = _scipy_resample(audio.astype(np.float32), new_length, axis=0)
|
||||
return np.asarray(resampled, dtype=np.float32)
|
||||
|
||||
|
||||
class MagiHumanAudioDecodingStage(PipelineStage):
|
||||
"""Decode `batch.audio_latents` to a waveform using Stable Audio's VAE.
|
||||
|
||||
The VAE is loaded lazily by `SAAudioVAEModel.sa_audio_vae_model` — the
|
||||
first call triggers a snapshot_download (requires HF token + accepted
|
||||
terms on stabilityai/stable-audio-open-1.0).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
audio_vae,
|
||||
time_stretching: float = _UPSTREAM_AUDIO_TIME_STRETCH,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.audio_vae = audio_vae
|
||||
self.time_stretching = time_stretching
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
latent_audio = getattr(batch, "audio_latents", None)
|
||||
if latent_audio is None:
|
||||
# Joint AV: missing audio latents means the denoising stage broke.
|
||||
raise ValueError("MagiHumanAudioDecodingStage requires batch.audio_latents to be set. "
|
||||
"Did the denoising stage produce them? Joint AV pipeline expects "
|
||||
"both video and audio latents from MagiHumanDenoisingStage.")
|
||||
|
||||
# Upstream shape: `[B, L, C_latent]` from the DiT; AutoencoderOobleck
|
||||
# expects `[B, C_latent, L]`. MagiEvaluator.post_process does
|
||||
# `latent_audio.squeeze(0); audio_vae.decode(latent_audio.T)`
|
||||
# (which yields `[C_latent, L]`, implicit batch=1). We keep the
|
||||
# batch dim and transpose L<->C.
|
||||
latent_bcl = latent_audio.permute(0, 2, 1).contiguous()
|
||||
|
||||
# Decode: [B, C_latent, L] -> [B, audio_channels, samples]
|
||||
audio_out = self.audio_vae.decode(latent_bcl)
|
||||
|
||||
audio_np = audio_out.squeeze(0).T.float().cpu().numpy()
|
||||
audio_np = _resample_sinc(audio_np, self.time_stretching)
|
||||
|
||||
# Conform to FastVideo convention: VideoGenerator._mux_audio
|
||||
# reads these two keys and muxes via PyAV.
|
||||
if batch.extra is None:
|
||||
batch.extra = {}
|
||||
batch.extra["audio"] = audio_np
|
||||
batch.extra["audio_sample_rate"] = int(getattr(self.audio_vae, "sampling_rate", _SA_AUDIO_OPEN_SAMPLE_RATE))
|
||||
return batch
|
||||
@@ -0,0 +1,228 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Joint-modality denoising stage for daVinci-MagiHuman base text-to-AV.
|
||||
|
||||
Runs the FlowUniPC denoise loop with CFG=2 over video + audio latents
|
||||
jointly. Text embeddings are already pad-or-trimmed to `t5_gemma_target_length`
|
||||
by `MagiHumanLatentPreparationStage`; the original context lengths are
|
||||
stashed on the batch as `magi_original_text_lens` / `magi_original_neg_text_lens`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.hooks.activation_trace import trace_step
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.pipelines.basic.magi_human.stages.latent_preparation import (
|
||||
StaticPackedInputs,
|
||||
assemble_packed_inputs,
|
||||
build_static_packed_inputs,
|
||||
unpack_tokens,
|
||||
)
|
||||
|
||||
|
||||
def _dit_forward(
|
||||
dit,
|
||||
video_latent: torch.Tensor,
|
||||
audio_feat_len: int,
|
||||
txt_feat: torch.Tensor,
|
||||
txt_feat_len: int,
|
||||
static_packed: StaticPackedInputs,
|
||||
coords_style: str,
|
||||
video_in_channels: int,
|
||||
audio_in_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
x, coords, mm = assemble_packed_inputs(
|
||||
static=static_packed,
|
||||
txt_feat=txt_feat,
|
||||
txt_feat_len=txt_feat_len,
|
||||
coords_style=coords_style,
|
||||
)
|
||||
video_token_num = static_packed.video_token_num
|
||||
out = dit(x, coords, mm)
|
||||
return unpack_tokens(
|
||||
out,
|
||||
video_token_num=video_token_num,
|
||||
audio_feat_len=audio_feat_len,
|
||||
video_in_channels=video_in_channels,
|
||||
audio_in_channels=audio_in_channels,
|
||||
latent_shape=tuple(video_latent.shape),
|
||||
patch_size=patch_size,
|
||||
)
|
||||
|
||||
|
||||
def _overwrite_first_frame(
|
||||
video_latent: torch.Tensor,
|
||||
image_latent: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
if image_latent is not None:
|
||||
video_latent[:, :, :1] = image_latent.to(
|
||||
device=video_latent.device,
|
||||
dtype=video_latent.dtype,
|
||||
)[:, :, :1]
|
||||
return video_latent
|
||||
|
||||
|
||||
class MagiHumanDenoisingStage(PipelineStage):
|
||||
"""UniPC-flow joint denoising with CFG=2 over (video, audio) latents."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer,
|
||||
scheduler,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
video_in_channels: int = 192,
|
||||
audio_in_channels: int = 64,
|
||||
video_txt_guidance_scale: float = 5.0,
|
||||
audio_txt_guidance_scale: float = 5.0,
|
||||
cfg_number: int = 2,
|
||||
coords_style: str = "v2",
|
||||
video_guidance_high_t_threshold: int = 500,
|
||||
video_guidance_low_t_value: float = 2.0,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.patch_size = patch_size
|
||||
self.video_in_channels = video_in_channels
|
||||
self.audio_in_channels = audio_in_channels
|
||||
self.video_txt_guidance_scale = video_txt_guidance_scale
|
||||
self.audio_txt_guidance_scale = audio_txt_guidance_scale
|
||||
self.cfg_number = cfg_number
|
||||
self.coords_style = coords_style
|
||||
self.video_guidance_high_t_threshold = video_guidance_high_t_threshold
|
||||
self.video_guidance_low_t_value = video_guidance_low_t_value
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = batch.latents.device
|
||||
shift = fastvideo_args.pipeline_config.flow_shift
|
||||
# Video and audio use independent FlowUniPC state (upstream
|
||||
# inference/pipeline/video_generate.py:404-407 instantiates two
|
||||
# separate schedulers). Sharing one scheduler causes the
|
||||
# `model_outputs` buffer for the video step to pollute the audio
|
||||
# step's diff calculation (different shapes -> broadcast error).
|
||||
video_scheduler = copy.deepcopy(self.scheduler)
|
||||
audio_scheduler = copy.deepcopy(self.scheduler)
|
||||
video_scheduler.set_timesteps(
|
||||
batch.num_inference_steps,
|
||||
device=device,
|
||||
shift=shift,
|
||||
)
|
||||
audio_scheduler.set_timesteps(
|
||||
batch.num_inference_steps,
|
||||
device=device,
|
||||
shift=shift,
|
||||
)
|
||||
timesteps = video_scheduler.timesteps
|
||||
|
||||
video_latent = batch.latents
|
||||
audio_latent = batch.audio_latents
|
||||
image_latent = getattr(batch, "image_latent", None)
|
||||
|
||||
# Expect [1, L, 3584] text embeds plus a list of original lengths.
|
||||
txt_feat = batch.prompt_embeds[0]
|
||||
txt_feat_len = int(batch.magi_original_text_lens[0])
|
||||
|
||||
neg_txt_feat: torch.Tensor | None = None
|
||||
neg_txt_feat_len: int = 0
|
||||
if self.cfg_number == 2:
|
||||
neg_list = batch.negative_prompt_embeds or []
|
||||
if not neg_list:
|
||||
raise ValueError("CFG=2 requires negative prompt embeddings; got None. "
|
||||
"Did the prompt encoding stage run?")
|
||||
else:
|
||||
neg_txt_feat = neg_list[0]
|
||||
neg_txt_feat_len = int(batch.magi_original_neg_text_lens[0])
|
||||
|
||||
audio_feat_len = int(audio_latent.shape[1])
|
||||
|
||||
disable_tqdm = not getattr(fastvideo_args, "log_level_progress", True)
|
||||
for idx, t in enumerate(tqdm(timesteps, disable=disable_tqdm)):
|
||||
video_latent = _overwrite_first_frame(video_latent, image_latent)
|
||||
# Precompute packed video+audio tokens after any TI2V first-frame
|
||||
# overwrite. Text varies per cond/uncond call and is attached in
|
||||
# _dit_forward via assemble_packed_inputs.
|
||||
static_packed = build_static_packed_inputs(
|
||||
video_latent=video_latent,
|
||||
audio_latent=audio_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
patch_size=self.patch_size,
|
||||
coords_style=self.coords_style,
|
||||
layout=getattr(batch, "magi_static_packed_layout", None),
|
||||
)
|
||||
with trace_step(idx), set_forward_context(
|
||||
current_timestep=int(t.item()) if torch.is_tensor(t) else int(t),
|
||||
attn_metadata=None,
|
||||
):
|
||||
v_cond_video, v_cond_audio = _dit_forward(
|
||||
self.transformer,
|
||||
video_latent=video_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
txt_feat=txt_feat,
|
||||
txt_feat_len=txt_feat_len,
|
||||
static_packed=static_packed,
|
||||
coords_style=self.coords_style,
|
||||
video_in_channels=self.video_in_channels,
|
||||
audio_in_channels=self.audio_in_channels,
|
||||
patch_size=self.patch_size,
|
||||
)
|
||||
|
||||
if self.cfg_number == 2:
|
||||
v_uncond_video, v_uncond_audio = _dit_forward(
|
||||
self.transformer,
|
||||
video_latent=video_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
txt_feat=neg_txt_feat,
|
||||
txt_feat_len=neg_txt_feat_len,
|
||||
static_packed=static_packed,
|
||||
coords_style=self.coords_style,
|
||||
video_in_channels=self.video_in_channels,
|
||||
audio_in_channels=self.audio_in_channels,
|
||||
patch_size=self.patch_size,
|
||||
)
|
||||
else:
|
||||
v_uncond_video = None
|
||||
v_uncond_audio = None
|
||||
|
||||
if self.cfg_number == 2:
|
||||
video_guidance = (self.video_txt_guidance_scale
|
||||
if t > self.video_guidance_high_t_threshold else self.video_guidance_low_t_value)
|
||||
assert v_uncond_video is not None and v_uncond_audio is not None
|
||||
v_video = v_uncond_video + video_guidance * (v_cond_video - v_uncond_video)
|
||||
v_audio = v_uncond_audio + self.audio_txt_guidance_scale * (v_cond_audio - v_uncond_audio)
|
||||
else:
|
||||
v_video = v_cond_video
|
||||
v_audio = v_cond_audio
|
||||
|
||||
# Independent scheduler state per modality (see comment above).
|
||||
video_latent = video_scheduler.step(
|
||||
v_video,
|
||||
t,
|
||||
video_latent,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
audio_latent = audio_scheduler.step(
|
||||
v_audio,
|
||||
t,
|
||||
audio_latent,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
video_latent = _overwrite_first_frame(video_latent, image_latent)
|
||||
batch.latents = video_latent
|
||||
batch.audio_latents = audio_latent
|
||||
return batch
|
||||
@@ -0,0 +1,590 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Latent preparation stage for daVinci-MagiHuman base text-to-AV.
|
||||
|
||||
Produces:
|
||||
- random video latent of shape `[1, z_dim, latent_T, latent_H, latent_W]`,
|
||||
- random audio latent of shape `[1, num_frames, 64]` (the DiT jointly
|
||||
denoises both modalities),
|
||||
- padded T5-Gemma text embedding (target length 640) plus the original
|
||||
(pre-pad) context length, which the UniPC + CFG loop needs so the
|
||||
unconditional path sees the same padded length.
|
||||
|
||||
Also stakes out the per-token coords / modality map that the DiT consumes
|
||||
(replicates the reference `MagiDataProxy.process_input`).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
# Matches inference/common/sequence_schema.py in the reference.
|
||||
MODALITY_VIDEO = 0
|
||||
MODALITY_AUDIO = 1
|
||||
MODALITY_TEXT = 2
|
||||
|
||||
# Audio temporal compression ratio: 1 audio frame → 1/4 latent frame.
|
||||
# Mirrors data_proxy.py:206 `(audio_feat_len - 1) // 4 + 1` where 4 is
|
||||
# the audio VAE's temporal stride (same as vae_stride[0] for video).
|
||||
_AUDIO_TEMPORAL_COMPRESSION = 4
|
||||
|
||||
# v1 text-coord reference shape: (T=2, H=1, W=1).
|
||||
# Mirrors data_proxy.py:202 `ref_feat_shape=(2, 1, 1)` for coords_style=="v1".
|
||||
_V1_TEXT_REF_SHAPE: tuple[int, int, int] = (2, 1, 1)
|
||||
|
||||
|
||||
def _build_coords(
|
||||
shape: tuple[int, int, int],
|
||||
ref_feat_shape: tuple[int, int, int],
|
||||
offset_thw: tuple[int, int, int] = (0, 0, 0),
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> torch.Tensor:
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
ori_t, ori_h, ori_w = shape
|
||||
ref_t, ref_h, ref_w = ref_feat_shape
|
||||
offset_t, offset_h, offset_w = offset_thw
|
||||
time_rng = torch.arange(ori_t, device=device, dtype=dtype) + offset_t
|
||||
h_rng = torch.arange(ori_h, device=device, dtype=dtype) + offset_h
|
||||
w_rng = torch.arange(ori_w, device=device, dtype=dtype) + offset_w
|
||||
tg, hg, wg = torch.meshgrid(time_rng, h_rng, w_rng, indexing="ij")
|
||||
coords = torch.stack([tg, hg, wg], dim=-1).reshape(-1, 3)
|
||||
meta = torch.tensor(
|
||||
[ori_t, ori_h, ori_w, ref_t, ref_h, ref_w],
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
).expand(coords.size(0), -1)
|
||||
return torch.cat([coords, meta], dim=-1)
|
||||
|
||||
|
||||
def _pad_or_trim_dim1(t: torch.Tensor, target: int) -> tuple[torch.Tensor, int]:
|
||||
"""Pad-or-trim along dim 1. Returns (new_tensor, original_length)."""
|
||||
current = t.size(1)
|
||||
if current < target:
|
||||
pad = [0, 0, 0, target - current]
|
||||
return F.pad(t, pad, "constant", 0.0), current
|
||||
return t[:, :target], target
|
||||
|
||||
|
||||
def _img2tokens(x_t: torch.Tensor, t_patch: int, patch: int) -> torch.Tensor:
|
||||
"""Pack a video latent [B, C, T, H, W] -> [B, L, C * t_patch * patch^2].
|
||||
|
||||
Per-token feature ordering is channel-major ``(C pT pH pW)``: the DiT's
|
||||
``video_embedder`` weight was trained on the layout produced by
|
||||
upstream's grouped-conv ``UnfoldNd`` packer (channel slowest, patch
|
||||
elements fastest). Spatial-major ``(pT pH pW C)`` silently permutes the
|
||||
in-features and produces noise output. Asymmetric with
|
||||
``unpack_tokens`` which uses ``(pT pH pW C)`` to match
|
||||
``final_linear_video``'s trained output layout.
|
||||
"""
|
||||
B, C, T, H, W = x_t.shape
|
||||
assert T % t_patch == 0 and H % patch == 0 and W % patch == 0, (
|
||||
f"Latent dims {T,H,W} must divide ({t_patch}, {patch}, {patch})")
|
||||
return rearrange(
|
||||
x_t,
|
||||
"B C (T pT) (H pH) (W pW) -> B (T H W) (C pT pH pW)",
|
||||
pT=t_patch,
|
||||
pH=patch,
|
||||
pW=patch,
|
||||
).contiguous()
|
||||
|
||||
|
||||
class MagiHumanLatentPreparationStage(PipelineStage):
|
||||
"""Prepare latents, coords, modality maps, and padded text embed."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vae_stride: tuple[int, int, int] = (4, 16, 16),
|
||||
z_dim: int = 48,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
fps: int = 25,
|
||||
t5_gemma_target_length: int = 640,
|
||||
coords_style: Literal["v1", "v2"] = "v2",
|
||||
text_offset: int = 0,
|
||||
audio_in_channels: int = 64,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.vae_stride = vae_stride
|
||||
self.z_dim = z_dim
|
||||
self.patch_size = patch_size
|
||||
self.fps = fps
|
||||
self.t5_gemma_target_length = t5_gemma_target_length
|
||||
self.coords_style = coords_style
|
||||
self.text_offset = text_offset
|
||||
self.audio_in_channels = audio_in_channels
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
fps = self.fps
|
||||
# Prefer the caller-provided `batch.num_frames` (the standard
|
||||
# SamplingParam knob — production preset and SSIM tests both set
|
||||
# it). Fall back to `batch.num_seconds * fps + 1` when num_frames
|
||||
# is unset or the image-default sentinel (1). This matches
|
||||
# upstream MagiDataProxy.process_input which derives `num_frames
|
||||
# = seconds * fps + 1` and rejects values that don't satisfy
|
||||
# `(num_frames - 1) % vae_temporal_stride == 0`.
|
||||
requested_num_frames = int(getattr(batch, "num_frames", None) or 0)
|
||||
if requested_num_frames > 1:
|
||||
num_frames = requested_num_frames
|
||||
else:
|
||||
seconds = int(getattr(batch, "num_seconds", None) or 4)
|
||||
num_frames = seconds * fps + 1
|
||||
latent_T = (num_frames - 1) // 4 + 1
|
||||
|
||||
# Match upstream pipeline.py:61-64 + video_generate.py:254-261:
|
||||
# the requested 272p height snaps to 256, while width stays 480.
|
||||
br_h = int(batch.height) if batch.height else 256
|
||||
br_w = int(batch.width) if batch.width else 480
|
||||
pT, pH, pW = self.patch_size
|
||||
vt, vh, vw = self.vae_stride
|
||||
# Snap to patch granularity (matches reference).
|
||||
latent_H = (br_h // vh // pH) * pH
|
||||
latent_W = (br_w // vw // pW) * pW
|
||||
actual_H = latent_H * vh
|
||||
actual_W = latent_W * vw
|
||||
batch.height = actual_H
|
||||
batch.width = actual_W
|
||||
|
||||
generator = torch.Generator(device=device)
|
||||
if batch.seed is not None:
|
||||
generator.manual_seed(int(batch.seed))
|
||||
|
||||
# Video latent: [1, z_dim, latent_T, latent_H, latent_W]
|
||||
video_latent = torch.randn(
|
||||
(1, self.z_dim, latent_T, latent_H, latent_W),
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
image_latent = getattr(batch, "image_latent", None)
|
||||
if image_latent is not None:
|
||||
video_latent[:, :, :1] = image_latent.to(
|
||||
device=video_latent.device,
|
||||
dtype=video_latent.dtype,
|
||||
)[:, :, :1]
|
||||
# Audio latent: [1, num_frames, audio_in_channels]
|
||||
audio_latent = torch.randn(
|
||||
(1, num_frames, self.audio_in_channels),
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
# Prompt embeds: the upstream TextEncodingStage already ran. It
|
||||
# produced a list of [1, L, D] tensors per prompt. Pad/trim each
|
||||
# to the target length and store the original length so the DiT
|
||||
# stage can build the correct modality-map slices.
|
||||
padded_prompt_embeds: list[torch.Tensor] = []
|
||||
padded_prompt_lens: list[int] = []
|
||||
for embed in batch.prompt_embeds:
|
||||
# embed: [1, L, 3584]
|
||||
padded, original = _pad_or_trim_dim1(
|
||||
embed.to(torch.float32),
|
||||
target=self.t5_gemma_target_length,
|
||||
)
|
||||
padded_prompt_embeds.append(padded)
|
||||
padded_prompt_lens.append(original)
|
||||
batch.prompt_embeds = padded_prompt_embeds
|
||||
# Stash the original text length list on the batch for the denoise
|
||||
# stage — FastVideo's ForwardBatch doesn't have a first-class field
|
||||
# for this so we attach it.
|
||||
batch.magi_original_text_lens = padded_prompt_lens
|
||||
|
||||
# Matching negative prompts.
|
||||
if batch.negative_prompt_embeds is not None and batch.negative_prompt_embeds:
|
||||
padded_neg: list[torch.Tensor] = []
|
||||
padded_neg_lens: list[int] = []
|
||||
for embed in batch.negative_prompt_embeds:
|
||||
padded, original = _pad_or_trim_dim1(
|
||||
embed.to(torch.float32),
|
||||
target=self.t5_gemma_target_length,
|
||||
)
|
||||
padded_neg.append(padded)
|
||||
padded_neg_lens.append(original)
|
||||
batch.negative_prompt_embeds = padded_neg
|
||||
batch.magi_original_neg_text_lens = padded_neg_lens
|
||||
|
||||
batch.latents = video_latent
|
||||
batch.audio_latents = audio_latent
|
||||
batch.num_frames = num_frames
|
||||
batch.magi_latent_T = latent_T
|
||||
batch.magi_latent_H = latent_H
|
||||
batch.magi_latent_W = latent_W
|
||||
# Precompute the step-invariant packed layout (coords / modality
|
||||
# maps / channel-padding width) once; the denoise loop reuses it
|
||||
# every step instead of rebuilding meshgrids on each call.
|
||||
batch.magi_static_packed_layout = precompute_static_packed_layout(
|
||||
latent_shape=tuple(video_latent.shape), # type: ignore[arg-type]
|
||||
audio_feat_len=int(audio_latent.shape[1]),
|
||||
z_dim=self.z_dim,
|
||||
audio_in_channels=self.audio_in_channels,
|
||||
patch_size=self.patch_size,
|
||||
coords_style=self.coords_style,
|
||||
device=video_latent.device,
|
||||
)
|
||||
return batch
|
||||
|
||||
|
||||
class StaticPackedInputs:
|
||||
"""Step-invariant packed inputs: video+audio tokens, coords, modality map.
|
||||
|
||||
Computed once before the denoise loop; reused for every cond/uncond call.
|
||||
Text tokens are NOT included here because cond/uncond have different lengths.
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"video_tokens",
|
||||
"audio_tokens",
|
||||
"video_coords",
|
||||
"audio_coords",
|
||||
"video_mm",
|
||||
"audio_mm",
|
||||
"video_token_num",
|
||||
"audio_feat_len",
|
||||
"max_ch",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
video_tokens: torch.Tensor,
|
||||
audio_tokens: torch.Tensor,
|
||||
video_coords: torch.Tensor,
|
||||
audio_coords: torch.Tensor,
|
||||
video_mm: torch.Tensor,
|
||||
audio_mm: torch.Tensor,
|
||||
max_ch: int,
|
||||
) -> None:
|
||||
self.video_tokens = video_tokens
|
||||
self.audio_tokens = audio_tokens
|
||||
self.video_coords = video_coords
|
||||
self.audio_coords = audio_coords
|
||||
self.video_mm = video_mm
|
||||
self.audio_mm = audio_mm
|
||||
self.video_token_num = video_tokens.size(0)
|
||||
self.audio_feat_len = audio_tokens.size(0)
|
||||
self.max_ch = max_ch
|
||||
|
||||
|
||||
class StaticPackedLayout:
|
||||
"""Step- and value-invariant portion of the static packed inputs.
|
||||
|
||||
Coords, modality maps, and the channel-padding width depend only on the
|
||||
latent shape, audio length, channel widths, and patch sizes — all fixed
|
||||
for a single generation. Precompute once before the denoise loop and
|
||||
reuse on every step. Only the per-step token tensors must be rebuilt.
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"video_coords",
|
||||
"audio_coords",
|
||||
"video_mm",
|
||||
"audio_mm",
|
||||
"max_ch",
|
||||
"video_token_num",
|
||||
"audio_feat_len",
|
||||
)
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
video_coords: torch.Tensor,
|
||||
audio_coords: torch.Tensor,
|
||||
video_mm: torch.Tensor,
|
||||
audio_mm: torch.Tensor,
|
||||
max_ch: int,
|
||||
video_token_num: int,
|
||||
audio_feat_len: int,
|
||||
) -> None:
|
||||
self.video_coords = video_coords
|
||||
self.audio_coords = audio_coords
|
||||
self.video_mm = video_mm
|
||||
self.audio_mm = audio_mm
|
||||
self.max_ch = max_ch
|
||||
self.video_token_num = video_token_num
|
||||
self.audio_feat_len = audio_feat_len
|
||||
|
||||
|
||||
def precompute_static_packed_layout(
|
||||
latent_shape: tuple[int, int, int, int, int],
|
||||
audio_feat_len: int,
|
||||
z_dim: int,
|
||||
audio_in_channels: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
coords_style: Literal["v1", "v2"] = "v2",
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> StaticPackedLayout:
|
||||
"""Precompute the invariant fields used by ``build_static_packed_inputs``.
|
||||
|
||||
Arguments are derived from configs and latent shape — none depend on
|
||||
the current denoising-step values. Call this once in the latent
|
||||
preparation stage (or any pre-loop site) and pass the result via the
|
||||
``layout=`` arg of ``build_static_packed_inputs`` to skip the
|
||||
meshgrid/full() work on every step.
|
||||
"""
|
||||
pT, pH, pW = patch_size
|
||||
_, _, T, H, W = latent_shape
|
||||
if device is None:
|
||||
device = torch.device("cpu")
|
||||
|
||||
video_token_num = (T // pT) * (H // pH) * (W // pW)
|
||||
# `_img2tokens` packs to channel `z_dim * pT * pH * pW`; audio tokens
|
||||
# are `audio_in_channels` wide — both are config constants.
|
||||
max_ch = max(z_dim * pT * pH * pW, audio_in_channels)
|
||||
|
||||
video_ref_shape = (T // pT, H // pH, W // pW)
|
||||
video_coords = _build_coords(
|
||||
shape=video_ref_shape,
|
||||
ref_feat_shape=video_ref_shape,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if coords_style == "v2":
|
||||
audio_ref_t = (audio_feat_len - 1) // _AUDIO_TEMPORAL_COMPRESSION + 1
|
||||
audio_coords = _build_coords(
|
||||
shape=(audio_feat_len, 1, 1),
|
||||
ref_feat_shape=(audio_ref_t // pT, 1, 1),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
else:
|
||||
audio_coords = _build_coords(
|
||||
shape=(audio_feat_len, 1, 1),
|
||||
ref_feat_shape=(T // pT, 1, 1),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
video_mm = torch.full((video_token_num, ), MODALITY_VIDEO, dtype=torch.int64, device=device)
|
||||
audio_mm = torch.full((audio_feat_len, ), MODALITY_AUDIO, dtype=torch.int64, device=device)
|
||||
|
||||
return StaticPackedLayout(
|
||||
video_coords=video_coords,
|
||||
audio_coords=audio_coords,
|
||||
video_mm=video_mm,
|
||||
audio_mm=audio_mm,
|
||||
max_ch=max_ch,
|
||||
video_token_num=video_token_num,
|
||||
audio_feat_len=audio_feat_len,
|
||||
)
|
||||
|
||||
|
||||
def build_static_packed_inputs(
|
||||
video_latent: torch.Tensor,
|
||||
audio_latent: torch.Tensor,
|
||||
audio_feat_len: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
coords_style: Literal["v1", "v2"] = "v2",
|
||||
layout: StaticPackedLayout | None = None,
|
||||
) -> StaticPackedInputs:
|
||||
"""Build the step-invariant portion of the packed token stream.
|
||||
|
||||
Returns video+audio tokens (padded to a common channel width), their
|
||||
coords, and their modality slices. Text is excluded because cond/uncond
|
||||
differ in length; call assemble_packed_inputs to attach text per call.
|
||||
|
||||
Mirrors SingleData.token_sequence / coords_mapping / modality_mapping in
|
||||
inference/pipeline/data_proxy.py, minus the text portion.
|
||||
|
||||
When ``layout`` is provided, coords / modality maps / max_ch are taken
|
||||
from the precomputed values and only the per-step token tensors are
|
||||
rebuilt; this is the hot-path call from the denoising loop. When
|
||||
``layout`` is None the function recomputes everything from scratch
|
||||
(e.g. for one-shot tests via ``build_packed_inputs``).
|
||||
"""
|
||||
pT, pH, pW = patch_size
|
||||
assert video_latent.size(0) == 1, "batch size 1 required for MagiHuman base"
|
||||
|
||||
video_tokens = _img2tokens(video_latent, t_patch=pT, patch=pH)[0]
|
||||
audio_tokens = audio_latent[0, :audio_feat_len].contiguous()
|
||||
|
||||
if layout is not None:
|
||||
max_ch = layout.max_ch
|
||||
video_tokens = F.pad(video_tokens, (0, max_ch - video_tokens.size(-1)))
|
||||
audio_tokens = F.pad(audio_tokens, (0, max_ch - audio_tokens.size(-1)))
|
||||
return StaticPackedInputs(
|
||||
video_tokens=video_tokens,
|
||||
audio_tokens=audio_tokens,
|
||||
video_coords=layout.video_coords,
|
||||
audio_coords=layout.audio_coords,
|
||||
video_mm=layout.video_mm,
|
||||
audio_mm=layout.audio_mm,
|
||||
max_ch=max_ch,
|
||||
)
|
||||
|
||||
# Slow path: rebuild every invariant from scratch. Kept for the
|
||||
# ``build_packed_inputs`` one-shot wrapper used by tests/parity helpers.
|
||||
_, z_dim, T, H, W = video_latent.shape
|
||||
|
||||
max_ch = max(video_tokens.size(-1), audio_tokens.size(-1))
|
||||
video_tokens = F.pad(video_tokens, (0, max_ch - video_tokens.size(-1)))
|
||||
audio_tokens = F.pad(audio_tokens, (0, max_ch - audio_tokens.size(-1)))
|
||||
|
||||
device = video_tokens.device
|
||||
dtype = video_tokens.dtype
|
||||
video_token_num = video_tokens.size(0)
|
||||
|
||||
video_mm = torch.full((video_token_num, ), MODALITY_VIDEO, dtype=torch.int64, device=device)
|
||||
audio_mm = torch.full((audio_feat_len, ), MODALITY_AUDIO, dtype=torch.int64, device=device)
|
||||
|
||||
video_ref_shape = (T // pT, H // pH, W // pW)
|
||||
video_coords = _build_coords(
|
||||
shape=video_ref_shape,
|
||||
ref_feat_shape=video_ref_shape,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if coords_style == "v2":
|
||||
audio_ref_t = (audio_feat_len - 1) // _AUDIO_TEMPORAL_COMPRESSION + 1
|
||||
audio_coords = _build_coords(
|
||||
shape=(audio_feat_len, 1, 1),
|
||||
ref_feat_shape=(audio_ref_t // pT, 1, 1),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
else:
|
||||
audio_coords = _build_coords(
|
||||
shape=(audio_feat_len, 1, 1),
|
||||
ref_feat_shape=(T // pT, 1, 1),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
return StaticPackedInputs(
|
||||
video_tokens=video_tokens,
|
||||
audio_tokens=audio_tokens,
|
||||
video_coords=video_coords,
|
||||
audio_coords=audio_coords,
|
||||
video_mm=video_mm,
|
||||
audio_mm=audio_mm,
|
||||
max_ch=max_ch,
|
||||
)
|
||||
|
||||
|
||||
def assemble_packed_inputs(
|
||||
static: StaticPackedInputs,
|
||||
txt_feat: torch.Tensor,
|
||||
txt_feat_len: int,
|
||||
coords_style: Literal["v1", "v2"] = "v2",
|
||||
text_offset: int = 0,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Attach per-call text tokens to the precomputed static packed inputs.
|
||||
|
||||
Returns (token_seq, coords, modality_map) ready for the DiT.
|
||||
"""
|
||||
text_tokens = txt_feat[0, :txt_feat_len].contiguous()
|
||||
max_ch = max(static.max_ch, text_tokens.size(-1))
|
||||
|
||||
video_tokens = F.pad(static.video_tokens, (0, max_ch - static.video_tokens.size(-1)))
|
||||
audio_tokens = F.pad(static.audio_tokens, (0, max_ch - static.audio_tokens.size(-1)))
|
||||
text_tokens = F.pad(text_tokens, (0, max_ch - text_tokens.size(-1)))
|
||||
token_seq = torch.cat([video_tokens, audio_tokens, text_tokens], dim=0)
|
||||
|
||||
device = token_seq.device
|
||||
dtype = token_seq.dtype
|
||||
text_mm = torch.full((txt_feat_len, ), MODALITY_TEXT, dtype=torch.int64, device=device)
|
||||
mm = torch.cat([static.video_mm, static.audio_mm, text_mm], dim=0)
|
||||
|
||||
if coords_style == "v2":
|
||||
text_coords = _build_coords(
|
||||
shape=(txt_feat_len, 1, 1),
|
||||
ref_feat_shape=(1, 1, 1),
|
||||
offset_thw=(-txt_feat_len, 0, 0),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
else:
|
||||
text_coords = _build_coords(
|
||||
shape=(txt_feat_len, 1, 1),
|
||||
ref_feat_shape=_V1_TEXT_REF_SHAPE,
|
||||
offset_thw=(text_offset, 0, 0),
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
coords = torch.cat([static.video_coords, static.audio_coords, text_coords], dim=0)
|
||||
return token_seq, coords, mm
|
||||
|
||||
|
||||
def build_packed_inputs(
|
||||
video_latent: torch.Tensor,
|
||||
audio_latent: torch.Tensor,
|
||||
audio_feat_len: int,
|
||||
txt_feat: torch.Tensor,
|
||||
txt_feat_len: int,
|
||||
patch_size: tuple[int, int, int],
|
||||
coords_style: Literal["v1", "v2"] = "v2",
|
||||
text_offset: int = 0,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Build the full packed token stream in one call (backwards-compat wrapper).
|
||||
|
||||
Equivalent to assemble_packed_inputs(build_static_packed_inputs(...), ...).
|
||||
Prefer calling the two helpers separately when the static portion can be
|
||||
reused across multiple calls (e.g. cond/uncond in the denoise loop).
|
||||
"""
|
||||
static = build_static_packed_inputs(
|
||||
video_latent=video_latent,
|
||||
audio_latent=audio_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
patch_size=patch_size,
|
||||
coords_style=coords_style,
|
||||
)
|
||||
return assemble_packed_inputs(
|
||||
static=static,
|
||||
txt_feat=txt_feat,
|
||||
txt_feat_len=txt_feat_len,
|
||||
coords_style=coords_style,
|
||||
text_offset=text_offset,
|
||||
)
|
||||
|
||||
|
||||
def unpack_tokens(
|
||||
output: torch.Tensor, # [L, max(V_ch, A_ch)]
|
||||
video_token_num: int,
|
||||
audio_feat_len: int,
|
||||
video_in_channels: int,
|
||||
audio_in_channels: int,
|
||||
latent_shape: tuple[int, int, int, int, int], # [1, z_dim, T, H, W]
|
||||
patch_size: tuple[int, int, int],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Inverse of `build_packed_inputs` for the DiT output.
|
||||
|
||||
Splits the flat output back into a video latent (un-patched into
|
||||
B C T H W) and an audio latent (B, L, 64).
|
||||
"""
|
||||
pT, pH, pW = patch_size
|
||||
_, z_dim, T, H, W = latent_shape
|
||||
tH, tW = H // pH, W // pW
|
||||
|
||||
video_flat = output[:video_token_num, :video_in_channels]
|
||||
video_latent = rearrange(
|
||||
video_flat,
|
||||
"(T H W) (pT pH pW C) -> C (T pT) (H pH) (W pW)",
|
||||
H=tH,
|
||||
W=tW,
|
||||
pT=pT,
|
||||
pH=pH,
|
||||
pW=pW,
|
||||
).contiguous().unsqueeze(0)
|
||||
|
||||
audio_latent = output[
|
||||
video_token_num:video_token_num + audio_feat_len,
|
||||
:audio_in_channels,
|
||||
].unsqueeze(0)
|
||||
|
||||
return video_latent, audio_latent
|
||||
@@ -0,0 +1,101 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Reference-image encoding for MagiHuman TI2V.
|
||||
|
||||
The upstream daVinci-MagiHuman TI2V path encodes the user image through the
|
||||
Wan VAE and overwrites the first denoising latent frame with that clean latent
|
||||
at every step. This stage mirrors `MagiEvaluator.encode_image` and stashes the
|
||||
normalized latent on `batch.image_latent` for the latent-prep and denoise stages.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from diffusers.utils import load_image
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
def _resizecrop(image: Image.Image, height: int, width: int) -> Image.Image:
|
||||
"""Mirror upstream `resizecrop`: center-crop to target aspect ratio."""
|
||||
current_width, current_height = image.size
|
||||
if current_width == width and current_height == height:
|
||||
return image
|
||||
if current_height / current_width > height / width:
|
||||
new_width = int(current_width)
|
||||
new_height = int(new_width * height / width)
|
||||
else:
|
||||
new_height = int(current_height)
|
||||
new_width = int(new_height * width / height)
|
||||
left = (current_width - new_width) / 2
|
||||
top = (current_height - new_height) / 2
|
||||
right = (current_width + new_width) / 2
|
||||
bottom = (current_height + new_height) / 2
|
||||
return image.crop((left, top, right, bottom))
|
||||
|
||||
|
||||
class MagiHumanReferenceImageStage(PipelineStage):
|
||||
"""Encode a TI2V reference image into the first-frame video latent."""
|
||||
|
||||
def __init__(self, vae: Any, vae_scale_factor: int = 16) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=vae_scale_factor)
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
image = getattr(batch, "image", None) or batch.pil_image
|
||||
if image is None and batch.image_path is not None:
|
||||
image = load_image(batch.image_path)
|
||||
if image is None:
|
||||
raise ValueError("MagiHuman TI2V requires `image_path` or `pil_image`.")
|
||||
if not isinstance(image, Image.Image):
|
||||
raise TypeError(f"MagiHuman TI2V expects a PIL image or image path, got {type(image)}")
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("MagiHuman TI2V requires concrete height and width before image encoding.")
|
||||
|
||||
height = int(batch.height)
|
||||
width = int(batch.width)
|
||||
device = get_local_torch_device()
|
||||
|
||||
image = _resizecrop(image.convert("RGB"), height, width)
|
||||
image_tensor = self.video_processor.preprocess(
|
||||
image,
|
||||
height=height,
|
||||
width=width,
|
||||
).to(device=device, dtype=torch.float32)
|
||||
image_tensor = image_tensor.unsqueeze(2)
|
||||
|
||||
self.vae = self.vae.to(device)
|
||||
encoded = self.vae.encode(image_tensor)
|
||||
image_latent = encoded.mean if hasattr(encoded, "mean") else encoded
|
||||
|
||||
# FastVideo's Wan VAE returns unnormalized posterior means; upstream
|
||||
# `WanVAE.encode` applies `(mu - mean) / std` before returning.
|
||||
shift_factor = getattr(self.vae, "shift_factor", None)
|
||||
if shift_factor is not None:
|
||||
if isinstance(shift_factor, torch.Tensor):
|
||||
image_latent = image_latent - shift_factor.to(image_latent.device, image_latent.dtype)
|
||||
else:
|
||||
image_latent = image_latent - shift_factor
|
||||
scaling_factor = getattr(self.vae, "scaling_factor", None)
|
||||
if scaling_factor is not None:
|
||||
if isinstance(scaling_factor, torch.Tensor):
|
||||
image_latent = image_latent * scaling_factor.to(image_latent.device, image_latent.dtype)
|
||||
else:
|
||||
image_latent = image_latent * scaling_factor
|
||||
|
||||
batch.image_latent = image_latent.to(torch.float32)
|
||||
return batch
|
||||
@@ -0,0 +1,156 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""SR video-only denoising stage for daVinci-MagiHuman SR-540p."""
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.hooks.activation_trace import trace_step
|
||||
from fastvideo.pipelines.basic.magi_human.stages.denoising import (
|
||||
_dit_forward,
|
||||
_overwrite_first_frame,
|
||||
)
|
||||
from fastvideo.pipelines.basic.magi_human.stages.latent_preparation import (
|
||||
build_static_packed_inputs, )
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
|
||||
class MagiHumanSRDenoisingStage(PipelineStage):
|
||||
"""Denoise only the SR video latent; audio passes through unchanged."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer,
|
||||
scheduler,
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2),
|
||||
video_in_channels: int = 192,
|
||||
audio_in_channels: int = 64,
|
||||
sr_num_inference_steps: int = 5,
|
||||
sr_video_txt_guidance_scale: float = 3.5,
|
||||
use_cfg_trick: bool = True,
|
||||
cfg_trick_start_frame: int = 13,
|
||||
cfg_trick_value: float = 2.0,
|
||||
cfg_number: int = 2,
|
||||
coords_style: str = "v1",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.patch_size = patch_size
|
||||
self.video_in_channels = video_in_channels
|
||||
self.audio_in_channels = audio_in_channels
|
||||
self.sr_num_inference_steps = sr_num_inference_steps
|
||||
self.sr_video_txt_guidance_scale = sr_video_txt_guidance_scale
|
||||
self.use_cfg_trick = use_cfg_trick
|
||||
self.cfg_trick_start_frame = cfg_trick_start_frame
|
||||
self.cfg_trick_value = cfg_trick_value
|
||||
self.cfg_number = cfg_number
|
||||
self.coords_style = coords_style
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch, fastvideo_args):
|
||||
return VerificationResult()
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
device = batch.latents.device
|
||||
shift = fastvideo_args.pipeline_config.flow_shift
|
||||
video_scheduler = copy.deepcopy(self.scheduler)
|
||||
video_scheduler.set_timesteps(
|
||||
self.sr_num_inference_steps,
|
||||
device=device,
|
||||
shift=shift,
|
||||
)
|
||||
|
||||
video_latent = batch.latents
|
||||
audio_latent = batch.audio_latents
|
||||
audio_feat_len = int(audio_latent.shape[1])
|
||||
image_latent = getattr(batch, "image_latent", None)
|
||||
|
||||
txt_feat = batch.prompt_embeds[0]
|
||||
txt_feat_len = int(batch.magi_original_text_lens[0])
|
||||
|
||||
neg_txt_feat: torch.Tensor | None = None
|
||||
neg_txt_feat_len = 0
|
||||
if self.cfg_number == 2:
|
||||
neg_list = batch.negative_prompt_embeds or []
|
||||
if not neg_list:
|
||||
raise ValueError("SR CFG=2 requires negative prompt embeddings.")
|
||||
neg_txt_feat = neg_list[0]
|
||||
neg_txt_feat_len = int(batch.magi_original_neg_text_lens[0])
|
||||
|
||||
latent_length = video_latent.shape[2]
|
||||
guidance = torch.tensor(
|
||||
self.sr_video_txt_guidance_scale,
|
||||
device=device,
|
||||
dtype=video_latent.dtype,
|
||||
).expand(1, 1, latent_length, 1, 1).clone()
|
||||
if self.use_cfg_trick:
|
||||
guidance[:, :, :self.cfg_trick_start_frame] = min(
|
||||
self.cfg_trick_value,
|
||||
self.sr_video_txt_guidance_scale,
|
||||
)
|
||||
|
||||
disable_tqdm = not getattr(fastvideo_args, "log_level_progress", True)
|
||||
for idx, t in enumerate(tqdm(video_scheduler.timesteps, disable=disable_tqdm)):
|
||||
video_latent = _overwrite_first_frame(video_latent, image_latent)
|
||||
static_packed = build_static_packed_inputs(
|
||||
video_latent=video_latent,
|
||||
audio_latent=audio_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
patch_size=self.patch_size,
|
||||
coords_style=self.coords_style,
|
||||
layout=getattr(batch, "magi_static_packed_layout", None),
|
||||
)
|
||||
with trace_step(idx), set_forward_context(
|
||||
current_timestep=int(t.item()) if torch.is_tensor(t) else int(t),
|
||||
attn_metadata=None,
|
||||
):
|
||||
v_cond_video, _ = _dit_forward(
|
||||
self.transformer,
|
||||
video_latent=video_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
txt_feat=txt_feat,
|
||||
txt_feat_len=txt_feat_len,
|
||||
static_packed=static_packed,
|
||||
coords_style=self.coords_style,
|
||||
video_in_channels=self.video_in_channels,
|
||||
audio_in_channels=self.audio_in_channels,
|
||||
patch_size=self.patch_size,
|
||||
)
|
||||
if self.cfg_number == 2:
|
||||
assert neg_txt_feat is not None
|
||||
v_uncond_video, _ = _dit_forward(
|
||||
self.transformer,
|
||||
video_latent=video_latent,
|
||||
audio_feat_len=audio_feat_len,
|
||||
txt_feat=neg_txt_feat,
|
||||
txt_feat_len=neg_txt_feat_len,
|
||||
static_packed=static_packed,
|
||||
coords_style=self.coords_style,
|
||||
video_in_channels=self.video_in_channels,
|
||||
audio_in_channels=self.audio_in_channels,
|
||||
patch_size=self.patch_size,
|
||||
)
|
||||
v_video = v_uncond_video + guidance * (v_cond_video - v_uncond_video)
|
||||
else:
|
||||
v_video = v_cond_video
|
||||
|
||||
video_latent = video_scheduler.step(
|
||||
v_video,
|
||||
t,
|
||||
video_latent,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
batch.latents = _overwrite_first_frame(video_latent, image_latent)
|
||||
batch.audio_latents = audio_latent
|
||||
return batch
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user