Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c2e89f22d3 | ||
|
|
e47a3c5aad | ||
|
|
d40fbfc534 | ||
|
|
126a52ce32 | ||
|
|
419e1c68f1 | ||
|
|
938bc3c972 | ||
|
|
381a7aae66 | ||
|
|
7fa4fb50bb | ||
|
|
de3cd6aab2 | ||
|
|
a64e5e62ab | ||
|
|
f4200cc3f3 | ||
|
|
363230b372 | ||
|
|
1dcdacc35a | ||
|
|
0a7a4a21f6 | ||
|
|
c30d731992 | ||
|
|
f8eaff4292 | ||
|
|
76e7048b3d | ||
|
|
b1386a7a78 | ||
|
|
664d6b3c23 | ||
|
|
0688c5131a | ||
|
|
0b605d0a41 | ||
|
|
8f5ebe2aeb | ||
|
|
84803076a0 | ||
|
|
5cc337fdb6 | ||
|
|
2b8e5a56a8 | ||
|
|
e117167eb9 | ||
|
|
4f9af56b78 | ||
|
|
33efe7a673 | ||
|
|
b07cc3f9d4 | ||
|
|
29eb4109cc | ||
|
|
d07a7691fe | ||
|
|
960485519f | ||
|
|
412c95b1a0 | ||
|
|
f6275f8005 |
@@ -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,5 +6,3 @@
|
||||
{"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).
|
||||
|
||||
@@ -1,343 +0,0 @@
|
||||
---
|
||||
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,41 +1,26 @@
|
||||
---
|
||||
name: seed-ssim-references
|
||||
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.
|
||||
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.
|
||||
---
|
||||
|
||||
# Seed SSIM Reference Artefacts (mp4 or pt)
|
||||
# Seed SSIM Reference Videos
|
||||
|
||||
## Purpose
|
||||
|
||||
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
|
||||
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.
|
||||
|
||||
reference videos exist on the HF dataset (`FastVideo/ssim-reference-videos`).
|
||||
This skill:
|
||||
|
||||
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
|
||||
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
|
||||
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 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).
|
||||
and writes an mp4 without crashing. The skill does not re-test locally; it
|
||||
goes straight to Modal L40S (which is what CI uses).
|
||||
|
||||
## When to use
|
||||
|
||||
@@ -84,7 +69,7 @@ Fail fast if the token env var is missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Ask for the test file, then detect artefact type
|
||||
### 1. Ask for the test file
|
||||
|
||||
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`)"*.
|
||||
@@ -95,22 +80,6 @@ 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
|
||||
@@ -123,19 +92,9 @@ TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
```
|
||||
|
||||
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.
|
||||
Then launch the Modal run:
|
||||
|
||||
```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)" \
|
||||
@@ -147,19 +106,6 @@ 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.
|
||||
@@ -197,59 +143,17 @@ get` preserves that trailing `generated_videos/` segment.
|
||||
|
||||
### 4. PAUSE — user reviews quality
|
||||
|
||||
Type-aware verification.
|
||||
|
||||
**For `ARTEFACT_TYPE = pixel`** — list the downloaded mp4s and ask the user to
|
||||
open them in a video player:
|
||||
Print the list of downloaded mp4s and their paths, then stop. Tell the user:
|
||||
|
||||
> "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 artefacts. Single command works for both
|
||||
artefact types because `_iter_reference_files` walks `.mp4` and `.pt`:
|
||||
Scoped copy — only the new test's mp4s. Loop over each `<model_id>` extracted
|
||||
in step 1:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
@@ -259,13 +163,12 @@ 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,pt}`
|
||||
downloaded tree; `copy-local` walks all `<model>/<backend>/*.mp4`
|
||||
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 for pixel: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
Result for latent: same path with `.pt` extension.
|
||||
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
|
||||
### 6. Upload to HF — scoped per model_id, with overwrite guard
|
||||
|
||||
@@ -298,54 +201,33 @@ 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. 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>`)
|
||||
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>`)
|
||||
and retry from step 2.
|
||||
- **`./generated_videos_modal/default/L40S_reference_videos/` missing after
|
||||
`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.
|
||||
`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).
|
||||
- **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 artefacts stay on disk for
|
||||
- **Quality looks wrong in step 4.** Abort. The mp4s 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 (pixel SSIM drifts across SKUs; latent cosine has
|
||||
tighter cross-SKU bf16 drift but the configured tolerances assume
|
||||
same-SKU seed → same-SKU verify).
|
||||
L40S runs can't match (SSIM drifts across SKUs).
|
||||
- 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
|
||||
|
||||
@@ -354,17 +236,10 @@ 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` — 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`.
|
||||
- `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`.
|
||||
|
||||
## Changelog
|
||||
|
||||
@@ -373,4 +248,3 @@ 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). |
|
||||
|
||||
@@ -4,21 +4,24 @@ description: How to develop, validate, and register a new evaluation metric
|
||||
|
||||
# Evaluation Development SOP
|
||||
|
||||
Standard procedure for adding new video quality evaluation metrics to the
|
||||
FastVideo agent toolkit.
|
||||
Standard procedure for adding new video quality evaluation metrics to
|
||||
the FastVideo agent toolkit.
|
||||
|
||||
## When to Use
|
||||
## When to use
|
||||
|
||||
- You need a metric that doesn't exist in `.agents/memory/evaluation-registry/README.md`.
|
||||
- You need a metric that does not exist in
|
||||
`.agents/memory/evaluation-registry/README.md`.
|
||||
- An existing metric needs significant changes to its methodology.
|
||||
- You're exploring a new evaluation approach.
|
||||
- You are exploring a new evaluation approach.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Research
|
||||
|
||||
- Search `.agents/memory/related-work/` for existing evaluation approaches.
|
||||
- Check the `evaluation_registry.md` for current metrics and their limitations.
|
||||
- Search `.agents/memory/related-work/` for existing evaluation
|
||||
approaches.
|
||||
- Check `.agents/memory/evaluation-registry/README.md` for current
|
||||
metrics and their limitations.
|
||||
- Review literature: FVD, CLIP-Score, human preference, etc.
|
||||
|
||||
### 2. Prototype
|
||||
@@ -29,21 +32,25 @@ FastVideo agent toolkit.
|
||||
|
||||
### 3. Validate
|
||||
|
||||
- **Known-good test**: Metric should score high on reference-quality videos.
|
||||
- **Known-bad test**: Metric should score low on degraded/unrelated videos.
|
||||
- **Sensitivity test**: Small quality differences should produce meaningful
|
||||
score differences.
|
||||
- **Known-good test**: metric should score high on reference-quality
|
||||
videos.
|
||||
- **Known-bad test**: metric should score low on degraded or unrelated
|
||||
videos.
|
||||
- **Sensitivity test**: small quality differences should produce
|
||||
meaningful score differences.
|
||||
- Document thresholds and their justification.
|
||||
|
||||
### 4. Register
|
||||
|
||||
Update `.agents/memory/evaluation-registry/README.md`:
|
||||
|
||||
- Add the metric with status `Active`.
|
||||
- Document location, thresholds, and trust level.
|
||||
|
||||
### 5. Integrate
|
||||
|
||||
Update `.agents/skills/evaluate-video-quality.md`:
|
||||
Update `.agents/skills/evaluate-video-quality/SKILL.md`:
|
||||
|
||||
- Add the new metric as a section.
|
||||
- Include code examples and interpretation guide.
|
||||
|
||||
@@ -52,3 +59,35 @@ Update `.agents/skills/evaluate-video-quality.md`:
|
||||
- Move the exploration log content into the skill.
|
||||
- Clean up the exploration file or mark it as `promoted`.
|
||||
- If anything went wrong during development, create a lesson.
|
||||
|
||||
## Where the metrics live
|
||||
|
||||
The eval suite is `fastvideo/eval/`. New metrics register themselves
|
||||
via `@register("<group>.<name>")` and are auto-discovered when
|
||||
`fastvideo.eval.metrics` is imported.
|
||||
|
||||
- **Native metrics** (SSIM, PSNR, LPIPS, optical flow, VLM): add a
|
||||
file under the appropriate group dir
|
||||
(`fastvideo/eval/metrics/common/`, `optical_flow/`, `videoscore2/`,
|
||||
`physics_iq/`).
|
||||
- **Metrics that wrap upstream research code**: follow the vbench
|
||||
pattern in `fastvideo/eval/metrics/vbench/`. The contract is:
|
||||
- Upstream lives as a git submodule under
|
||||
`fastvideo/third_party/eval/<bench>/`, pinned to a SHA in repo-root
|
||||
`.gitmodules`.
|
||||
- The metric package's `__init__.py` inserts the submodule path on
|
||||
`sys.path` and installs runtime compat shims (attribute-level
|
||||
monkey-patches) for any modern-dep drift. Do not modify upstream
|
||||
files on disk, and do not ship a `setup.sh`.
|
||||
- See `fastvideo/eval/README.md` for the worked vbench example.
|
||||
- Full porting guide:
|
||||
[`docs/contributing/eval-metrics.md`](../../docs/contributing/eval-metrics.md).
|
||||
|
||||
## Out of scope of the initial eval port
|
||||
|
||||
The following land in follow-up PRs:
|
||||
|
||||
- **MIND** metrics (depends on a separate `vipe` submodule).
|
||||
- **VBench-2.0** sibling package.
|
||||
- Native conversion of **FVD** under `fastvideo/eval/metrics/fvd/`.
|
||||
- The training-time `EvalCallback`.
|
||||
|
||||
@@ -29,8 +29,8 @@
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
],
|
||||
"run_config": {
|
||||
"num_warmup_runs": 2,
|
||||
"num_measurement_runs": 5,
|
||||
"num_warmup_runs": 1,
|
||||
"num_measurement_runs": 3,
|
||||
"required_gpus": 2
|
||||
},
|
||||
"thresholds": {
|
||||
|
||||
@@ -15,21 +15,8 @@ log "Project root: $PROJECT_ROOT"
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
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
|
||||
|
||||
python3 -m pip install modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
@@ -76,72 +63,7 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
upload_performance_artifacts() {
|
||||
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
|
||||
LOCAL_DIR="downloaded_reports"
|
||||
|
||||
_download_reports() {
|
||||
log "Downloading perf_reports/ from Modal Volume..."
|
||||
mkdir -p "$LOCAL_DIR"
|
||||
if ! modal volume get hf-model-weights "perf_reports/" "$LOCAL_DIR"; then
|
||||
log "Error: Failed to download perf_reports/ from Modal Volume."
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_dashboard() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_${SHORT_SHA}_*" | head -n 1)
|
||||
log "TARGET dashboard: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
log "Found dashboard: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
buildkite-agent annotate --style info --context "perf-dashboard" < "$target"
|
||||
else
|
||||
log "Warning: Could not find a dashboard file matching $SHORT_SHA"
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_perf_summary() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "perf_${SHORT_SHA}_*" | head -n 1)
|
||||
log "TARGET perf summary: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
log "Found perf summary: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
buildkite-agent annotate --style info --context "perf-summary" < "$target"
|
||||
else
|
||||
log "Warning: Could not find a perf summary file matching $SHORT_SHA"
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_modal_volume() {
|
||||
log "Cleaning up perf_reports/ from Modal Volume..."
|
||||
if modal volume rm hf-model-weights "perf_reports/" --recursive; then
|
||||
log "Successfully deleted perf_reports/ from Modal Volume."
|
||||
else
|
||||
log "Warning: Failed to delete perf_reports/ from Modal Volume. Manual cleanup may be required."
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_local() {
|
||||
log "Cleaning up local download directory..."
|
||||
rm -rf "$LOCAL_DIR"
|
||||
}
|
||||
|
||||
# --- Main flow ---
|
||||
_download_reports || { _cleanup_local; return 1; }
|
||||
_upload_dashboard
|
||||
_upload_perf_summary
|
||||
_cleanup_modal_volume
|
||||
_cleanup_local
|
||||
}
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
@@ -202,9 +124,8 @@ case "$TEST_TYPE" in
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_lora_extraction_tests"
|
||||
;;
|
||||
"performance")
|
||||
log "Running performance tests on Modal..."
|
||||
log "Running performance tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_performance_tests"
|
||||
POST_RUN_HOOK="upload_performance_artifacts"
|
||||
;;
|
||||
"api_server")
|
||||
log "Running API server integration tests..."
|
||||
@@ -226,10 +147,5 @@ else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
if [ -n "$POST_RUN_HOOK" ]; then
|
||||
log "Executing post-run hook: $POST_RUN_HOOK"
|
||||
"$POST_RUN_HOOK"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
|
||||
@@ -13,21 +13,8 @@ log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
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
|
||||
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
|
||||
@@ -37,11 +37,10 @@ jobs:
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv pip install --system -r requirements-mkdocs.txt
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
@@ -56,11 +56,10 @@ jobs:
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install build dependencies
|
||||
run: uv pip install --system build twine wheel
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install build twine wheel
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
|
||||
@@ -131,13 +131,11 @@ 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: |
|
||||
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}}
|
||||
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}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
@@ -147,20 +145,20 @@ jobs:
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
uv pip install --system setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
|
||||
pip install 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
|
||||
uv pip install --system auditwheel
|
||||
pip install auditwheel
|
||||
# Point auditwheel at torch libs, but do not vendor them into the wheel.
|
||||
TORCH_LIB_DIR=$(python - <<'PY'
|
||||
import os
|
||||
@@ -213,13 +211,10 @@ jobs:
|
||||
pattern: 'fastvideo_kernel-py*'
|
||||
merge-multiple: true
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
uv pip install --system build scikit-build-core cmake ninja
|
||||
|
||||
pip install 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,11 +92,3 @@ 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/
|
||||
|
||||
@@ -4,3 +4,6 @@
|
||||
[submodule "fastvideo-kernel/include/cutlass"]
|
||||
path = fastvideo-kernel/include/cutlass
|
||||
url = https://github.com/NVIDIA/cutlass.git
|
||||
[submodule "fastvideo/third_party/eval/vbench"]
|
||||
path = fastvideo/third_party/eval/vbench
|
||||
url = https://github.com/Vchitect/VBench.git
|
||||
|
||||
@@ -7,13 +7,20 @@ 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
|
||||
.github/workflows/_template-build-image.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
)
|
||||
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,8 +23,7 @@
|
||||
- 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`.
|
||||
- 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).
|
||||
- Target line length is 80.
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
@@ -55,31 +54,3 @@ 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: uv pip install transformers")
|
||||
raise ImportError("Please install transformers: 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: uv pip install transformers")
|
||||
raise ImportError("Please install transformers: 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"uv pip install huggingface_hub") from e
|
||||
f"pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
uv pip install -q opencv-python-headless transformers huggingface_hub
|
||||
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
|
||||
uv pip install -q opencv-python-headless
|
||||
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 with `uv`.
|
||||
Currently, the only dependency is `fastvideo`, which can be installed using pip.
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
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.9.4/flash_attn-2.8.3+cu128torch2.11-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.7.16/flash_attn-2.8.3+cu128torch2.10-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.9.4/flash_attn-2.8.3+cu128torch2.11-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.7.16/flash_attn-2.8.3+cu128torch2.10-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.9.4/flash_attn-2.8.3+cu128torch2.11-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.7.16/flash_attn-2.8.3+cu128torch2.10-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
|
||||
uv pip install -r requirements-mkdocs.txt
|
||||
pip install -r requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -1,210 +0,0 @@
|
||||
# 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. |
|
||||
@@ -296,10 +296,8 @@ Action:
|
||||
|
||||
- Add or reuse a numerical parity test that loads the official model and the
|
||||
FastVideo model and compares outputs.
|
||||
- See examples in `tests/local_tests/` organized by model family
|
||||
(e.g., `tests/local_tests/sd35/`, `tests/local_tests/ltx2/`,
|
||||
`tests/local_tests/stable_audio/`) and the navigation index in
|
||||
`tests/local_tests/README.md`.
|
||||
- See examples in `tests/local_tests/` (e.g., `tests/local_tests/upsamplers/`)
|
||||
and the commands in `tests/local_tests/README.md`.
|
||||
- If there are discrepancies, add opt‑in logging to both models and compare
|
||||
activation summaries (layer output sums, per‑stage logs).
|
||||
- First align the loaded weights (validate `param_names_mapping`).
|
||||
@@ -350,8 +348,7 @@ Purpose:
|
||||
|
||||
Action:
|
||||
|
||||
- Add a pipeline parity test under `tests/local_tests/<family>/`
|
||||
(e.g., `tests/local_tests/<family>/test_<family>_pipeline_parity.py`).
|
||||
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
|
||||
- See the [Testing Guide](testing.md) for test conventions.
|
||||
|
||||
### 7) Add user‑facing examples
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -0,0 +1,559 @@
|
||||
# Porting Eval Metrics into `fastvideo.eval`
|
||||
|
||||
This guide is for contributors adding new evaluation metrics to
|
||||
FastVideo's eval suite. To run the existing metrics, see
|
||||
[`fastvideo/eval/README.md`](../../fastvideo/eval/README.md).
|
||||
|
||||
## When to use this guide
|
||||
|
||||
Use this guide when you are:
|
||||
|
||||
- Adding a new metric (native or wrapping a third-party library).
|
||||
- Porting a benchmark (e.g. VBench, MIND, EvalCrafter) whose Python
|
||||
code needs to be importable from a pinned upstream.
|
||||
- Adding a new metric group (audio, vlm, etc.).
|
||||
|
||||
## TL;DR
|
||||
|
||||
Metrics are auto-discovered from
|
||||
`fastvideo/eval/metrics/<group>/<name>/metric.py`. Each declares itself
|
||||
with `@register("<group>.<name>")` and subclasses `BaseMetric`. Three
|
||||
recipes:
|
||||
|
||||
1. **Native metric** (pure-PyTorch, no submodule). Drop a file,
|
||||
declare deps, implement `compute(sample)`.
|
||||
2. **Library-wrapped metric** (CLIP, torch.hub, transformers, pyiqa).
|
||||
Same as above, plus route the library's cache through
|
||||
`get_cache_dir()` if it has a `download_root=` / `cache_dir=`
|
||||
kwarg.
|
||||
3. **Upstream-submodule-wrapped metric** (vbench-style). Pin upstream
|
||||
as a git submodule under `fastvideo/third_party/eval/<bench>/`. The
|
||||
adapter `__init__.py` does the `sys.path` insert and any runtime
|
||||
compat shims for modern dep versions. Patches live as Python in
|
||||
that file rather than as on-disk patches to the submodule.
|
||||
|
||||
The full recipes are below.
|
||||
|
||||
---
|
||||
|
||||
## 0) Layout and auto-discovery
|
||||
|
||||
```
|
||||
fastvideo/eval/metrics/
|
||||
├── base.py # BaseMetric + lifecycle contract
|
||||
├── common/ # group: SSIM, PSNR, LPIPS
|
||||
├── optical_flow/ # group: gt_optical_flow, synthetic_optical_flow
|
||||
├── vlm/ # group: VideoScore-2
|
||||
├── physics_iq/ # group + sub-metrics
|
||||
└── vbench/ # group: 16 sub-metrics
|
||||
├── __init__.py # sys.path bootstrap + runtime compat shims
|
||||
├── _grit_helper.py # shared upstream-touching helpers
|
||||
└── <sub_metric>/metric.py
|
||||
```
|
||||
|
||||
Auto-discovery (`fastvideo/eval/metrics/__init__.py`) walks each group
|
||||
dir and imports every `metric.py` it finds, which fires the
|
||||
`@register` decorators. Names starting with `_` are skipped. Use that
|
||||
prefix for shared helpers or vendored code that should not register
|
||||
itself.
|
||||
|
||||
---
|
||||
|
||||
## 1) The `BaseMetric` contract
|
||||
|
||||
Every metric subclasses `fastvideo.eval.metrics.base.BaseMetric` and
|
||||
declares:
|
||||
|
||||
```python
|
||||
class YourMetric(BaseMetric):
|
||||
name: str = "common.your_metric" # must match @register
|
||||
requires_reference: bool = True # needs sample["reference"]
|
||||
higher_is_better: bool = True # for ranking / aggregates
|
||||
dependencies: list[str] = [] # importable module names;
|
||||
# registry surfaces a clean
|
||||
# ImportError if missing
|
||||
needs_gpu: bool = False
|
||||
backbone: str | None = None # e.g. "clip_vit_l14"
|
||||
```
|
||||
|
||||
You must implement:
|
||||
|
||||
```python
|
||||
def compute(self, sample: dict) -> list[MetricResult]:
|
||||
"""sample['video'] is (1, T, C, H, W). Return a one-element list.
|
||||
|
||||
The leading 1 is preserved for forward-compat with batched eval;
|
||||
today :class:`EvalWorker` always invokes metrics with B=1.
|
||||
"""
|
||||
```
|
||||
|
||||
You may override:
|
||||
|
||||
- `setup(self) -> None`. Eager model loading. Called once by
|
||||
`create_evaluator`. Idempotent (re-entrant). Use the `if self._model
|
||||
is not None: return` pattern.
|
||||
- `to(self, device)`. Move the metric and its submodels to `device`.
|
||||
|
||||
If a required input is missing (e.g. an fps-aware metric called
|
||||
without `fps`), return `self._skip(sample, reason)` instead of
|
||||
raising.
|
||||
|
||||
---
|
||||
|
||||
## 2) Recipe A: native metric (no external deps)
|
||||
|
||||
Smallest case. Pixel math, simple closed-form.
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/common/your_metric/metric.py
|
||||
from __future__ import annotations
|
||||
import torch
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult
|
||||
|
||||
|
||||
@register("common.your_metric")
|
||||
class YourMetric(BaseMetric):
|
||||
name = "common.your_metric"
|
||||
requires_reference = True
|
||||
higher_is_better = True
|
||||
needs_gpu = False
|
||||
dependencies: list[str] = [] # nothing extra
|
||||
|
||||
def compute(self, sample: dict) -> list[MetricResult]:
|
||||
gen, ref = sample["video"], sample["reference"] # (B,T,C,H,W) each
|
||||
per_video = ((gen - ref) ** 2).mean(dim=(1, 2, 3, 4)).sqrt()
|
||||
return [
|
||||
MetricResult(name=self.name, score=float(s), details={})
|
||||
for s in per_video
|
||||
]
|
||||
```
|
||||
|
||||
That is the whole recipe. Drop the file and the registry picks it up.
|
||||
|
||||
---
|
||||
|
||||
## 3) Recipe B: library-wrapped metric (CLIP, torch.hub, transformers, pyiqa)
|
||||
|
||||
If your metric loads a backbone from a Python package, route the
|
||||
library at the eval cache so users get one knob (`FASTVIDEO_EVAL_CACHE`)
|
||||
to redirect everything.
|
||||
|
||||
### Cache routing rules
|
||||
|
||||
| Library | How to route | Location after redirect |
|
||||
|---|---|---|
|
||||
| `clip.load("ViT-X")` | pass `download_root=str(get_cache_dir() / "clip")` | `${FASTVIDEO_EVAL_CACHE}/clip/` |
|
||||
| `torch.hub.load(...)` | nothing; `TORCH_HOME` is redirected at `fastvideo.eval` import time | `${FASTVIDEO_EVAL_CACHE}/torch/hub/` |
|
||||
| `transformers.from_pretrained(...)` | nothing; leave HF's default cache (`~/.cache/huggingface/hub/`) so users dedupe with other ML projects | `~/.cache/huggingface/hub/` |
|
||||
| `huggingface_hub.snapshot_download` / `hf_hub_download` | use `ensure_checkpoint(...)` (it wraps these with filelock) | same as above |
|
||||
| `pyiqa.create_metric(...)` | no env var or kwarg honored; document in metric docstring | pyiqa-internal |
|
||||
| `lpips`, `ptlflow` | torch.hub-based, auto-redirected | `${FASTVIDEO_EVAL_CACHE}/torch/hub/` |
|
||||
| Raw URL (no HF Hub) | use `ensure_checkpoint(name, source="https://...")` | `${FASTVIDEO_EVAL_CACHE}/models/<name>` |
|
||||
| Dataset asset (raw video/mask/image) auto-fetched from a public bucket | download into `get_cache_dir() / "datasets" / "<bench>"`, mirroring upstream's relative layout. Vendor any small manifest (CSV/JSON ≤1 MB) under the metric folder so the dataset can be used without external setup. | `${FASTVIDEO_EVAL_CACHE}/datasets/<bench>/` |
|
||||
|
||||
### Dataset assets: vendor the manifest, auto-fetch the rest
|
||||
|
||||
If your metric ships with its own paired-reference dataset (Physics-IQ
|
||||
is the canonical example), follow this layout:
|
||||
|
||||
- **Manifest** (CSV/JSON ≤1 MB): vendor it under
|
||||
`fastvideo/eval/metrics/<bench>/_vendored/<manifest>.<ext>`, with a
|
||||
sibling `_vendored/LICENSE` recording attribution and provenance.
|
||||
The `_vendored/` subdir is the project-wide convention for
|
||||
upstream-provenance files: it is auto-skipped by metric discovery
|
||||
(the `_` prefix) and by codespell (one `*/_vendored/*` glob in
|
||||
`[tool.codespell].skip`), so dropping in a new vendored file
|
||||
requires no further config. Read the manifest from the dataset
|
||||
module via a `Path(__file__)`-relative resolver. Mirror
|
||||
`_VENDORED_DESCRIPTIONS_CSV` in
|
||||
`fastvideo/eval/datasets/physics_iq.py`.
|
||||
- **Heavy assets** (videos, masks, images): do not vendor. Auto-fetch
|
||||
on first miss into `get_cache_dir() / "datasets" / "<bench>"`,
|
||||
mirroring upstream's relative directory layout one-for-one so a
|
||||
pre-downloaded mirror at any path works as a drop-in
|
||||
`dataset_root=`. Use atomic `.part` then final-rename to be safe
|
||||
under concurrent SLURM ranks.
|
||||
- **Bucket override**: expose `FASTVIDEO_<BENCH>_BUCKET_URL` so users
|
||||
with internal mirrors can redirect.
|
||||
- **Opt-out**: accept `auto_download: bool = True` in the dataset
|
||||
constructor; on `False`, raise `FileNotFoundError` instead of
|
||||
fetching. This covers air-gapped runs and CI.
|
||||
|
||||
The end-state is `get_dataset("<bench>")` with no kwargs.
|
||||
|
||||
### Example: CLIP backbone + LAION head
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/your_group/your_metric/metric.py
|
||||
from __future__ import annotations
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult
|
||||
|
||||
|
||||
@register("your_group.your_metric")
|
||||
class YourMetric(BaseMetric):
|
||||
name = "your_group.your_metric"
|
||||
requires_reference = False
|
||||
needs_gpu = True
|
||||
dependencies = ["clip"] # "openai-clip" PyPI; importable as `clip`
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._clip = None
|
||||
self._head = None
|
||||
|
||||
def setup(self) -> None:
|
||||
if self._clip is not None:
|
||||
return
|
||||
import clip
|
||||
from fastvideo.eval.models import ensure_checkpoint, get_cache_dir
|
||||
|
||||
# Backbone: route CLIP's cache through our root.
|
||||
self._clip, _ = clip.load(
|
||||
"ViT-L/14",
|
||||
device=self.device,
|
||||
download_root=str(get_cache_dir() / "clip"),
|
||||
)
|
||||
self._clip.eval()
|
||||
|
||||
# URL-fetched head: ensure_checkpoint downloads to
|
||||
# ${FASTVIDEO_EVAL_CACHE}/models/ with filelock + atomic rename.
|
||||
ckpt = ensure_checkpoint(
|
||||
"your_head.pth",
|
||||
source="https://example.com/path/to/your_head.pth",
|
||||
)
|
||||
self._head = nn.Linear(768, 1)
|
||||
self._head.load_state_dict(
|
||||
torch.load(ckpt, map_location="cpu", weights_only=True)
|
||||
)
|
||||
self._head.to(self.device).eval()
|
||||
|
||||
def to(self, device):
|
||||
super().to(device)
|
||||
if self._clip is not None:
|
||||
self._clip = self._clip.to(self.device)
|
||||
if self._head is not None:
|
||||
self._head = self._head.to(self.device)
|
||||
return self
|
||||
|
||||
def compute(self, sample: dict) -> list[MetricResult]:
|
||||
...
|
||||
```
|
||||
|
||||
### Do not redirect other `~/.cache/...` dirs
|
||||
|
||||
If a third-party library hard-codes `~/.cache/<lib>/` and offers no
|
||||
override, document the exception in the metric's docstring. Forcing
|
||||
redirection by setting `os.environ` or patching `os.path.expanduser`
|
||||
is fragile and breaks user expectations of where the library's cache
|
||||
lives.
|
||||
|
||||
---
|
||||
|
||||
## 4) Recipe C: upstream-submodule-wrapped metric (vbench pattern)
|
||||
|
||||
Use this when the upstream benchmark ships Python code (`vbench/`,
|
||||
`MIND/`, etc.) that is not pip-installable cleanly. See
|
||||
`fastvideo/eval/metrics/vbench/__init__.py` for the worked example.
|
||||
|
||||
### 4.1 Pin the upstream as a submodule
|
||||
|
||||
```bash
|
||||
git submodule add <upstream-url> fastvideo/third_party/eval/<bench>
|
||||
cd fastvideo/third_party/eval/<bench>
|
||||
git checkout <pinned-sha>
|
||||
cd -
|
||||
git add .gitmodules fastvideo/third_party/eval/<bench>
|
||||
```
|
||||
|
||||
The submodule pulls under the standard `git submodule update --init
|
||||
--recursive` flow that users already run for kernel deps.
|
||||
|
||||
### 4.2 Bootstrap on `sys.path`
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/__init__.py
|
||||
from __future__ import annotations
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# fastvideo/eval/metrics/<bench>/__init__.py → ../../../../third_party/eval/<bench>
|
||||
_UPSTREAM = Path(__file__).resolve().parents[3] / "third_party" / "eval" / "<bench>"
|
||||
if _UPSTREAM.is_dir() and str(_UPSTREAM) not in sys.path:
|
||||
sys.path.insert(0, str(_UPSTREAM))
|
||||
```
|
||||
|
||||
We do not `pip install` the upstream because its egg-link/.pth would
|
||||
just re-do this `sys.path.insert`, and skipping the install also skips
|
||||
the upstream's `setup.py` (which often gates on a specific CUDA
|
||||
version).
|
||||
|
||||
### 4.3 Modern-dep compat: runtime shims
|
||||
|
||||
Upstream code pinned to e.g. `transformers==4.33.2`, `numpy<2`
|
||||
typically breaks against modern versions in 3-4 known places (API
|
||||
renames). Fix those at import time, in the same `__init__.py`:
|
||||
|
||||
```python
|
||||
def _install_compat_shims() -> None:
|
||||
# Example: transformers.modeling_utils API moved.
|
||||
try:
|
||||
import transformers.modeling_utils as _mu
|
||||
import transformers.pytorch_utils as _pu
|
||||
for _n in ("apply_chunking_to_forward",
|
||||
"find_pruneable_heads_and_indices",
|
||||
"prune_linear_layer"):
|
||||
if not hasattr(_mu, _n) and hasattr(_pu, _n):
|
||||
setattr(_mu, _n, getattr(_pu, _n))
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Example: numpy.lib.function_base.disp removed in numpy>=2.
|
||||
try:
|
||||
import types, numpy.lib as _nl
|
||||
if not hasattr(_nl, "function_base"):
|
||||
_stub = types.ModuleType("numpy.lib.function_base")
|
||||
_stub.disp = lambda *a, **k: None
|
||||
sys.modules["numpy.lib.function_base"] = _stub
|
||||
_nl.function_base = _stub
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
_install_compat_shims()
|
||||
```
|
||||
|
||||
For function-level patches that cannot be expressed as attribute
|
||||
writes (e.g. wrapping a model factory function), use a
|
||||
`sys.meta_path` finder that wraps the loader. See
|
||||
`_install_modeling_finetune_hook()` in
|
||||
`fastvideo/eval/metrics/vbench/__init__.py` for the pattern (about 30
|
||||
lines).
|
||||
|
||||
Why shims rather than `git apply` patches: patches go stale when the
|
||||
upstream SHA changes; shims are versioned Python code in our repo,
|
||||
they are grep-able, and they only run if the targeted module is
|
||||
imported.
|
||||
|
||||
### 4.4 Per-sub-metric files
|
||||
|
||||
Each sub-metric is a normal `BaseMetric` subclass that imports from
|
||||
the upstream:
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/<sub>/metric.py
|
||||
from fastvideo.eval.metrics.base import BaseMetric
|
||||
from fastvideo.eval.registry import register
|
||||
from fastvideo.eval.types import MetricResult
|
||||
|
||||
|
||||
@register("<bench>.<sub>")
|
||||
class YourSubMetric(BaseMetric):
|
||||
...
|
||||
|
||||
def setup(self) -> None:
|
||||
if self._model is not None:
|
||||
return
|
||||
# The sys.path bootstrap fired when fastvideo.eval.metrics.<bench>
|
||||
# was imported (which auto-discovery does before importing this
|
||||
# sub-package). Upstream imports just work:
|
||||
from <bench>.something import SomeModel
|
||||
...
|
||||
```
|
||||
|
||||
### 4.5 Conditional registration when the upstream is missing
|
||||
|
||||
If a user installed `fastvideo[eval]` but did not run `git submodule
|
||||
update --init`, `<bench>.*` metrics should not register. The
|
||||
auto-discovery walker imports each sub-package's `metric` module; have
|
||||
that import bail out cleanly:
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/__init__.py — at the bottom
|
||||
_AVAILABLE = (_UPSTREAM / "<bench>" / "__init__.py").is_file()
|
||||
```
|
||||
|
||||
```python
|
||||
# fastvideo/eval/metrics/<bench>/<sub>/__init__.py
|
||||
from fastvideo.eval.metrics.<bench> import _AVAILABLE
|
||||
if _AVAILABLE:
|
||||
from .metric import YourSubMetric # noqa
|
||||
```
|
||||
|
||||
`fastvideo eval list` then reflects what the user actually has rather
|
||||
than what they could have.
|
||||
|
||||
### 4.6 Leave the upstream alone unless it blocks the metric
|
||||
|
||||
The upstream is pinned. If a metric works against the pinned SHA,
|
||||
leave the upstream files untouched. If it actively breaks against
|
||||
modern deps (the import-drift cases above), shim it. Avoid
|
||||
fastvideo-side forks of upstream code; they make patches go stale and
|
||||
parity drift.
|
||||
|
||||
---
|
||||
|
||||
## 5) Model checkpoints: `ensure_checkpoint`
|
||||
|
||||
Use `ensure_checkpoint(name, source, filename=None)` for any
|
||||
non-package weights. It resolves a local path, downloading on miss,
|
||||
with filelock safety across processes and SLURM ranks.
|
||||
|
||||
| `source` form | What happens |
|
||||
|---|---|
|
||||
| `"/abs/path/to/file.pth"` | passthrough, returned unchanged |
|
||||
| `"https://..."` | downloaded to `${FASTVIDEO_EVAL_CACHE}/models/<name>` via `huggingface_hub.http_get`, atomic rename, filelock |
|
||||
| `"org/repo"` (no `filename`) | `snapshot_download(repo_id)` → `~/.cache/huggingface/hub/` |
|
||||
| `"org/repo"` (with `filename`) | `hf_hub_download(repo_id, filename)` → `~/.cache/huggingface/hub/` |
|
||||
|
||||
`name` is only used as the local filename for URL sources. HF sources
|
||||
ignore it (HF manages its own cache key by content hash).
|
||||
|
||||
```python
|
||||
from fastvideo.eval.models import ensure_checkpoint
|
||||
|
||||
# URL: name matters
|
||||
ckpt = ensure_checkpoint(
|
||||
"amt-s.pth",
|
||||
source="https://huggingface.co/lalala125/AMT/resolve/main/amt-s.pth",
|
||||
)
|
||||
|
||||
# HF single file: name is decorative
|
||||
ckpt = ensure_checkpoint(
|
||||
"raft-things.pth", # ignored; HF cache uses repo+sha
|
||||
source="OpenGVLab/VBench_Used_Models",
|
||||
filename="raft-things.pth",
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 6) Declaring `dependencies`
|
||||
|
||||
Set `dependencies = ["pkg1", "pkg2"]` on your metric class with
|
||||
importable module names (not PyPI distribution names). The registry
|
||||
checks each via `importlib.util.find_spec` at instantiation time and
|
||||
raises a clean `ImportError` pointing the user at the right install
|
||||
extra:
|
||||
|
||||
```python
|
||||
class YourMetric(BaseMetric):
|
||||
dependencies = ["clip", "timm"] # importable as `import clip`, `import timm`
|
||||
```
|
||||
|
||||
If a dep is in `[project.optional-dependencies.eval-<group>]`, you do
|
||||
not need to do anything more. If it is a new dep, add it to that
|
||||
group in `pyproject.toml`.
|
||||
|
||||
---
|
||||
|
||||
## 7) Common gotchas
|
||||
|
||||
- The standard `git submodule update --init --recursive` is enough
|
||||
for a benchmark; do not write a `setup.sh`. Modern-dep compat goes
|
||||
into your `__init__.py` as runtime shims.
|
||||
- Do not modify upstream files on disk. The submodule should always
|
||||
match its pinned SHA. Compat lives in our `__init__.py`.
|
||||
- Do not pip-install the upstream. The egg-link is a glorified
|
||||
`sys.path.insert`, which we do directly in `__init__.py`.
|
||||
- Do not call `torch.hub.set_dir(...)` from your metric. It is done
|
||||
globally in `fastvideo/eval/__init__.py`.
|
||||
- Do not put cache-redirection env vars in your metric's `setup()`.
|
||||
By the time `setup()` runs, the library has likely already cached
|
||||
the default-location decision. Set env vars at package-init time.
|
||||
- Skip rather than raise when an input is missing. Use
|
||||
`self._skip(sample, reason)` for any expected-missing input. It
|
||||
returns a list of `MetricResult(score=None)` so other metrics in
|
||||
the same evaluator continue.
|
||||
- Watch for upstream re-registration conflicts. If the upstream uses
|
||||
a global registry (detectron2's `META_ARCH_REGISTRY`, MMCV, etc.),
|
||||
loading the same model twice in the same process will throw. The
|
||||
evaluator already loads each metric once; if you write a custom
|
||||
setup-then-call pattern, mirror that single-load discipline.
|
||||
|
||||
---
|
||||
|
||||
## 8) Training-time eval: keep evaluators hot, free caches between calls
|
||||
|
||||
When wiring eval into a training loop, the working pattern is:
|
||||
|
||||
1. Construct the `Evaluator` once and attach it to the pipeline
|
||||
(`self._eval = create_evaluator(...)`). Do not recreate it per
|
||||
validation round; that re-pays the model load cost.
|
||||
2. Save validation videos to disk (the diffusion path already does
|
||||
this). Pass paths to `evaluator.evaluate`, not in-memory tensors
|
||||
that share GPU memory with the training model.
|
||||
3. Run validation only on rank 0 of each sequence-parallel group.
|
||||
Gather paths from other ranks and let rank 0 score everything.
|
||||
4. After every `evaluate(...)` call, call
|
||||
`evaluator.release_cuda_memory()` in a `finally` block. That runs
|
||||
`gc.collect()` + `torch.cuda.empty_cache()` +
|
||||
`torch.cuda.ipc_collect()`. The eval model stays loaded; only
|
||||
transient activation buffers from the just-finished call get
|
||||
freed:
|
||||
|
||||
```python
|
||||
for video_path in batch:
|
||||
try:
|
||||
scores = self._eval.evaluate(video=load_video(video_path))
|
||||
finally:
|
||||
self._eval.release_cuda_memory()
|
||||
```
|
||||
|
||||
5. If memory pressure spikes (rare on H200), call
|
||||
`evaluator.unload()` to drop every metric reference and let the
|
||||
GPU memory be GC'd. `unload` is reversible:
|
||||
`evaluator.reload()` rebuilds the same metrics with the original
|
||||
config (re-paying the model load cost). Calling `evaluate`
|
||||
between `unload` and `reload` raises a clear `RuntimeError`.
|
||||
|
||||
For most metrics (sub-1 GB backbones, e.g. CLIP/DINO/RAFT/AMT) the
|
||||
eval model can stay co-resident with the training model in
|
||||
`transformer.eval()` mode without any swap. For larger ones
|
||||
(VideoScore2 at 14 GB), measure first; if it fits on the rank-0 GPU
|
||||
during validation (training model in eval mode means no
|
||||
grads/optimizer updates), keep it hot. If not, `unload` between
|
||||
rounds.
|
||||
|
||||
## 9) Local verification
|
||||
|
||||
Native and library-wrapped metrics: a single-GPU smoke is enough.
|
||||
|
||||
```python
|
||||
import torch
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
ev = create_evaluator(metrics=["<group>.<your_metric>"], device="cuda")
|
||||
video = torch.randn(1, 49, 3, 256, 256, device="cuda").clamp(0, 1)
|
||||
print(ev.evaluate(video=video))
|
||||
```
|
||||
|
||||
Submodule-wrapped metrics: also do a parity check against the
|
||||
upstream once. Clone upstream into a separate venv, run the same
|
||||
video through both, and expect an exact match on bit-deterministic
|
||||
metrics and ≤1% drift on backbone-heavy ones (driven by
|
||||
transformers/torch version differences).
|
||||
|
||||
For quick parity in CI: pin a tiny test video, record expected
|
||||
scores ± tolerance, and add a calibration test under
|
||||
`fastvideo/tests/eval/`.
|
||||
|
||||
---
|
||||
|
||||
## 10) When not to add a metric
|
||||
|
||||
- **Set-vs-set distribution metrics** (FVD, FID-style) do not fit
|
||||
`BaseMetric.compute(sample)` cleanly; they need a population.
|
||||
Adding them requires a stateful accumulator interface that does
|
||||
not exist yet. Open an issue first.
|
||||
- **Metrics requiring a single-GPU model larger than available
|
||||
memory.** Eval is not the place for tensor-parallel sharding;
|
||||
metrics are expected to fit on one GPU.
|
||||
- **Metrics that need `mmcv` with a conflicting CUDA ABI.** Document
|
||||
the affected sub-metrics as unsupported and skip them. Building
|
||||
isolation infrastructure (subprocess engine, per-metric venv) is
|
||||
out of scope.
|
||||
@@ -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,30 +306,6 @@ 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."
|
||||
@@ -404,13 +380,6 @@ 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
|
||||
|
||||
uv pip install fastvideo
|
||||
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 (still drives installs through `uv`):
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
pip install -e .
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
@@ -58,16 +58,14 @@ 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
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install FlashAttention:
|
||||
|
||||
```bash
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -89,7 +87,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
@@ -103,7 +101,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
|
||||
@@ -57,10 +57,8 @@ 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
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -82,7 +80,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
- Install MoGe:
|
||||
|
||||
```bash
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
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
|
||||
uv pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
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
|
||||
uv pip install ninja
|
||||
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 uv pip install -e .
|
||||
python setup.py install # or 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
|
||||
uv pip install vsa
|
||||
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
|
||||
uv pip install vsa
|
||||
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
|
||||
uv pip install vsa
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -7,7 +7,7 @@ and the GEN3C diffusion model.
|
||||
|
||||
Requirements:
|
||||
1. Install MoGe:
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
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:
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,53 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,34 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,45 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,34 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,37 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,34 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,34 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,77 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — text-to-audio (baseline) example.
|
||||
|
||||
User story (game-audio designer, prototyping):
|
||||
"I'm prototyping a level and I need 6 seconds of background
|
||||
ambience — gentle wind, distant thunder, a hint of birdsong. I
|
||||
don't want to dig through a sound library; I want to type what I
|
||||
hear in my head and get a wav back. If it's wrong I'll iterate
|
||||
on the prompt. This is the first stop."
|
||||
|
||||
User story (musician sketching ideas):
|
||||
"I want to bounce a 30s lo-fi drum loop to use as a placeholder
|
||||
bed while I build the rest of the track. Type prompt, get audio,
|
||||
drop into the DAW. The actual production beat I'll record
|
||||
myself, but I need *something* to write the chords against."
|
||||
|
||||
User story (researcher exploring the model):
|
||||
"First time touching Stable Audio Open — what does it sound
|
||||
like at default settings? This is the smallest amount of code
|
||||
that goes from prompt to mp4."
|
||||
|
||||
How it works:
|
||||
Pure text-to-audio (T2A). The pipeline runs:
|
||||
T5 + NumberConditioner -> StableAudioDiT -> Oobleck VAE
|
||||
via the `dpmpp-3m-sde` k-diffusion sampler. All components are
|
||||
FastVideo-native — no diffusers / transformers model imports at
|
||||
runtime (see REVIEW item 30). Mirrors upstream
|
||||
`stable_audio_tools.inference.generation.generate_diffusion_cond`
|
||||
bit-for-bit (~0.2% abs_mean drift on 25 steps).
|
||||
|
||||
Tunable knobs (the "creative dials"):
|
||||
audio_end_in_s
|
||||
1–6 — quick ideation (sub-10s wall clock at 100 steps)
|
||||
10–30 — full musical phrase / loop length (the README example
|
||||
uses 30s)
|
||||
47.5 — model maximum (full sample_size = 2097152 / 44100 Hz)
|
||||
num_inference_steps
|
||||
25 — fast preview, occasional artifacts
|
||||
100 — preset default (matches the HF model card)
|
||||
250 — diminishing returns past here
|
||||
guidance_scale
|
||||
3 — looser, more variation per seed
|
||||
7 — preset default; matches README
|
||||
12+ — sharper but can sound "fried"
|
||||
|
||||
Prerequisites:
|
||||
1. Accept the terms on https://huggingface.co/stabilityai/stable-audio-open-1.0
|
||||
and export your HF token in the shell:
|
||||
export HF_TOKEN=hf_...
|
||||
2. Install optional inference deps (one-time):
|
||||
uv pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_audio/stable_audio_basic/output_stable_audio.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# 6-second clip; the model max is ~47.5s.
|
||||
audio_end_in_s=6.0,
|
||||
# The registered preset gives 100 steps + CFG=7.0 by default;
|
||||
# override num_inference_steps / guidance_scale here for QA.
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,77 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — audio-to-audio variation example.
|
||||
|
||||
User story (musician, late at night):
|
||||
"I generated this 12-second lo-fi loop earlier and I love the chord
|
||||
progression and overall vibe, but the snare hit at 0:08 sounds wrong
|
||||
and the rhythm feels stiff. I don't want to start over from scratch
|
||||
and lose what's working — I want the model to keep the harmony and
|
||||
mood but reroll the percussion + groove."
|
||||
|
||||
User story (sound designer, on a deadline):
|
||||
"I have one good 'sword clang' SFX. The art director wants 8 sibling
|
||||
variations that all feel like the same sword from different angles —
|
||||
same metal, same weight, slightly different impact. I'd rather
|
||||
refine my one good take than text-prompt my way through 50 misses."
|
||||
|
||||
Pass `init_audio=path/to/clip` (any wav/mp3/mp4/m4a/flac the standard
|
||||
deps decode) and the model will use it as a starting point for the
|
||||
text prompt instead of pure noise.
|
||||
|
||||
Picking `init_audio_strength` (0.0 to 1.0):
|
||||
|
||||
Higher = closer to the source clip. Lower = more transformation.
|
||||
(Same convention as the "Input Audio Strength" slider in
|
||||
Stability's commercial Stable Audio web UI, so values transfer
|
||||
directly.)
|
||||
|
||||
| strength | what you get |
|
||||
|----------|----------------------------------------------------|
|
||||
| 1.00 | Output ≈ reference. No transformation. |
|
||||
| 0.85 | Texture micro-variation only. |
|
||||
| 0.70 | Light reroll, same instruments. |
|
||||
| 0.60 | Default. Instrument identity is replaceable |
|
||||
| | (cello can take over from piano on the same notes).|
|
||||
| 0.50 | Heavy — only melody / chord progression survives. |
|
||||
| 0.30 | Reference acts as a loose mood prompt. |
|
||||
| 0.00 | Plain T2A — reference ignored. |
|
||||
|
||||
Rule of thumb by intent:
|
||||
* "Fix one part of this clip" -> 0.75 .. 0.85
|
||||
* "Same notes, different instrument" -> 0.55 .. 0.65
|
||||
* "Same chord progression, new content" -> 0.40 .. 0.55
|
||||
* "Use this as a loose mood prompt" -> 0.20 .. 0.35
|
||||
|
||||
If the reference timbre is bleeding through more than you want,
|
||||
lower it; if the structure is gone, raise it.
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Change the piano to a cello playing the same notes"
|
||||
# Path to any audio-bearing file (wav, mp3, mp4, m4a, flac, ...).
|
||||
# Set to `None` to skip A2A and run plain T2A.
|
||||
INIT_AUDIO_PATH: str | None = None
|
||||
# Reference fidelity in [0, 1] -- higher = closer to source.
|
||||
INIT_AUDIO_STRENGTH = 0.6
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_a2a/output_a2a.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=6.0,
|
||||
init_audio=INIT_AUDIO_PATH,
|
||||
init_audio_strength=INIT_AUDIO_STRENGTH,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,84 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open 1.0 — inpainting / outpainting (loop extension) example.
|
||||
|
||||
User story (loop extension — the killer app):
|
||||
"I have a 6-second drum loop my client likes. They want it as
|
||||
background bed for a 30-second ad. I need it to loop seamlessly,
|
||||
but a hard cut every 6s sounds bad. Let me extend it to 30s,
|
||||
keeping the first 6s exactly as-is and letting the model continue
|
||||
the groove for the remaining 24s."
|
||||
|
||||
User story (audio repair):
|
||||
"There's a microphone bump at 0:14 in this 30-second field
|
||||
recording — really obvious in headphones. Mask out 0:13 to 0:15
|
||||
and let the model regenerate plausible ambience that blends in.
|
||||
Everything else stays exactly as I recorded it."
|
||||
|
||||
User story (transition smoothing):
|
||||
"I have two 10-second clips I want to crossfade. Mask out a 1s
|
||||
overlap region in the middle and let the model invent a coherent
|
||||
transition between the two."
|
||||
|
||||
How it works (RePaint-style blending):
|
||||
Stable Audio Open 1.0 wasn't trained as an inpainting model
|
||||
(`model_type=diffusion_cond`, not `diffusion_cond_inpaint`), so we
|
||||
can't use the upstream's mask-conditioned approach directly. We
|
||||
use the RePaint trick instead, which works on any v-prediction
|
||||
diffusion model:
|
||||
|
||||
1. Encode the reference clip into latent space.
|
||||
2. At every denoising step `i`, replace the kept region of the
|
||||
in-flight latent (where mask == 1) with the reference
|
||||
re-noised to the next timestep's sigma. Only the unkept
|
||||
region (mask == 0) is freely denoised.
|
||||
3. After the loop, the kept region is exactly the reference;
|
||||
the unkept region is freshly generated content.
|
||||
|
||||
This is approximate compared to a properly trained inpainting
|
||||
checkpoint — the seam between kept/unkept can have slight EQ
|
||||
discontinuity — but it works on the existing public model.
|
||||
|
||||
Tunable: the mask is a 1-D tensor in {0, 1} at the model's sample
|
||||
rate. Conventions:
|
||||
1.0 = keep this sample from the reference
|
||||
0.0 = regenerate this sample
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`.
|
||||
"""
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Steady lo-fi hip hop drum loop with vinyl crackle."
|
||||
# Required: path to the reference audio file (wav, mp3, mp4, m4a, flac,
|
||||
# ...) you want to extend or repair. The pipeline raises if a mask is
|
||||
# passed without a reference, so this must be a real path.
|
||||
REFERENCE_AUDIO_PATH = "path/to/your/loop.wav"
|
||||
KEEP_SECONDS = 6.0 # first KEEP_SECONDS preserved exactly
|
||||
TOTAL_SECONDS = 12.0 # extend the loop to this duration
|
||||
|
||||
|
||||
def main() -> None:
|
||||
if not os.path.isfile(REFERENCE_AUDIO_PATH):
|
||||
raise FileNotFoundError(
|
||||
f"REFERENCE_AUDIO_PATH={REFERENCE_AUDIO_PATH!r} does not exist. "
|
||||
"Edit this script to point at a real audio file (wav/mp3/mp4/"
|
||||
"m4a/flac) before running.")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path="outputs_audio/stable_audio_inpaint/output_inpaint.wav",
|
||||
save_video=True,
|
||||
audio_end_in_s=TOTAL_SECONDS,
|
||||
inpaint_audio=REFERENCE_AUDIO_PATH,
|
||||
# Tuple form: keep first KEEP_SECONDS, regenerate the rest.
|
||||
inpaint_mask=(KEEP_SECONDS, TOTAL_SECONDS),
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,53 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Stable Audio Open Small — fast / lightweight T2A example.
|
||||
|
||||
User story (interactive UI builder):
|
||||
"I'm building a sound-design UI where the user types a prompt and
|
||||
we want sub-2-second feedback so the experience feels like
|
||||
autocomplete, not a render queue. The full Stable Audio Open 1.0
|
||||
takes ~8s on a single GPU; the small variant takes a fraction of
|
||||
that — quality is lower but completely usable for real-time
|
||||
iteration."
|
||||
|
||||
User story (overnight batch jobs):
|
||||
"I'm generating 10,000 short SFX variants for a procedural game.
|
||||
Wall-clock matters more than per-clip polish — give me the small
|
||||
model so I can fit the run in one night instead of a week."
|
||||
|
||||
How it works:
|
||||
The small variant is a separate Stability AI checkpoint
|
||||
(`stabilityai/stable-audio-open-small`) that ships the same Oobleck
|
||||
VAE as the 1.0 base model but a smaller / faster DiT (`embed_dim=1024`,
|
||||
`depth=16`, `qk_norm="ln"`) and only one duration conditioner
|
||||
(`seconds_total`, no `seconds_start`). FastVideo loads from the
|
||||
converted Diffusers-format repo `FastVideo/stable-audio-open-small-Diffusers`
|
||||
via the standard component loader; per-variant arch fields come
|
||||
from `transformer/config.json` and `conditioner/config.json`.
|
||||
|
||||
Prerequisites: same as `basic_stable_audio.py`. The converted repo is
|
||||
public so no gated-access flow is required.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
PROMPT = "Lo-fi hip hop instrumental with vinyl crackle and gentle piano."
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/stable-audio-open-small-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
output_path = "outputs_audio/stable_audio_small/output_stable_audio_small.wav"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
# Small variant trains on a ~11.9s window — keep `audio_end_in_s`
|
||||
# at or below that.
|
||||
audio_end_in_s=6.0,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,100 @@
|
||||
"""Generate one LTX2 video and score it with VBench metrics.
|
||||
|
||||
The generation block is the same as
|
||||
``examples/inference/basic/basic_ltx2.py`` — same prompt, same model,
|
||||
same shape, same num_frames. After ``shutdown()`` the script loads the
|
||||
mp4 back, builds a single :class:`fastvideo.eval.Evaluator`, and runs
|
||||
the prompt-aware VBench subset that's meaningful for an arbitrary
|
||||
text→video sample.
|
||||
|
||||
The first run downloads CLIP / DINO / RAFT / AMT / ViCLIP / MUSIQ
|
||||
weights to ``~/.cache/fastvideo/eval/`` (~few GB total).
|
||||
|
||||
GPU memory caveat
|
||||
-----------------
|
||||
Scoring 1088×1920×121 with all 8 metrics needs a dedicated GPU (~80 GB).
|
||||
On a shared GPU, ``vbench.motion_smoothness`` (AMT correlation volume)
|
||||
will OOM — its memory autoscale reads ``total_memory`` rather than
|
||||
``mem_get_info()`` free memory and therefore underestimates the
|
||||
required scale-down. Drop ``motion_smoothness`` from ``METRICS`` if
|
||||
sharing, or run on a smaller-resolution generation.
|
||||
"""
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import Evaluator
|
||||
from fastvideo.eval.io import build_eval_kwargs
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
# VBench sub-metrics meaningful for an arbitrary text→video sample
|
||||
# (just the generated frames, optionally fps + the source prompt).
|
||||
# Structured-prompt metrics (vbench.color, vbench.multiple_objects,
|
||||
# vbench.scene, ...) are excluded — they need prompts built to a
|
||||
# specific schema.
|
||||
METRICS = [
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.background_consistency", # DINO on background patches
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
]
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# ----- generation (matches examples/inference/basic/basic_ltx2.py) -----
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
)
|
||||
generator.shutdown()
|
||||
# Free residual CUDA memory the generator left behind so the
|
||||
# evaluator can grab the largest possible workspace for AMT/RAFT.
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# ----- scoring -----
|
||||
print(f"\n[eval] building evaluator: {METRICS}")
|
||||
evaluator = Evaluator(metrics=METRICS)
|
||||
|
||||
# LTX2 outputs at 24 fps by default.
|
||||
sample = build_eval_kwargs({"prompt": PROMPT}, output_path, fps=24.0)
|
||||
print(f"[eval] running ({sample['video'].shape[1]} frames @ 24 fps)...")
|
||||
results = evaluator.evaluate(**sample)
|
||||
|
||||
print("\n=== VBench scores ===")
|
||||
for name in METRICS:
|
||||
r = results[name]
|
||||
if r.score is None:
|
||||
reason = r.details.get("skipped", "no score")
|
||||
print(f" {name}: SKIPPED ({reason})")
|
||||
else:
|
||||
print(f" {name}: {r.score:.4f}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,157 @@
|
||||
"""End-to-end Physics-IQ: dataset → generate → score → aggregate.
|
||||
|
||||
Generates one video per take-1 scenario with LTX2 (using the scenario
|
||||
caption as the prompt), scores each generated video against the take-1
|
||||
reference and the take-2 "physical-variance" reference, and prints
|
||||
aggregate scores using :meth:`PhysicsIQMetric.aggregate_components` —
|
||||
the official scoring recipe from the upstream benchmark.
|
||||
|
||||
Reference videos / masks / switch-frames auto-fetch on first miss into
|
||||
``${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/``; pass ``--dataset-root``
|
||||
to point at a pre-downloaded mirror instead.
|
||||
|
||||
Quick smoke run on 4 scenarios across 2 GPUs::
|
||||
|
||||
python examples/inference/eval/bench_physics_iq.py \\
|
||||
--limit 4 --num-gpus 2 \\
|
||||
--videos-dir outputs_video/physics_iq_smoke
|
||||
|
||||
Re-score existing generations without regenerating::
|
||||
|
||||
python examples/inference/eval/bench_physics_iq.py \\
|
||||
--videos-dir outputs_video/physics_iq_smoke \\
|
||||
--skip-generation
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator, get_metric
|
||||
from fastvideo.eval.datasets import get_dataset
|
||||
|
||||
|
||||
def _expected_filename(row: dict) -> str:
|
||||
"""Filename Physics-IQ expects for the generated video for *row*.
|
||||
|
||||
Uses the dataset's own ``expected_gen_filename`` annotation so the
|
||||
output filenames match the benchmark's manifest convention.
|
||||
"""
|
||||
return row["auxiliary_info"]["expected_gen_filename"]
|
||||
|
||||
|
||||
def _generate_videos(rows: list[dict], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
todo = [(row, videos_dir / _expected_filename(row)) for row in rows]
|
||||
todo = [(row, out) for (row, out) in todo if not out.is_file()]
|
||||
if not todo:
|
||||
print(f"[gen] all {len(rows)} videos already present; skipping.")
|
||||
return
|
||||
|
||||
print(f"[gen] {len(todo)}/{len(rows)} scenarios to render with {model} "
|
||||
f"({num_frames}x{height}x{width})...")
|
||||
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
|
||||
try:
|
||||
for row, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=row["prompt"], output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
)
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dataset-root", type=Path, default=None,
|
||||
help="Path to a pre-downloaded Physics-IQ release. "
|
||||
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
|
||||
"auto-fetching missing assets from the public bucket.")
|
||||
p.add_argument("--videos-dir", type=Path,
|
||||
default=Path("outputs_video/bench_physics_iq"),
|
||||
help="Where to read/write generated videos.")
|
||||
p.add_argument("--limit", type=int, default=None,
|
||||
help="Truncate to first N scenarios for smoke runs.")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the text→video generator to use.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Re-score existing videos under --videos-dir.")
|
||||
p.add_argument("--scores-out", type=Path, default=None,
|
||||
help="Where to write per-scenario scores (JSON). "
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
args = p.parse_args()
|
||||
|
||||
# 1. Walk the Physics-IQ corpus. Pass --limit to the dataset
|
||||
# constructor so auto-download only fetches the assets we'll use.
|
||||
ds = get_dataset("physics_iq", dataset_root=args.dataset_root, limit=args.limit)
|
||||
rows = list(ds)
|
||||
print(f"[load] Physics-IQ: {len(rows)} scenarios from {ds.dataset_dir}")
|
||||
|
||||
# 2. Generate (or reuse) one mp4 per scenario.
|
||||
if not args.skip_generation:
|
||||
_generate_videos(
|
||||
rows, args.videos_dir, args.model, args.num_gpus,
|
||||
args.num_frames, args.height, args.width,
|
||||
)
|
||||
|
||||
# 3. Score each scenario. The metric reads file paths directly out
|
||||
# of the row dict (reference, reference_take2, masks), so we
|
||||
# just attach the generated video path and forward.
|
||||
evaluator = create_evaluator(metrics=["physics_iq"], num_gpus=args.num_gpus)
|
||||
|
||||
samples: list[dict] = []
|
||||
matched: list[dict] = []
|
||||
for row in rows:
|
||||
video_path = args.videos_dir / _expected_filename(row)
|
||||
if not video_path.is_file():
|
||||
print(f"[eval] missing {video_path}; skipping.")
|
||||
continue
|
||||
# The physics_iq metric accepts file paths via its polymorphic
|
||||
# input handling — no need to load the tensors here.
|
||||
samples.append({"video": str(video_path), **row})
|
||||
matched.append(row)
|
||||
|
||||
all_results = evaluator.evaluate(samples=samples)
|
||||
evaluator.shutdown()
|
||||
|
||||
# 4. Aggregate per the upstream scoring recipe.
|
||||
metric = get_metric("physics_iq")
|
||||
components = metric.aggregate_components(
|
||||
[r["physics_iq"] for r in all_results]
|
||||
)
|
||||
|
||||
print()
|
||||
print("=== Physics-IQ aggregate ===")
|
||||
for name, value in components.items():
|
||||
print(f" {name:24s} {value:.4f}")
|
||||
|
||||
detailed = [
|
||||
{
|
||||
"scenario": row["auxiliary_info"]["scenario_id"],
|
||||
"view": row["view"],
|
||||
"scenario_name": row["auxiliary_info"]["scenario_name"],
|
||||
"score": results["physics_iq"].score,
|
||||
}
|
||||
for row, results in zip(matched, all_results)
|
||||
]
|
||||
out = args.scores_out or (args.videos_dir / "scores.json")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(
|
||||
{"aggregate": components, "per_scenario": detailed},
|
||||
indent=2,
|
||||
))
|
||||
print(f"\n[done] per-scenario scores → {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,155 @@
|
||||
"""End-to-end VBench: dataset → generate → score → aggregate.
|
||||
|
||||
Iterates the VBench prompt corpus, generates one video per prompt with
|
||||
LTX2, scores each generated video against the requested ``vbench.*``
|
||||
sub-metrics, and prints per-metric averages over the run.
|
||||
|
||||
Re-running with ``--skip-generation`` reuses any mp4 already on disk
|
||||
under ``--videos-dir``, so you can iterate on metric selection without
|
||||
re-paying the generation cost.
|
||||
|
||||
Example — quick smoke run on 4 prompts from the ``aesthetic_quality``
|
||||
dimension across 2 GPUs::
|
||||
|
||||
python examples/inference/eval/bench_vbench.py \\
|
||||
--dimensions aesthetic_quality \\
|
||||
--limit 4 --num-gpus 2 \\
|
||||
--videos-dir outputs_video/vbench_smoke
|
||||
|
||||
Full benchmark on a single dimension::
|
||||
|
||||
python examples/inference/eval/bench_vbench.py \\
|
||||
--dimensions subject_consistency --num-gpus 8
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.datasets import get_dataset
|
||||
|
||||
|
||||
def _slugify(prompt: str, max_len: int = 100) -> str:
|
||||
"""Filesystem-safe filename stem; mirrors VBench's official convention."""
|
||||
s = re.sub(r'[\\/:*?"<>|]', "", prompt[:max_len]).strip().strip(".")
|
||||
return re.sub(r"\s+", " ", s) or "output"
|
||||
|
||||
|
||||
def _generate_videos(prompts: list[str], videos_dir: Path,
|
||||
model: str, num_gpus: int,
|
||||
num_frames: int, height: int, width: int) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
videos_dir.mkdir(parents=True, exist_ok=True)
|
||||
todo = [(p, videos_dir / f"{_slugify(p)}.mp4") for p in prompts]
|
||||
todo = [(p, out) for (p, out) in todo if not out.is_file()]
|
||||
if not todo:
|
||||
print(f"[gen] all {len(prompts)} videos already present; skipping.")
|
||||
return
|
||||
|
||||
print(f"[gen] {len(todo)}/{len(prompts)} prompts to render with {model} "
|
||||
f"({num_frames}x{height}x{width})...")
|
||||
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
|
||||
try:
|
||||
for prompt, out_path in todo:
|
||||
gen.generate_video(
|
||||
prompt=prompt, output_path=str(out_path), save_video=True,
|
||||
num_frames=num_frames, height=height, width=width,
|
||||
)
|
||||
finally:
|
||||
gen.shutdown()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--dimensions", default="aesthetic_quality,subject_consistency",
|
||||
help="Comma-separated VBench dimensions (or 'all').")
|
||||
p.add_argument("--limit", type=int, default=None,
|
||||
help="Truncate to first N prompts for smoke runs.")
|
||||
p.add_argument("--videos-dir", type=Path,
|
||||
default=Path("outputs_video/bench_vbench"))
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the text→video generator to use.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--fps", type=float, default=24.0,
|
||||
help="Frame-rate annotation passed to fps-aware metrics.")
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Re-score existing videos under --videos-dir without "
|
||||
"regenerating.")
|
||||
p.add_argument("--scores-out", type=Path, default=None,
|
||||
help="Where to dump per-prompt scores as JSON. "
|
||||
"Defaults to <videos-dir>/scores.json.")
|
||||
args = p.parse_args()
|
||||
|
||||
# 1. Pull prompts from VBench.
|
||||
dims_arg: list[str] | str = (
|
||||
args.dimensions if args.dimensions == "all"
|
||||
else [d.strip() for d in args.dimensions.split(",") if d.strip()]
|
||||
)
|
||||
ds = get_dataset("vbench", dimensions=dims_arg)
|
||||
rows = list(ds)[: args.limit]
|
||||
print(f"[load] VBench: {len(rows)} prompts across {ds.dimensions}")
|
||||
|
||||
# 2. Generate (or reuse) one mp4 per prompt.
|
||||
if not args.skip_generation:
|
||||
_generate_videos(
|
||||
[row["prompt"] for row in rows],
|
||||
args.videos_dir, args.model, args.num_gpus,
|
||||
args.num_frames, args.height, args.width,
|
||||
)
|
||||
|
||||
# 3. Score each video against the requested vbench sub-metrics.
|
||||
metric_names = sorted(set(f"vbench.{d}" for d in ds.dimensions))
|
||||
print(f"[eval] metrics: {metric_names}")
|
||||
evaluator = create_evaluator(metrics=metric_names, num_gpus=args.num_gpus)
|
||||
|
||||
samples: list[dict] = []
|
||||
matched_rows: list[dict] = []
|
||||
for row in rows:
|
||||
video_path = args.videos_dir / f"{_slugify(row['prompt'])}.mp4"
|
||||
if not video_path.is_file():
|
||||
print(f"[eval] missing {video_path}; skipping this row.")
|
||||
continue
|
||||
# Pass the path; the worker decodes lazily so memory stays bounded.
|
||||
samples.append({
|
||||
"video": str(video_path),
|
||||
"fps": args.fps,
|
||||
**row, # prompt / aux / dims
|
||||
})
|
||||
matched_rows.append(row)
|
||||
|
||||
all_results = evaluator.evaluate(samples=samples)
|
||||
evaluator.shutdown()
|
||||
|
||||
# 4. Aggregate per-metric.
|
||||
by_metric: dict[str, list[float]] = defaultdict(list)
|
||||
detailed: list[dict] = []
|
||||
for row, results in zip(matched_rows, all_results):
|
||||
scores = {name: r.score for name, r in results.items()}
|
||||
detailed.append({"prompt": row["prompt"], "scores": scores})
|
||||
for name, score in scores.items():
|
||||
if score is not None:
|
||||
by_metric[name].append(score)
|
||||
|
||||
print()
|
||||
print("=== per-metric averages ===")
|
||||
for name in sorted(by_metric):
|
||||
avg = sum(by_metric[name]) / len(by_metric[name])
|
||||
print(f" {name:42s} {avg:.4f} (n={len(by_metric[name])})")
|
||||
|
||||
out = args.scores_out or (args.videos_dir / "scores.json")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(detailed, indent=2))
|
||||
print(f"\n[done] per-prompt scores → {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,167 @@
|
||||
"""End-to-end: generate a video with LTX2 and score it with VBench.
|
||||
|
||||
Pipeline:
|
||||
prompt → LTX2-Base → mp4 → fastvideo.eval → vbench scores
|
||||
|
||||
Run::
|
||||
|
||||
pip install -e .[eval]
|
||||
git submodule update --init fastvideo/third_party/eval/vbench
|
||||
|
||||
python examples/inference/eval/eval_ltx2_vbench.py
|
||||
# or with 4 GPUs and the distilled checkpoint:
|
||||
python examples/inference/eval/eval_ltx2_vbench.py \
|
||||
--model FastVideo/LTX2-Distilled-Diffusers --num-gpus 4
|
||||
|
||||
The default metric set covers the vbench sub-metrics that are
|
||||
meaningful for an arbitrary text→video sample — i.e. those that need
|
||||
only the generated video (and optionally fps + the source prompt).
|
||||
Structured-prompt metrics like ``vbench.color``, ``vbench.scene``,
|
||||
``vbench.multiple_objects`` etc. are *not* on by default — they only
|
||||
make sense when the prompt is built to a specific schema, and they
|
||||
require GRiT/detectron2 setup. Pass them via ``--metrics`` if you have
|
||||
a matching prompt.
|
||||
|
||||
First-time runs download CLIP, DINO, RAFT, AMT, ViCLIP, and MUSIQ
|
||||
weights to ``~/.cache/fastvideo/eval/models/`` and
|
||||
``~/.cache/torch/hub/`` (~few GB total). Subsequent runs are fast.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
|
||||
PROMPT = (
|
||||
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
|
||||
"of a woman and a man in their 30s, facing each other with serious "
|
||||
"expressions. The woman, emotional and dramatic, says softly, \"That's "
|
||||
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
|
||||
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
|
||||
"then mutters defensively, \"He's just having fun.\" The camera slowly "
|
||||
"pans right, revealing the grandfather in the garden wearing enormous "
|
||||
"butterfly wings, waving his arms in the air like he's trying to take "
|
||||
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
|
||||
"The woman covers her face, on the verge of tears. The tone is deadpan, "
|
||||
"absurd, and quietly tragic."
|
||||
)
|
||||
|
||||
DEFAULT_METRICS = [
|
||||
# No-input metrics: just need the generated frames.
|
||||
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
|
||||
"vbench.subject_consistency", # DINO frame-to-first cosine
|
||||
"vbench.background_consistency", # DINO on background patches
|
||||
"vbench.imaging_quality", # pyiqa MUSIQ
|
||||
"vbench.temporal_flickering", # pixel-wise frame deltas
|
||||
"vbench.motion_smoothness", # AMT frame interpolator residual
|
||||
# Need fps annotation:
|
||||
"vbench.dynamic_degree", # RAFT optical-flow magnitude
|
||||
# Need the source prompt:
|
||||
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
|
||||
]
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
|
||||
help="HF repo id of the LTX2 checkpoint.")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4",
|
||||
help="Where to save the generated mp4.")
|
||||
p.add_argument("--num-frames", type=int, default=121)
|
||||
p.add_argument("--height", type=int, default=1088)
|
||||
p.add_argument("--width", type=int, default=1920)
|
||||
p.add_argument("--prompt", default=PROMPT)
|
||||
p.add_argument("--fps", type=float, default=24.0,
|
||||
help="Frame-rate annotation passed to fps-aware metrics "
|
||||
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
|
||||
"by default.")
|
||||
p.add_argument("--metrics", default=",".join(DEFAULT_METRICS),
|
||||
help="Comma-separated metric names. Pass 'all' for every "
|
||||
"registered metric, or e.g. 'vbench' for the whole group.")
|
||||
p.add_argument("--scores-out", default="outputs_video/ltx2_eval/scores.json")
|
||||
p.add_argument("--skip-generation", action="store_true",
|
||||
help="Reuse an existing --output video instead of regenerating.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def generate(args: argparse.Namespace) -> Path:
|
||||
out = Path(args.output)
|
||||
if args.skip_generation and out.is_file():
|
||||
print(f"[gen] reusing existing video at {out}")
|
||||
return out
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
print(f"[gen] loading {args.model} ({args.num_gpus} GPU)...")
|
||||
generator = VideoGenerator.from_pretrained(args.model, num_gpus=args.num_gpus)
|
||||
try:
|
||||
print(f"[gen] generating to {out}...")
|
||||
generator.generate_video(
|
||||
prompt=args.prompt,
|
||||
output_path=str(out),
|
||||
save_video=True,
|
||||
num_frames=args.num_frames,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
return out
|
||||
|
||||
|
||||
def evaluate_video(video_path: Path, prompt: str, fps: float,
|
||||
metric_names) -> dict:
|
||||
print(f"[eval] loading video from {video_path}...")
|
||||
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
|
||||
video = video.unsqueeze(0) # → (1, T, C, H, W)
|
||||
|
||||
print(f"[eval] building evaluator: {metric_names}")
|
||||
evaluator = create_evaluator(metrics=metric_names, device="cuda")
|
||||
|
||||
print(f"[eval] running ({video.shape[1]} frames @ {fps} fps)...")
|
||||
results = evaluator.evaluate(
|
||||
video=video,
|
||||
text_prompt=[prompt],
|
||||
fps=fps,
|
||||
)
|
||||
|
||||
if isinstance(results, list):
|
||||
results = results[0] # batch of 1
|
||||
|
||||
return {
|
||||
name: {"score": r.score, "details": r.details}
|
||||
for name, r in results.items()
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
if args.metrics.strip() == "all":
|
||||
metric_names = "all"
|
||||
else:
|
||||
metric_names = [m.strip() for m in args.metrics.split(",") if m.strip()]
|
||||
|
||||
video_path = generate(args)
|
||||
scores = evaluate_video(video_path, args.prompt, args.fps, metric_names)
|
||||
|
||||
print("\n=== VBench scores ===")
|
||||
for name, payload in scores.items():
|
||||
print(f" {name}: {payload['score']}")
|
||||
|
||||
out = Path(args.scores_out)
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(
|
||||
{"video": str(video_path), "prompt": args.prompt, "scores": scores},
|
||||
indent=2,
|
||||
))
|
||||
print(f"[done] scores written to {out}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Score a folder of videos in parallel across multiple GPUs.
|
||||
|
||||
Uses :meth:`Evaluator.evaluate(samples=[...])`, which round-robins each
|
||||
sample dict across the GPU replicas the evaluator was built with.
|
||||
|
||||
Example::
|
||||
|
||||
python examples/inference/eval/score_folder.py \\
|
||||
--videos generated/ \\
|
||||
--metrics vbench.aesthetic_quality,vbench.subject_consistency \\
|
||||
--num-gpus 4 \\
|
||||
--output scores.json
|
||||
|
||||
Pair each generated video with a same-name reference video (e.g.
|
||||
``ref/<stem>.mp4``) by passing ``--reference-dir``::
|
||||
|
||||
python examples/inference/eval/score_folder.py \\
|
||||
--videos generated/ --reference-dir ref/ \\
|
||||
--metrics common.psnr,common.ssim,common.lpips \\
|
||||
--num-gpus 4
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
|
||||
|
||||
def _list_videos(directory: Path) -> list[Path]:
|
||||
exts = {".mp4", ".avi", ".mov", ".mkv", ".gif"}
|
||||
return sorted(p for p in directory.iterdir() if p.suffix.lower() in exts)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--videos", type=Path, required=True,
|
||||
help="Directory of generated videos.")
|
||||
p.add_argument("--reference-dir", type=Path, default=None,
|
||||
help="Directory of reference videos with matching stems.")
|
||||
p.add_argument("--metrics", default="vbench.aesthetic_quality")
|
||||
p.add_argument("--num-gpus", type=int, default=1)
|
||||
p.add_argument("--fps", type=float, default=None,
|
||||
help="Frame-rate annotation for fps-aware metrics.")
|
||||
p.add_argument("--output", type=Path, default=Path("scores.json"))
|
||||
args = p.parse_args()
|
||||
|
||||
video_paths = _list_videos(args.videos)
|
||||
if not video_paths:
|
||||
raise SystemExit(f"No videos under {args.videos}")
|
||||
print(f"Found {len(video_paths)} videos in {args.videos}")
|
||||
|
||||
metrics: list[str] | str = (
|
||||
args.metrics if args.metrics == "all"
|
||||
else [m.strip() for m in args.metrics.split(",") if m.strip()]
|
||||
)
|
||||
evaluator = create_evaluator(metrics=metrics, num_gpus=args.num_gpus)
|
||||
|
||||
# Build per-video sample dicts holding *paths*, not pre-loaded
|
||||
# tensors. Each path is decoded inside the worker thread that picks
|
||||
# up its sample, so peak resident memory is bounded by num_gpus
|
||||
# rather than scaling with the size of the folder.
|
||||
samples: list[dict] = []
|
||||
for vp in video_paths:
|
||||
sample: dict = {"video": str(vp)}
|
||||
if args.reference_dir is not None:
|
||||
ref_path = args.reference_dir / vp.name
|
||||
if not ref_path.is_file():
|
||||
raise FileNotFoundError(f"Missing reference for {vp.name} at {ref_path}")
|
||||
sample["reference"] = str(ref_path)
|
||||
if args.fps is not None:
|
||||
sample["fps"] = args.fps
|
||||
samples.append(sample)
|
||||
|
||||
print(f"Scoring with {len(evaluator.metric_names)} metric(s) "
|
||||
f"on {evaluator.num_gpus} GPU(s)...")
|
||||
all_results = evaluator.evaluate(samples=samples)
|
||||
evaluator.shutdown()
|
||||
|
||||
payload = [
|
||||
{
|
||||
"video": str(vp),
|
||||
"scores": {name: r.score for name, r in results.items()},
|
||||
}
|
||||
for vp, results in zip(video_paths, all_results)
|
||||
]
|
||||
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
args.output.write_text(json.dumps(payload, indent=2))
|
||||
print(f"Wrote {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Score one video on one GPU.
|
||||
|
||||
Smallest possible use of ``fastvideo.eval``: load an mp4, build an
|
||||
:class:`Evaluator` for the requested metric set, run it.
|
||||
|
||||
Examples::
|
||||
|
||||
# Reference-free (just the generated video):
|
||||
python examples/inference/eval/score_video.py \\
|
||||
--video clip.mp4 \\
|
||||
--metrics vbench.aesthetic_quality,vbench.imaging_quality
|
||||
|
||||
# Reference-paired (compare against ground truth):
|
||||
python examples/inference/eval/score_video.py \\
|
||||
--video gen.mp4 --reference ref.mp4 \\
|
||||
--metrics common.psnr,common.ssim,common.lpips
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
|
||||
def main() -> None:
|
||||
p = argparse.ArgumentParser(description=__doc__,
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter)
|
||||
p.add_argument("--video", required=True, help="Path to the generated mp4.")
|
||||
p.add_argument("--reference", default=None,
|
||||
help="Optional path to a reference mp4 (for paired metrics).")
|
||||
p.add_argument("--metrics", default="common.psnr,common.ssim",
|
||||
help="Comma-separated metric names, or a group name like 'vbench'.")
|
||||
p.add_argument("--device", default="cuda:0")
|
||||
p.add_argument("--text-prompt", default=None,
|
||||
help="Text prompt for prompt-aware metrics "
|
||||
"(vbench.overall_consistency, etc.).")
|
||||
p.add_argument("--fps", type=float, default=None,
|
||||
help="Frame-rate annotation for fps-aware metrics "
|
||||
"(vbench.dynamic_degree, etc.).")
|
||||
args = p.parse_args()
|
||||
|
||||
metrics: list[str] | str = (
|
||||
args.metrics if args.metrics in ("all",)
|
||||
else [m.strip() for m in args.metrics.split(",") if m.strip()]
|
||||
)
|
||||
evaluator = create_evaluator(metrics=metrics, device=args.device)
|
||||
|
||||
sample: dict = {"video": load_video(args.video)}
|
||||
if args.reference is not None:
|
||||
sample["reference"] = load_video(args.reference)
|
||||
if args.text_prompt is not None:
|
||||
sample["text_prompt"] = args.text_prompt
|
||||
if args.fps is not None:
|
||||
sample["fps"] = args.fps
|
||||
|
||||
results = evaluator.evaluate(**sample)
|
||||
evaluator.shutdown()
|
||||
|
||||
print(json.dumps(
|
||||
{name: r.score for name, r in results.items()},
|
||||
indent=2,
|
||||
))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,73 +0,0 @@
|
||||
# Cosmos Predict2 2B T2V finetune config.
|
||||
#
|
||||
# Data must be preprocessed with Cosmos VAE + T5 text encoder
|
||||
# into parquet format before training.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: nvidia/Cosmos-Predict2-2B-Video2World
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 8
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos_preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
# Cosmos VAE: 4x temporal, 8x spatial compression.
|
||||
# 93 frames -> 24 latent frames, 480x832 -> 60x104
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 1.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 5000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos_finetune
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos
|
||||
run_name: cosmos_finetune
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos_pipeline.Cosmos2VideoToWorldPipeline
|
||||
dataset_file: data/cosmos_preprocessed/validation_prompts.json
|
||||
every_steps: 100
|
||||
sampling_steps: [50]
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -1,79 +0,0 @@
|
||||
# Cosmos-Predict2.5-2B Text-to-World overfitting test config.
|
||||
#
|
||||
# Overfits on a few short videos (480x832, 93 frames) to verify the
|
||||
# Cosmos 2.5 training plugin works end-to-end.
|
||||
#
|
||||
# Preprocess data first:
|
||||
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_cosmos25_overfit.py
|
||||
#
|
||||
# Run:
|
||||
# bash examples/train/run.sh examples/train/configs/overfit_cosmos25_t2w.yaml
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.cosmos.CosmosModel
|
||||
init_from: KyleShao/Cosmos-Predict2.5-2B-Diffusers
|
||||
trainable: true
|
||||
enable_gradient_checkpointing_type: full
|
||||
flow_shift: 1.0
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 1
|
||||
|
||||
data:
|
||||
data_path: data/cosmos25_overfit_preprocessed
|
||||
dataloader_num_workers: 0
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 24
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 93
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 300
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/cosmos25_overfit
|
||||
training_state_checkpointing_steps: 50
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: fastvideo_cosmos25
|
||||
run_name: cosmos25_overfit
|
||||
|
||||
model:
|
||||
precondition_outputs: false
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.cosmos.cosmos2_5_pipeline.Cosmos2_5Pipeline
|
||||
dataset_file: data/cosmos25_overfit_preprocessed/validation_prompts.json
|
||||
every_steps: 150
|
||||
sampling_steps: [35]
|
||||
guidance_scale: 7.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 1.0
|
||||
@@ -1,7 +1,7 @@
|
||||
cmake_minimum_required(VERSION 3.26 FATAL_ERROR)
|
||||
project(fastvideo-kernel LANGUAGES CXX)
|
||||
|
||||
# Prefer environment variable (used by CI or uv pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
# Prefer environment variable (used by CI or 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., `uv pip install flash-attn`."
|
||||
"flash-attn is not installed. Please install it, e.g., `pip install flash-attn`."
|
||||
)
|
||||
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
# `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.
|
||||
@@ -15,7 +15,6 @@ class GenerationResult:
|
||||
samples: Any | None = None
|
||||
frames: Any | None = None
|
||||
audio: Any | None = None
|
||||
audio_sample_rate: int | None = None
|
||||
size: tuple[int, int, int] | None = None
|
||||
generation_time: float | None = None
|
||||
logging_info: Any | None = None
|
||||
@@ -45,7 +44,6 @@ class GenerationResult:
|
||||
"samples",
|
||||
"frames",
|
||||
"audio",
|
||||
"audio_sample_rate",
|
||||
"size",
|
||||
"generation_time",
|
||||
"logging_info",
|
||||
@@ -64,7 +62,6 @@ class GenerationResult:
|
||||
samples=result.get("samples"),
|
||||
frames=result.get("frames"),
|
||||
audio=result.get("audio"),
|
||||
audio_sample_rate=result.get("audio_sample_rate"),
|
||||
size=result.get("size"),
|
||||
generation_time=result.get("generation_time"),
|
||||
logging_info=result.get("logging_info"),
|
||||
@@ -83,7 +80,6 @@ class GenerationResult:
|
||||
"samples": self.samples,
|
||||
"frames": self.frames,
|
||||
"audio": self.audio,
|
||||
"audio_sample_rate": self.audio_sample_rate,
|
||||
"size": self.size,
|
||||
"generation_time": self.generation_time,
|
||||
"logging_info": self.logging_info,
|
||||
|
||||
@@ -112,33 +112,6 @@ class SamplingParam:
|
||||
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
|
||||
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
|
||||
|
||||
# Stable Audio (T2A): clip start/end in seconds. Honored by
|
||||
# `StableAudioConditioningStage` + `StableAudioDecodingStage`. Other
|
||||
# families ignore them.
|
||||
audio_start_in_s: float | None = None
|
||||
audio_end_in_s: float | None = None
|
||||
|
||||
# Stable Audio audio-to-audio (variation):
|
||||
# `init_audio` -- a path or `[B, C, samples]` waveform at the model
|
||||
# sample rate; the pipeline encodes it via the VAE
|
||||
# and uses it as the starting latent.
|
||||
# `init_audio_strength` -- 0..1, higher = closer to the reference
|
||||
# (matches the convention of Stability's
|
||||
# commercial Stable Audio 2.0 UI). 1.0 ~=
|
||||
# VAE round-trip, 0.0 ~= plain T2A.
|
||||
# `init_noise_level` -- legacy raw `sigma_max` override (0.3..500,
|
||||
# higher = more freedom). Kept for callers
|
||||
# that already use it; prefer `init_audio_strength`.
|
||||
init_audio: Any = None
|
||||
init_audio_strength: float | None = None
|
||||
init_noise_level: float | None = None
|
||||
|
||||
# Stable Audio inpainting (RePaint-style): `inpaint_audio` is the
|
||||
# reference clip, `inpaint_mask` is a [samples] tensor in {0, 1} where
|
||||
# 1 means *keep the reference* and 0 means *regenerate*.
|
||||
inpaint_audio: Any = None
|
||||
inpaint_mask: Any = None
|
||||
|
||||
# Continuation state carried across streaming/multi-segment calls.
|
||||
continuation_state: ContinuationState | None = None
|
||||
# When True, the pipeline returns a ContinuationState on the result so
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
# `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,6 +2,7 @@
|
||||
|
||||
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:
|
||||
@@ -17,7 +18,6 @@ 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: uv pip install git+https://github.com/thu-ml/SpargeAttn.git")
|
||||
"Install with: 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}"
|
||||
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
# `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.
|
||||
@@ -48,6 +48,8 @@ class ModelConfig:
|
||||
for key, value in source_model_dict.items():
|
||||
if key in valid_fields:
|
||||
setattr(arch_config, key, value)
|
||||
else:
|
||||
raise AttributeError(f"{type(arch_config).__name__} has no field '{key}'")
|
||||
|
||||
if hasattr(arch_config, "__post_init__"):
|
||||
arch_config.__post_init__()
|
||||
|
||||
@@ -5,14 +5,11 @@ 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
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"MagiHumanVideoConfig", "StableAudioConfig"
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig"
|
||||
]
|
||||
|
||||
@@ -50,8 +50,7 @@ class CosmosArchConfig(DiTArchConfig):
|
||||
})
|
||||
|
||||
# Cosmos-specific config parameters based on transformer_cosmos.py
|
||||
# in_channels includes the condition_mask channel (16 latent + 1 cond = 17)
|
||||
in_channels: int = 17
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads: int = 16
|
||||
attention_head_dim: int = 128
|
||||
|
||||
@@ -1,110 +0,0 @@
|
||||
# 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"
|
||||
@@ -1,76 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the Stable Audio Open 1.0 DiT.
|
||||
|
||||
Note: the SA pipeline bypasses the standard `ComposedPipelineBase`
|
||||
component loader because the published HF repo ships a single monolithic
|
||||
`model.safetensors` (no Diffusers-style `model_index.json` or
|
||||
per-subfolder layout). The arch fields and `param_names_mapping` here
|
||||
document the architecture and key remap so the same conventions used by
|
||||
the rest of the DiT family apply (FSDP shard conditions, supported
|
||||
attention backends, future loader integrations) — they are not currently
|
||||
consumed by `fastvideo/models/loader/fsdp_load.py` for SA.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
# Matches `transformer.layers.{i}` in the SA DiT module tree.
|
||||
parts = n.split(".")
|
||||
return (len(parts) >= 3 and parts[-3] == "transformer" and parts[-2] == "layers" and parts[-1].isdigit())
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer])
|
||||
|
||||
# SA's checkpoint is `stable_audio_tools` raw format (not Diffusers),
|
||||
# so the only remaps are: strip the `model.model.` host-pipeline
|
||||
# prefix, and rename `nn.LayerNorm`'s `gamma`/`beta` to torch's
|
||||
# canonical `weight`/`bias`. Linear / cross-attention naming already
|
||||
# matches FastVideo's conventions, so no further remap is needed.
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^model\.model\.(.*?)\.gamma$": r"\1.weight",
|
||||
r"^model\.model\.(.*?)\.beta$": r"\1.bias",
|
||||
r"^model\.model\.(.*)$": r"\1",
|
||||
})
|
||||
|
||||
# SA only supports backends compatible with single-GPU LocalAttention.
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
# Architecture constants (from the published `model_config.json` for
|
||||
# `stabilityai/stable-audio-open-1.0`).
|
||||
io_channels: int = 64
|
||||
embed_dim: int = 1536
|
||||
depth: int = 24
|
||||
num_attention_heads: int = 24
|
||||
cond_token_dim: int = 768
|
||||
global_cond_dim: int = 1536
|
||||
project_cond_tokens: bool = False
|
||||
project_global_cond: bool = True
|
||||
# Set to "ln" to wrap attention Q/K in LayerNorm (used by
|
||||
# `stable-audio-open-small`; absent in the 1.0 base).
|
||||
qk_norm: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.hidden_size = self.embed_dim
|
||||
self.in_channels = self.io_channels
|
||||
self.out_channels = self.io_channels
|
||||
self.num_channels_latents = self.io_channels
|
||||
self.attention_head_dim = self.embed_dim // self.num_attention_heads
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=StableAudioArchConfig)
|
||||
|
||||
prefix: str = "StableAudio"
|
||||
@@ -7,13 +7,9 @@ from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
|
||||
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
|
||||
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", "T5GemmaEncoderConfig"
|
||||
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig"
|
||||
]
|
||||
|
||||
@@ -1,80 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the Stable Audio Open 1.0 multi-conditioner.
|
||||
|
||||
The conditioner bundles three sub-conditioners — a T5 text encoder
|
||||
(prompt) and two NumberConditioners (`seconds_start` / `seconds_total`)
|
||||
— into the (cross_attn_cond, cross_attn_mask, global_embed) triple the
|
||||
DiT consumes. The architecture is fully specified by the official
|
||||
`stable_audio_tools` `MultiConditioner` config; the constants here
|
||||
mirror that.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.base import ArchConfig
|
||||
from fastvideo.configs.models.encoders.base import (EncoderArchConfig, EncoderConfig)
|
||||
|
||||
|
||||
def _default_configs() -> list[dict]:
|
||||
"""Default = `stable-audio-open-1.0`'s three sub-conditioners."""
|
||||
return [
|
||||
{
|
||||
"id": "prompt",
|
||||
"type": "t5",
|
||||
"config": {
|
||||
"t5_model_name": "t5-base",
|
||||
"max_length": 128
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "seconds_start",
|
||||
"type": "number",
|
||||
"config": {
|
||||
"min_val": 0,
|
||||
"max_val": 512
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "seconds_total",
|
||||
"type": "number",
|
||||
"config": {
|
||||
"min_val": 0,
|
||||
"max_val": 512
|
||||
}
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioConditionerArchConfig(EncoderArchConfig):
|
||||
architectures: list[str] = field(default_factory=lambda: ["StableAudioMultiConditioner"])
|
||||
|
||||
# Shared embedding width across all sub-conditioners (T5 last-hidden
|
||||
# dim and NumberEmbedder feature dim both = `cond_dim`).
|
||||
cond_dim: int = 768
|
||||
|
||||
# Sub-conditioner identifiers. Order in `cross_attention_cond_ids`
|
||||
# is the concat order for the cross-attn token sequence; order in
|
||||
# `global_cond_ids` is the concat order for the global FiLM-style
|
||||
# embedding.
|
||||
cross_attention_cond_ids: tuple[str, ...] = ("prompt", "seconds_start", "seconds_total")
|
||||
global_cond_ids: tuple[str, ...] = ("seconds_start", "seconds_total")
|
||||
|
||||
# Per-sub-conditioner spec list (mirrors upstream
|
||||
# `model_config.json.model.conditioning.configs`). Each entry is
|
||||
# `{"id": ..., "type": "t5"|"number", "config": {...}}`. The default
|
||||
# matches `stable-audio-open-1.0`; SA-small overrides via the
|
||||
# `conditioner/config.json` shipped in the converted repo.
|
||||
configs: list = field(default_factory=_default_configs)
|
||||
|
||||
# Match official `stable_audio_tools/models/conditioners.py:334`:
|
||||
# T5 is loaded directly in fp16.
|
||||
t5_dtype: str = "float16"
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioConditionerConfig(EncoderConfig):
|
||||
arch_config: ArchConfig = field(default_factory=StableAudioConditionerArchConfig)
|
||||
|
||||
prefix: str = "stable_audio_conditioner"
|
||||
@@ -41,14 +41,6 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
text_len: int = 512
|
||||
dtype: str | None = None
|
||||
gradient_checkpointing: bool = False
|
||||
# Extra fields present in upstream HF T5Config but unused by FastVideo's
|
||||
# encoder. Declared here so `update_model_arch` doesn't reject them when
|
||||
# loading repos like `stabilityai/stable-audio-open-1.0` that ship the
|
||||
# full HF config.
|
||||
n_positions: int = 512
|
||||
decoder_start_token_id: int = 0
|
||||
output_past: bool = True
|
||||
task_specific_params: dict | None = None
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
# 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"
|
||||
@@ -5,7 +5,6 @@ from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
from fastvideo.configs.models.vaes.oobleck import OobleckVAEArchConfig, OobleckVAEConfig
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
|
||||
|
||||
__all__ = [
|
||||
@@ -17,6 +16,4 @@ __all__ = [
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
]
|
||||
|
||||
@@ -1,68 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Config for the Stable Audio Open 1.0 "Oobleck" VAE.
|
||||
|
||||
Mirrors the per-channel `vae/config.json` shipped in
|
||||
`stabilityai/stable-audio-open-1.0` 1:1 (see
|
||||
`fastvideo/models/vaes/oobleck.py::OobleckVAE.from_pretrained`, which
|
||||
constructs the VAE from these fields). Inherits the FastVideo VAEConfig
|
||||
base so the standard `load_encoder` / `load_decoder` flags + tiling
|
||||
knobs apply.
|
||||
|
||||
Naming: the VAE architecture is officially "Oobleck" (per Stability
|
||||
AI's stable-audio-tools) — the surrounding model family is "Stable
|
||||
Audio Open 1.0". This config is named after the architecture
|
||||
(`OobleckVAEConfig`) since the same VAE is shared across Stable Audio
|
||||
checkpoints; downstream pipelines reference it by its arch name, not
|
||||
by a host-pipeline name.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class OobleckVAEArchConfig(VAEArchConfig):
|
||||
"""Stable Audio Open 1.0 VAE architecture constants."""
|
||||
|
||||
architectures: list[str] = field(default_factory=lambda: ["AutoencoderOobleck"])
|
||||
|
||||
# From stabilityai/stable-audio-open-1.0/vae/config.json.
|
||||
encoder_hidden_size: int = 128
|
||||
downsampling_ratios: list[int] = field(default_factory=lambda: [2, 4, 4, 8, 8])
|
||||
channel_multiples: list[int] = field(default_factory=lambda: [1, 2, 4, 8, 16])
|
||||
decoder_channels: int = 128
|
||||
decoder_input_channels: int = 64
|
||||
audio_channels: int = 2 # stereo
|
||||
sampling_rate: int = 44100
|
||||
|
||||
|
||||
@dataclass
|
||||
class OobleckVAEConfig(VAEConfig):
|
||||
"""FastVideo VAE config wrapping the Oobleck arch.
|
||||
|
||||
Audio VAEs don't use the temporal/spatial tiling defaults that the
|
||||
base VAEConfig is shaped for (those exist for video VAEs); they are
|
||||
retained but irrelevant for audio.
|
||||
"""
|
||||
|
||||
arch_config: VAEArchConfig = field(default_factory=OobleckVAEArchConfig)
|
||||
|
||||
# Audio is 1-D, so the video-VAE tiling defaults are inert. Disable
|
||||
# them so callers don't accidentally trip on tile-stride math built
|
||||
# for spatial tensors.
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
# Where the FastVideo loader / pipeline-glue wrapper should fetch
|
||||
# weights from when no local path is supplied. Gated repo — caller's
|
||||
# HF token must have accepted terms on
|
||||
# https://huggingface.co/stabilityai/stable-audio-open-1.0.
|
||||
pretrained_path: str = "stabilityai/stable-audio-open-1.0"
|
||||
pretrained_subfolder: str = "vae"
|
||||
# Match official `stable_audio_tools`: VAE runs in fp16 (the
|
||||
# `pretransform.model_half` path in
|
||||
# `stable_audio_tools/models/pretransforms.py`).
|
||||
pretrained_dtype: str = "float16"
|
||||
@@ -1,67 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""`PipelineConfig` for Stable Audio Open 1.0."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import StableAudioConfig
|
||||
from fastvideo.configs.models.vaes import OobleckVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioT2AConfig(PipelineConfig):
|
||||
"""Stable Audio Open 1.0 pipeline config."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=StableAudioConfig)
|
||||
# Standard `TransformerLoader` reads `dit_precision`; default in
|
||||
# `PipelineConfig` is bf16, but we want fp16 to match official.
|
||||
dit_precision: str = "fp16"
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=OobleckVAEConfig)
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# `StableAudioMultiConditioner` owns its own T5; zero out the
|
||||
# parent's text-encoder slots so the length-equality validator passes.
|
||||
text_encoder_configs: tuple = field(default_factory=tuple)
|
||||
preprocess_text_funcs: tuple = field(default_factory=tuple)
|
||||
postprocess_text_funcs: tuple = field(default_factory=tuple)
|
||||
|
||||
num_inference_steps: int = 100
|
||||
guidance_scale: float = 7.0
|
||||
audio_end_in_s: float = 10.0 # short-clip default
|
||||
audio_start_in_s: float = 0.0
|
||||
sampling_rate: int = 44100
|
||||
audio_channels: int = 2
|
||||
# Stable Audio Open 1.0 was trained at a fixed 2,097,152-sample
|
||||
# window (= 2097152 / 44100 ≈ 47.55s). Anything past this is
|
||||
# silently truncated by the post-decode slice — validate up-front.
|
||||
sample_size: int = 2097152
|
||||
max_audio_duration_s: float = 2097152 / 44100
|
||||
|
||||
# Match the official `stable_audio_tools` defaults (`model_half=True`
|
||||
# in `run_gradio.py`), which loads the DiT, VAE, and T5 in fp16 and
|
||||
# wraps T5 forward in `autocast(fp16)`. fp16 is also a hard
|
||||
# requirement for FlashAttention-2 / FA-3.
|
||||
precision: str = "fp16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# A2A needs encode; load both halves for either path.
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class StableAudioOpenSmallConfig(StableAudioT2AConfig):
|
||||
"""`stable-audio-open-small` overrides: shorter training window
|
||||
(524288 samples ≈ 11.89s @ 44.1 kHz) and a faster default sampler
|
||||
config carried by the small preset.
|
||||
"""
|
||||
|
||||
sample_size: int = 524288
|
||||
max_audio_duration_s: float = 524288 / 44100
|
||||
audio_end_in_s: float = 6.0 # short-clip default suitable for the small window
|
||||
@@ -0,0 +1,216 @@
|
||||
"""``fastvideo eval`` CLI: list registered eval metrics and run them
|
||||
against a set of videos.
|
||||
|
||||
This is a thin wrapper around :mod:`fastvideo.eval`. Heavy lifting
|
||||
(metric loading, GPU handling, batching) lives in
|
||||
:func:`fastvideo.eval.create_evaluator`.
|
||||
|
||||
Examples::
|
||||
|
||||
fastvideo eval list
|
||||
fastvideo eval list --group vbench
|
||||
fastvideo eval run --videos path/to/videos/*.mp4 \\
|
||||
--metrics common.ssim --reference path/to/refs/
|
||||
fastvideo eval run --videos clip.mp4 --metrics vbench.aesthetic_quality \\
|
||||
--output scores.json
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class EvalSubcommand(CLISubcommand):
|
||||
"""The ``eval`` subcommand — entry point for the eval suite."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.name = "eval"
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
action = getattr(args, "eval_action", None)
|
||||
if action == "list":
|
||||
_cmd_list(args)
|
||||
elif action == "run":
|
||||
_cmd_run(args)
|
||||
else:
|
||||
# Re-print help if no action was given.
|
||||
self._parser.print_help() # type: ignore[attr-defined]
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
action = getattr(args, "eval_action", None)
|
||||
if action == "run" and not args.videos:
|
||||
raise SystemExit("`fastvideo eval run` requires --videos")
|
||||
|
||||
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
eval_parser = subparsers.add_parser(
|
||||
"eval",
|
||||
help="Run video-gen evaluation metrics",
|
||||
usage="fastvideo eval {list,run} [...]",
|
||||
)
|
||||
sub = eval_parser.add_subparsers(dest="eval_action", required=False)
|
||||
|
||||
# `eval list`
|
||||
list_p = sub.add_parser("list", help="List registered metrics")
|
||||
list_p.add_argument("--group", type=str, default=None, help="Filter to a metric group (e.g. 'vbench').")
|
||||
|
||||
# `eval run`
|
||||
run_p = sub.add_parser("run", help="Evaluate videos against one or more metrics")
|
||||
run_p.add_argument("--videos",
|
||||
type=str,
|
||||
nargs="+",
|
||||
required=False,
|
||||
help="Path, glob, or directory of generated videos.")
|
||||
run_p.add_argument("--reference",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path / glob / dir of reference videos (for paired metrics).")
|
||||
run_p.add_argument("--metrics", type=str, default="all", help="Comma-separated metric names, or 'all'.")
|
||||
run_p.add_argument("--device", type=str, default="cuda", help="Torch device (e.g. 'cuda', 'cuda:0', 'cpu').")
|
||||
run_p.add_argument("--text-prompt",
|
||||
type=str,
|
||||
nargs="*",
|
||||
default=None,
|
||||
help="Prompt(s) for text-conditioned metrics. One per video.")
|
||||
run_p.add_argument("--fps", type=float, default=None, help="Frame-rate annotation passed to fps-aware metrics.")
|
||||
run_p.add_argument("--output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Write results as JSON to this path (default: stdout).")
|
||||
|
||||
# Stash the parser so cmd() can re-print help on no-action.
|
||||
self._parser = eval_parser # type: ignore[attr-defined]
|
||||
return cast(FlexibleArgumentParser, eval_parser)
|
||||
|
||||
|
||||
def _cmd_list(args: argparse.Namespace) -> None:
|
||||
from fastvideo.eval import list_metrics
|
||||
names = list_metrics()
|
||||
if args.group:
|
||||
prefix = args.group.rstrip(".") + "."
|
||||
names = [n for n in names if n == args.group or n.startswith(prefix)]
|
||||
if not names:
|
||||
print(f"(no metrics matched group {args.group!r})")
|
||||
return
|
||||
for name in names:
|
||||
print(name)
|
||||
|
||||
|
||||
def _cmd_run(args: argparse.Namespace) -> None:
|
||||
from fastvideo.eval import create_evaluator
|
||||
from fastvideo.eval.io import load_video
|
||||
|
||||
video_paths = _expand_paths(args.videos)
|
||||
if not video_paths:
|
||||
raise SystemExit(f"No videos matched: {args.videos}")
|
||||
ref_paths = _expand_paths([args.reference]) if args.reference else None
|
||||
|
||||
metrics_arg: list[str] | str = ("all" if args.metrics == "all" else
|
||||
[m.strip() for m in args.metrics.split(",") if m.strip()])
|
||||
|
||||
evaluator = create_evaluator(metrics=metrics_arg, device=args.device)
|
||||
|
||||
all_results: list[dict] = []
|
||||
for i, vp in enumerate(video_paths):
|
||||
logger.info("Evaluating %s (%d/%d)", vp, i + 1, len(video_paths))
|
||||
kwargs: dict = {"video": load_video(vp)}
|
||||
if ref_paths is not None:
|
||||
ref = ref_paths[i] if i < len(ref_paths) else ref_paths[0]
|
||||
kwargs["reference"] = load_video(ref)
|
||||
if args.text_prompt is not None:
|
||||
prompt = (args.text_prompt[i] if i < len(args.text_prompt) else args.text_prompt[0])
|
||||
kwargs["text_prompt"] = [prompt]
|
||||
if args.fps is not None:
|
||||
kwargs["fps"] = args.fps
|
||||
|
||||
results = evaluator.evaluate(**kwargs)
|
||||
all_results.append({
|
||||
"video": str(vp),
|
||||
"scores": _serialize_results(results),
|
||||
})
|
||||
|
||||
payload = json.dumps(all_results, indent=2, default=_jsonable)
|
||||
if args.output:
|
||||
Path(args.output).write_text(payload)
|
||||
logger.info("Wrote results to %s", args.output)
|
||||
else:
|
||||
print(payload)
|
||||
|
||||
|
||||
def _expand_paths(patterns: list[str]) -> list[str]:
|
||||
out: list[str] = []
|
||||
for pat in patterns:
|
||||
p = Path(pat)
|
||||
if p.is_dir():
|
||||
for ext in (".mp4", ".avi", ".mov", ".mkv", ".gif"):
|
||||
out.extend(sorted(str(f) for f in p.iterdir() if f.suffix.lower() == ext))
|
||||
elif any(c in pat for c in "*?["):
|
||||
out.extend(sorted(glob.glob(pat)))
|
||||
else:
|
||||
out.append(pat)
|
||||
# de-dup, preserve order
|
||||
seen: set[str] = set()
|
||||
deduped: list[str] = []
|
||||
for x in out:
|
||||
if x not in seen:
|
||||
seen.add(x)
|
||||
deduped.append(x)
|
||||
return deduped
|
||||
|
||||
|
||||
def _serialize_results(results) -> dict | list:
|
||||
"""Turn evaluator output (dict or list-of-dicts) into JSON-friendly form."""
|
||||
if isinstance(results, list):
|
||||
return [_serialize_results(r) for r in results]
|
||||
if isinstance(results, dict):
|
||||
return {k: _serialize_metric_result(v) for k, v in results.items()}
|
||||
return _serialize_metric_result(results)
|
||||
|
||||
|
||||
def _serialize_metric_result(mr) -> dict:
|
||||
return {
|
||||
"name": getattr(mr, "name", None),
|
||||
"score": getattr(mr, "score", None),
|
||||
"details": getattr(mr, "details", None),
|
||||
}
|
||||
|
||||
|
||||
def _jsonable(obj):
|
||||
"""``json.dumps(default=...)`` coercer for metric outputs.
|
||||
|
||||
Metrics frequently land numpy scalars / arrays, torch tensors, and
|
||||
pathlib paths inside ``MetricResult.details`` (e.g. ``optical_flow``
|
||||
populates ``per_frame_metrics`` with numpy floats). The stdlib JSON
|
||||
encoder rejects all of those by default — this callback walks the
|
||||
leaves and coerces them to native Python types.
|
||||
"""
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
if isinstance(obj, np.integer):
|
||||
return int(obj)
|
||||
if isinstance(obj, np.floating):
|
||||
return float(obj)
|
||||
if isinstance(obj, np.bool_):
|
||||
return bool(obj)
|
||||
if isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
if isinstance(obj, torch.Tensor):
|
||||
return obj.detach().cpu().tolist()
|
||||
if isinstance(obj, Path):
|
||||
return str(obj)
|
||||
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
return [EvalSubcommand()]
|
||||
@@ -3,10 +3,9 @@
|
||||
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
|
||||
from fastvideo.entrypoints.cli.eval import cmd_init as eval_cmd_init
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
@@ -14,8 +13,8 @@ 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())
|
||||
commands.extend(eval_cmd_init())
|
||||
return commands
|
||||
|
||||
|
||||
|
||||
@@ -1,115 +0,0 @@
|
||||
# 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,28 +11,6 @@ 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,
|
||||
@@ -42,24 +20,12 @@ __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",
|
||||
]
|
||||
|
||||
@@ -1,542 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -1,122 +0,0 @@
|
||||
# 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()
|
||||
@@ -1,36 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -1,197 +0,0 @@
|
||||
# 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"]
|
||||
@@ -1,24 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -1,101 +0,0 @@
|
||||
# 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"]
|
||||
@@ -1,85 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -1,44 +0,0 @@
|
||||
# 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"]
|
||||
@@ -1,46 +0,0 @@
|
||||
# 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"]
|
||||
@@ -1,82 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -1,146 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
@@ -1,27 +0,0 @@
|
||||
# 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",
|
||||
]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user