Compare commits
48
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7683593c77 | ||
|
|
a253856147 | ||
|
|
6e25d94ebc | ||
|
|
cae8fa18dc | ||
|
|
821e5a0832 | ||
|
|
c1abc42782 | ||
|
|
ef15ea2391 | ||
|
|
1ea2517e22 | ||
|
|
0c63528c59 | ||
|
|
b063f8ca41 | ||
|
|
e7fff0173a | ||
|
|
d82abc271e | ||
|
|
970409962f | ||
|
|
055586703d | ||
|
|
5d89f86675 | ||
|
|
19a51a1fe6 | ||
|
|
d3232cea5a | ||
|
|
0c90c8c24d | ||
|
|
4c08ffce49 | ||
|
|
af4a77553c | ||
|
|
c096fda1eb | ||
|
|
8f47e85be0 | ||
|
|
90d3bd19eb | ||
|
|
afb4f7d3c5 | ||
|
|
02e1143f22 | ||
|
|
e2f4d1a7b5 | ||
|
|
f037351146 | ||
|
|
1ee11e08dc | ||
|
|
595f0ea60e | ||
|
|
d921832cd2 | ||
|
|
629697629a | ||
|
|
a25313beec | ||
|
|
dbde64385b | ||
|
|
9d909f5f04 | ||
|
|
76b0550c15 | ||
|
|
384c1e9493 | ||
|
|
b1dbcc93f6 | ||
|
|
b93833772e | ||
|
|
30b523edd6 | ||
|
|
6aab7f3832 | ||
|
|
9cd53fe5f8 | ||
|
|
6a32cf3a5e | ||
|
|
98be9b3da2 | ||
|
|
c53e85b767 | ||
|
|
40a8bd2d3b | ||
|
|
31aa115611 | ||
|
|
98ac10a528 | ||
|
|
a5a6d171e5 |
@@ -1,20 +1,23 @@
|
||||
---
|
||||
name: reseed-performance-baseline
|
||||
description: Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift using one or more reviewed normalized performance JSONs. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline in FastVideo/performance-tracking must be advanced from a consistent batch of reviewed source results. The workflow backs up existing history under /tmp, validates all source JSONs for the same (model_id, gpu_type), rejects internally inconsistent source batches, uploads one success=true reseed record per accepted source JSON, and offers to clean local temp state after a successful upload.
|
||||
description: Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, environment-caused benchmark shift, or reviewed v2 calibration using one or more reviewed normalized performance JSONs. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline in FastVideo/performance-tracking must be advanced, or when a new v2 exact comparable identity needs its first approved baseline. The workflow backs up existing history under /tmp, validates all source JSONs for the same legacy (model_id, gpu_type) target or the same v2 exact identity, rejects internally inconsistent source batches, uploads one success=true baseline record per accepted source JSON, and offers to clean local temp state after a successful upload.
|
||||
---
|
||||
|
||||
# Re-seed Performance Baseline
|
||||
|
||||
## Purpose
|
||||
|
||||
Replace or advance the rolling performance baseline for a single
|
||||
`(model_id, gpu_type)` pair in the HF dataset
|
||||
`FastVideo/performance-tracking`.
|
||||
Replace or advance the rolling performance baseline in the HF dataset
|
||||
`FastVideo/performance-tracking`. Legacy targets are scoped by
|
||||
`(model_id, gpu_type)`. V2 targets are scoped by exact comparable identity:
|
||||
`workload_id`, `variant_id`, `benchmark_version`, `hardware_profile_id`,
|
||||
`software_profile_id`, and `recipe_fingerprint`.
|
||||
|
||||
Performance comparison uses the median of up to the last 5 successful records
|
||||
for the same model and GPU. Failed records are useful audit history, but they
|
||||
do not move the future baseline because `compare_baseline.py` loads records
|
||||
with `successful_only=True`.
|
||||
Performance comparison uses the median of up to the last 5 successful,
|
||||
baseline-eligible records for the same target. Failed or calibration-only
|
||||
records are useful audit history, but they do not move the future baseline
|
||||
because `compare_baseline.py` loads records with `successful_only=True` and
|
||||
`baseline_eligible_only=True`.
|
||||
|
||||
This skill now reseeds from a reviewed batch of one or more source performance
|
||||
JSONs. It uploads one new `success=true` record per accepted source JSON; it
|
||||
@@ -22,11 +25,13 @@ does not blindly replicate one measurement into 3 or 5 records. The effective
|
||||
reseed size is therefore dynamic and equals the number of provided, validated,
|
||||
internally consistent source JSONs.
|
||||
|
||||
If the operator provides fewer than 3 records, call out that the last-5 rolling
|
||||
median may not move immediately. If the operator provides 3 consistent shifted
|
||||
records, the rolling median usually moves immediately. If the operator provides
|
||||
5 consistent shifted records, the last-5 window is effectively reset to the new
|
||||
runtime profile.
|
||||
For baseline shifts with existing history, if the operator provides fewer than
|
||||
3 records, call out that the last-5 rolling median may not move immediately. If
|
||||
the operator provides 3 consistent shifted records, the rolling median usually
|
||||
moves immediately. If the operator provides 5 consistent shifted records, the
|
||||
last-5 window is effectively reset to the new runtime profile. For the first
|
||||
approved v2 baseline of a new exact identity, one reviewed calibration seed is
|
||||
enough for the next comparable run to leave `CALIBRATION_NEEDED`.
|
||||
|
||||
These records are intentional operator-approved baseline resets, not ordinary
|
||||
independent main-branch persistence. Mark them clearly with provenance fields
|
||||
@@ -66,10 +71,10 @@ approval, then upload reviewed accepted baseline records.
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `model_id` | Legacy required; v2 inferred | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. For v2 records, use the `model_id` from each source artifact only as the upload directory; comparison is by exact identity. |
|
||||
| `gpu_type` | Legacy required; v2 inferred | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. V2 hardware matching uses `hardware_profile_id`; preserve `gpu_type` as display metadata. |
|
||||
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
@@ -78,10 +83,14 @@ Hardcoded defaults:
|
||||
supported by the code, but use the default unless the user explicitly asks).
|
||||
- Local sync root: `/tmp/perf-tracking` (`PERFORMANCE_TRACKING_ROOT` override
|
||||
is supported).
|
||||
- Prepared-record staging root: `/tmp/performance_reseed_prepared`
|
||||
(`PERFORMANCE_RESEED_STAGING_ROOT` override is supported). Keep it separate
|
||||
and non-nested from the sync root.
|
||||
- Backup root: `/tmp/performance_reseed_backup`.
|
||||
- Download scratch root for source artifact URLs: `/tmp/performance_reseed_source`.
|
||||
- Baseline window: last 5 `success=true` records for the same
|
||||
`(model_id, gpu_type)`.
|
||||
- Baseline window: last 5 `success=true`, `baseline_eligible=true` records
|
||||
for the same legacy `(model_id, gpu_type)` target or the same v2 exact
|
||||
comparable identity.
|
||||
- Reseed count: dynamic. Upload exactly one accepted seed record per validated
|
||||
source JSON.
|
||||
|
||||
@@ -115,12 +124,24 @@ with open(source_result, encoding="utf-8") as f:
|
||||
record = json.load(f)
|
||||
```
|
||||
|
||||
Stop if any normalized record's `model_id` or `gpu_type` does not match the
|
||||
requested `model_id` and `gpu_type`.
|
||||
Classify the source batch before continuing:
|
||||
|
||||
- **Legacy source records** have no v2 exact identity fields. Stop if any
|
||||
normalized record's `model_id` or `gpu_type` does not match the requested
|
||||
`model_id` and `gpu_type`.
|
||||
- **V2 source records** have exact identity fields. Stop unless every source
|
||||
record has all six comparable identity fields and they are identical across
|
||||
the batch: `workload_id`, `variant_id`, `benchmark_version`,
|
||||
`hardware_profile_id`, `software_profile_id`, and `recipe_fingerprint`.
|
||||
Do not fall back to legacy `(model_id, gpu_type)` matching for v2 records.
|
||||
|
||||
The source records may have `success: false` when they came from failed
|
||||
rolling baseline comparisons. That is expected; only the reviewed reseed
|
||||
records become new `success: true` baseline records after explicit approval.
|
||||
For a first v2 baseline seed, the source records must instead be successful
|
||||
scheduled-main full-suite `CALIBRATION_NEEDED` normalized artifacts. Reject PR,
|
||||
local, direct-run, non-main-branch, or non-full-suite calibration artifacts as
|
||||
seed sources.
|
||||
|
||||
Sort validated source records by their original `timestamp` ascending before
|
||||
preparing the seed records. If a source timestamp is missing or unparsable,
|
||||
@@ -148,8 +169,7 @@ For each metric with at least two non-null source values:
|
||||
4. Stop if any source record regresses against the batch median by more than
|
||||
`max_intra_batch_regression`.
|
||||
|
||||
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
|
||||
otherwise `0.05`. Print a table with per-source values, batch median, and
|
||||
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
|
||||
worst intra-batch regression.
|
||||
|
||||
This check prevents uploading a mixed batch where one JSON is materially
|
||||
@@ -183,7 +203,7 @@ present, that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records under /tmp
|
||||
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
@@ -192,16 +212,16 @@ The sync command pattern is:
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
For legacy records, back up the sanitized model directory under `/tmp`:
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
|
||||
from hf_store import sanitize
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.performance.hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
@@ -210,6 +230,16 @@ mkdir -p "$BACKUP_DIR"
|
||||
cp -R "${PERFORMANCE_TRACKING_ROOT}/${MODEL_SAFE}" "$BACKUP_DIR/" 2>/dev/null || true
|
||||
```
|
||||
|
||||
For v2 records, back up the full local tracking root after sync. Exact identity
|
||||
lookup scans across model directories, so a benchmark rename may have relevant
|
||||
history outside the current source artifact's `model_id` directory:
|
||||
|
||||
```bash
|
||||
BACKUP_DIR="/tmp/performance_reseed_backup/${TIMESTAMP}_${SHORT_COMMIT}_v2_exact_identity"
|
||||
mkdir -p "$BACKUP_DIR"
|
||||
cp -R "${PERFORMANCE_TRACKING_ROOT}" "$BACKUP_DIR/tracking-root"
|
||||
```
|
||||
|
||||
Write provenance next to the backup:
|
||||
|
||||
```bash
|
||||
@@ -232,10 +262,12 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
|
||||
### 3. Compute old baseline and candidate shift
|
||||
|
||||
Load the last 5 successful records for the target:
|
||||
Load the last 5 successful baseline records for the target.
|
||||
|
||||
For legacy targets:
|
||||
|
||||
```python
|
||||
from hf_store import load_records_for_model
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
@@ -243,6 +275,28 @@ records = load_records_for_model(
|
||||
"<gpu_type>",
|
||||
last_n=5,
|
||||
successful_only=True,
|
||||
baseline_eligible_only=True,
|
||||
)
|
||||
```
|
||||
|
||||
For v2 exact-identity targets:
|
||||
|
||||
```python
|
||||
from fastvideo.performance.hf_store import load_records_for_identity
|
||||
|
||||
records = load_records_for_identity(
|
||||
"/tmp/perf-tracking",
|
||||
{
|
||||
"workload_id": "<workload_id>",
|
||||
"variant_id": "<variant_id>",
|
||||
"benchmark_version": "<benchmark_version>",
|
||||
"hardware_profile_id": "<hardware_profile_id>",
|
||||
"software_profile_id": "<software_profile_id>",
|
||||
"recipe_fingerprint": "<recipe_fingerprint>",
|
||||
},
|
||||
last_n=5,
|
||||
successful_only=True,
|
||||
baseline_eligible_only=True,
|
||||
)
|
||||
```
|
||||
|
||||
@@ -258,7 +312,8 @@ medians after appending the proposed seed records, and source batch spread for:
|
||||
|
||||
Also print how many successful old records exist. Make clear:
|
||||
|
||||
- 1 seed record usually does not move a last-5 median by itself.
|
||||
- 1 seed record usually does not move an existing last-5 median by itself, but
|
||||
it is enough to establish the first v2 baseline for a new exact identity.
|
||||
- 3 consistent seed records usually move the last-5 median immediately.
|
||||
- 5 consistent seed records effectively reset the last-5 window.
|
||||
- The records are intentional approved baseline resets and must be labeled
|
||||
@@ -268,10 +323,10 @@ Also print how many successful old records exist. Make clear:
|
||||
|
||||
Require an explicit confirmation phrase before preparing the upload:
|
||||
|
||||
> About to RE-SEED performance baseline for `<model_id>` on `<gpu_type>`.
|
||||
> About to RE-SEED performance baseline for `<target description>`.
|
||||
> This will upload `<N>` new `success=true` records to
|
||||
> `FastVideo/performance-tracking/<sanitize(model_id)>/`, one per accepted
|
||||
> source JSON.
|
||||
> `FastVideo/performance-tracking/<sanitize(model_id)>/` or the source
|
||||
> artifact's v2 model directory, one per accepted source JSON.
|
||||
>
|
||||
> Reason: `<intent_rationale>`
|
||||
> Source results: `<source_results>`
|
||||
@@ -289,25 +344,63 @@ Do not continue unless the user types exactly `confirm performance reseed`.
|
||||
|
||||
### 5. Create the accepted seed records
|
||||
|
||||
Create one seed record from each normalized source result. Do not copy the
|
||||
Create one seed record from each normalized source result.
|
||||
|
||||
For first v2 baseline seeds, use the scoped utility. It validates exact
|
||||
identity, requires successful scheduled-main full-suite `CALIBRATION_NEEDED`
|
||||
source artifacts, preserves the normalized v2 identity and metadata fields,
|
||||
and writes seed records with `success=true`, `baseline_eligible=true`, and
|
||||
`comparison_status=PASS`:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/performance/seed_baseline.py \
|
||||
--source-result <normalized_perf_1.json> \
|
||||
--source-result <normalized_perf_2.json> \
|
||||
--intent-rationale "<intent_rationale>" \
|
||||
--max-intra-batch-regression 0.05 \
|
||||
--tracking-root "${PERFORMANCE_TRACKING_ROOT}" \
|
||||
--staging-root "${PERFORMANCE_RESEED_STAGING_ROOT:-/tmp/performance_reseed_prepared}"
|
||||
```
|
||||
|
||||
The utility is prepare-only and intentionally has no upload option. Upload the
|
||||
scoped records only after the separate confirmation in step 6.
|
||||
|
||||
The utility validates against an isolated fresh HF snapshot and leaves
|
||||
`PERFORMANCE_TRACKING_ROOT` untouched; that argument only proves the staging
|
||||
root is separate from the operator's tracking mirror. Before writing, it stops
|
||||
if the exact identity already has a successful baseline-eligible record or if
|
||||
the workload/variant/version already trusts another recipe. It atomically
|
||||
reserves the exact identity and writes a digest-protected upload manifest bound
|
||||
to the current HF endpoint, repository id, and repository type. Keep the
|
||||
prepared records, manifest, source files, and reservation unchanged until the
|
||||
operation is uploaded or explicitly cleaned up.
|
||||
|
||||
If the prepared seed records look correct, upload only those scoped records in
|
||||
step 7. Do not rerun the utility with a different source list after approval.
|
||||
|
||||
For legacy reseeds or accepted v2 baseline shifts from regression artifacts,
|
||||
create one seed record from each normalized source result. Do not copy the
|
||||
source JSON wholesale.
|
||||
|
||||
Infer the baseline field allowlist from all existing HF records for the target
|
||||
`(model_id, gpu_type)` after syncing, including both `success=true` and
|
||||
`success=false` records. Use the union of non-provenance keys present in those
|
||||
target records, preserving only fields that also exist in the normalized
|
||||
source record or are explicitly set by the reseed workflow. Always include
|
||||
`model_id`, `timestamp`, and `success` because the upload path and baseline
|
||||
loader depend on them. Always set `timestamp` to a fresh reseed timestamp and
|
||||
`success` to `true`. Do not include unrelated source-only fields that are
|
||||
absent from existing HF records.
|
||||
after syncing, including both `success=true` and `success=false` records. For
|
||||
legacy targets the target is `(model_id, gpu_type)`. For v2 baseline-shift
|
||||
reseeds the target is the exact comparable identity. Use the union of
|
||||
non-provenance keys present in those target records, preserving only fields
|
||||
that also exist in the normalized source record or are explicitly set by the
|
||||
reseed workflow. Always include `model_id`, `timestamp`, `success`,
|
||||
`baseline_eligible`, and `comparison_status` because the upload path and
|
||||
baseline loader depend on them. Always set `timestamp` to a fresh reseed
|
||||
timestamp, `success` to `true`, `baseline_eligible` to `true`, and
|
||||
`comparison_status` to `PASS`. Do not include unrelated source-only fields
|
||||
that are absent from existing HF records.
|
||||
|
||||
Exclude existing provenance or operator metadata from the inferred baseline
|
||||
field allowlist. At minimum, exclude keys prefixed with `baseline_reseed` and
|
||||
any fields known to be local-only audit metadata.
|
||||
|
||||
If there are no previous HF records for the target model/GPU, fall back to this
|
||||
default baseline field list:
|
||||
If there are no previous HF records for the target, fall back to this default
|
||||
baseline field list:
|
||||
|
||||
- `model_id`
|
||||
- `timestamp`
|
||||
@@ -320,6 +413,22 @@ default baseline field list:
|
||||
- `dit_time_s`
|
||||
- `vae_decode_time_s`
|
||||
- `success`
|
||||
- `baseline_eligible`
|
||||
- `comparison_status`
|
||||
|
||||
For v2 baseline-shift reseeds with no previous HF records for the exact
|
||||
identity, also preserve:
|
||||
|
||||
- `workload_id`
|
||||
- `variant_id`
|
||||
- `benchmark_version`
|
||||
- `recipe_fingerprint`
|
||||
- `hardware_profile_id`
|
||||
- `software_profile_id`
|
||||
- `recipe`
|
||||
- `hardware_profile`
|
||||
- `software_profile`
|
||||
- `software_comparison_profile`
|
||||
|
||||
Do not upload extra fields from the source artifact.
|
||||
|
||||
@@ -335,6 +444,22 @@ Optional provenance fields are allowed and useful:
|
||||
- `baseline_reseed_operator`
|
||||
- `baseline_reseed_max_intra_batch_regression`
|
||||
|
||||
The v2 calibration seed utility writes analogous first-seed provenance:
|
||||
|
||||
- `baseline_seed: true`
|
||||
- `baseline_seed_reason`
|
||||
- `baseline_seed_source_result`
|
||||
- `baseline_seed_source_status`
|
||||
- `baseline_seed_source_timestamp`
|
||||
- `baseline_seed_source_success`
|
||||
- `baseline_seed_source_run_source`
|
||||
- `baseline_seed_source_branch`
|
||||
- `baseline_seed_source_test_scope`
|
||||
- `baseline_seed_source_pr_number`
|
||||
- `baseline_seed_batch_size`
|
||||
- `baseline_seed_batch_index`
|
||||
- `baseline_seed_operator`
|
||||
|
||||
Use a fresh reseed timestamp for each seed record, not the original source
|
||||
result timestamp. This is required because
|
||||
`load_records_for_model(..., last_n=5)` keeps the last records after loading
|
||||
@@ -357,7 +482,8 @@ Prefer uploading new accepted seed records so failed history remains visible.
|
||||
Print:
|
||||
|
||||
- Backup directory path under `/tmp`.
|
||||
- Prepared local record paths under `PERFORMANCE_TRACKING_ROOT`.
|
||||
- Prepared local record paths under `PERFORMANCE_RESEED_STAGING_ROOT`.
|
||||
- Prepared upload-manifest path under the identity reservation.
|
||||
- HF paths that will receive the new records.
|
||||
- Old rolling medians.
|
||||
- Source batch medians, source batch spread, reseed count, and candidate
|
||||
@@ -369,22 +495,36 @@ prepared records plus backup on disk.
|
||||
|
||||
### 7. Upload only the scoped records
|
||||
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
For a first v2 calibration seed, use the manifest uploader after the user
|
||||
replies exactly `upload`:
|
||||
|
||||
```python
|
||||
from hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```bash
|
||||
python -c 'from fastvideo.tests.performance.seed_baseline import upload_prepared_seed_manifest; print(upload_prepared_seed_manifest("<prepared_manifest>"))'
|
||||
```
|
||||
|
||||
Run it once per prepared record. Each upload goes to:
|
||||
The uploader verifies the source and prepared-record digests, pins and scans
|
||||
the current HF revision, rechecks exact-identity and recipe-cohort conflicts,
|
||||
and writes the entire batch in one commit whose `parent_commit` must still be
|
||||
current. A concurrent Hub update makes the commit fail. Do not retry
|
||||
automatically: preserve staging, refresh/review remote state, and request a new
|
||||
explicit `upload` after the conflict is understood. Each record goes to:
|
||||
|
||||
```text
|
||||
FastVideo/performance-tracking/<sanitize(model_id)>/<record_filename>.json
|
||||
```
|
||||
|
||||
Never bulk upload the whole tracking root. Never modify another model's
|
||||
directory in the same operation.
|
||||
Never call `upload_record()` once per first-seed record: that can partially
|
||||
land the batch and has no compare-and-swap guard.
|
||||
|
||||
For a legacy reseed or an accepted v2 baseline shift, the first-seed manifest
|
||||
validator does not apply because an eligible baseline already exists. Upload
|
||||
only the individually reviewed records prepared in step 5 with the shared
|
||||
`upload_record(local_path, record, strict=True)` helper. Stop on the first
|
||||
failure and report exactly which records reached HF; do not silently rerun or
|
||||
replicate the remainder.
|
||||
|
||||
Never bulk upload the tracking or staging root, and never modify another
|
||||
model's directory in the same operation.
|
||||
|
||||
### 8. Report outcome and offer cleanup
|
||||
|
||||
@@ -406,9 +546,14 @@ distinguish an accepted baseline shift from a hidden regression.
|
||||
After the upload is verified, ask whether the user wants to clear temporary
|
||||
local state. Explain what each directory is for:
|
||||
|
||||
- `PERFORMANCE_TRACKING_ROOT`, usually `/tmp/perf-tracking`: local synced
|
||||
mirror of `FastVideo/performance-tracking` plus the prepared local seed
|
||||
records used for scoped upload.
|
||||
- `PERFORMANCE_TRACKING_ROOT`, usually `/tmp/perf-tracking`: read-only local
|
||||
synced mirror used for operator review and reporting. First-v2 preparation
|
||||
independently proves remote state from a fresh temporary HF snapshot.
|
||||
- `PERFORMANCE_RESEED_STAGING_ROOT`, usually
|
||||
`/tmp/performance_reseed_prepared`: prepared local seed records used for the
|
||||
scoped upload, plus the identity reservation and digest manifest. Keeping
|
||||
this separate prevents aborted preparations from appearing in later
|
||||
baseline reads.
|
||||
- `/tmp/performance_reseed_backup/<...>`: local backup of the target model's
|
||||
pre-reseed HF history plus `PROVENANCE.txt`, kept so a bad reseed can be
|
||||
audited or corrected.
|
||||
@@ -418,14 +563,18 @@ local state. Explain what each directory is for:
|
||||
Ask:
|
||||
|
||||
> Reseed succeeded. Do you want me to delete the local temp tracking mirror,
|
||||
> source downloads, and reseed backup under `/tmp`? These files are local
|
||||
> safety/audit artifacts only; HF already has the uploaded records.
|
||||
> this reseed's prepared staging records, source downloads, and reseed backup
|
||||
> under `/tmp`? These files are local safety/audit artifacts only; HF already
|
||||
> has the uploaded records.
|
||||
>
|
||||
> Reply `cleanup reseed temp` to delete them, anything else to keep them.
|
||||
|
||||
Do not delete anything unless the user replies exactly
|
||||
`cleanup reseed temp`. If cleanup is requested, remove only the specific
|
||||
directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
directories and prepared record paths created for this reseed. Do not remove
|
||||
the shared staging root when it contains other records. Remove this operation's
|
||||
identity reservation only with its prepared records and manifest, and never
|
||||
remove unrelated `/tmp` contents.
|
||||
|
||||
## Failure modes and handling
|
||||
|
||||
@@ -437,19 +586,34 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
against the source batch median by more than `max_intra_batch_regression`.
|
||||
Ask for cleaner sources or a reviewed explanation before continuing.
|
||||
- **Too few source records to move the median.** Continue only after making
|
||||
clear that one or two records may not immediately move the last-5 median.
|
||||
clear that one or two records may not immediately move an existing last-5
|
||||
median. This warning does not block a first v2 calibration seed for an exact
|
||||
identity with no eligible baseline yet.
|
||||
- **The source results are noisy or suspicious.** Stop. Reseeding amplifies
|
||||
those measurements into the baseline, so they must be reviewed first.
|
||||
- **HF sync fails.** Stop for destructive reseeds. A stale or empty sync can
|
||||
make the old baseline look missing.
|
||||
- **The exact v2 identity already has an eligible baseline.** Stop. The
|
||||
`CALIBRATION_NEEDED` artifact is stale; use the reviewed baseline-shift path
|
||||
instead of the first-seed utility.
|
||||
- **The workload/variant/version trusts another recipe.** Stop. The source is
|
||||
stale relative to the current recipe cohort and must not bypass
|
||||
`RECIPE_MISMATCH` by creating a second trusted recipe.
|
||||
- **The staging root already has a prepared seed for the exact identity.**
|
||||
Stop and reuse, upload, or explicitly clean that preparation. Do not prepare
|
||||
another copy of the same measurement.
|
||||
- **The conditional Hub commit loses its parent race.** Stop without retrying.
|
||||
Keep the preparation, refresh and review the new remote state, then request
|
||||
a new explicit `upload` only if the seed is still valid.
|
||||
- **Candidate still violates fixed thresholds.** Report that this skill only
|
||||
handles the rolling HF baseline; update benchmark JSON thresholds in code
|
||||
review if maintainers accept the new absolute limit.
|
||||
- **The user aborts at either confirmation.** Leave the backup and prepared
|
||||
records on disk. Nothing should be uploaded.
|
||||
- **The user declines cleanup.** Keep `/tmp/perf-tracking`, the source
|
||||
download directory if any, and `/tmp/performance_reseed_backup/<...>` in
|
||||
place for audit/debugging.
|
||||
- **The user declines cleanup.** Keep `/tmp/perf-tracking`, the prepared seed
|
||||
records under `/tmp/performance_reseed_prepared`, the source download
|
||||
directory if any, and `/tmp/performance_reseed_backup/<...>` in place for
|
||||
audit/debugging.
|
||||
- **A bad seed was uploaded.** Use the backup and HF history to identify the
|
||||
uploaded file, then remove or supersede it with an explicitly reviewed
|
||||
corrective record. Do not silently rewrite unrelated history.
|
||||
@@ -460,8 +624,9 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/performance/hf_store.py` — HF sync and record loading helpers.
|
||||
- `fastvideo/tests/performance/seed_baseline.py` — first-seed preparation,
|
||||
staging reservation, manifest validation, and conditional batch upload.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
- `.buildkite/performance-benchmarks/tests/*.json` — fixed absolute benchmark
|
||||
@@ -474,3 +639,4 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
| 2026-05-03 | Initial version. Sister workflow to `reseed-ssim-references`, scoped to one performance `(model_id, gpu_type)` baseline seed with backup, confirmation, provenance, and `success=true` upload. |
|
||||
| 2026-05-03 | Previous policy: replicate one approved shifted source result into 3 success records by default, or 5 only when explicitly requested. Add provenance marker for replicated-source reseeds. Superseded by the 2026-05-08 dynamic multi-source policy. |
|
||||
| 2026-05-08 | Replace fixed 3/5 replication with dynamic multi-source reseeding: upload one seed record per reviewed source JSON, validate intra-batch consistency, move backup/source scratch under `/tmp`, and ask whether to clean temp state after successful upload. |
|
||||
| 2026-07-13 | Keep first-v2-seed preparation outside the canonical mirror, reserve staging identities atomically, reject stale or replayed calibration seeds, and upload reviewed manifests with a single parent-guarded Hub commit. |
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
+28
-1
@@ -114,6 +114,17 @@ steps:
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Extraction Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
@@ -371,6 +382,21 @@ steps:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "scripts/lora_extraction/**"
|
||||
- "fastvideo/tests/lora_extraction/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/training/training_utils.py"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Extraction Tests"
|
||||
env:
|
||||
- TEST_TYPE=lora_extraction
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
@@ -410,7 +436,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
@@ -455,6 +481,7 @@ steps:
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/worker/**"
|
||||
- "fastvideo/entrypoints/**"
|
||||
- "fastvideo/performance/**"
|
||||
- "fastvideo/tests/performance/**"
|
||||
- ".buildkite/performance-benchmarks/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -76,10 +76,27 @@ 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:-} BUILDKITE_BUILD_URL=${BUILDKITE_BUILD_URL:-} BUILDKITE_BUILD_ID=${BUILDKITE_BUILD_ID:-} BUILDKITE_JOB_ID=${BUILDKITE_JOB_ID:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} BUILDKITE_SOURCE=${BUILDKITE_SOURCE:-} TEST_SCOPE=${TEST_SCOPE:-} BUILDKITE_BUILD_URL=${BUILDKITE_BUILD_URL:-} BUILDKITE_BUILD_ID=${BUILDKITE_BUILD_ID:-} BUILDKITE_JOB_ID=${BUILDKITE_JOB_ID:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
is_truthy() {
|
||||
case "${1:-}" in
|
||||
1|true|TRUE|yes|YES|on|ON) return 0 ;;
|
||||
*) return 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
ssim_bootstrap_args() {
|
||||
local title="${PR_TITLE:-}"
|
||||
local message="${BUILDKITE_MESSAGE:-}"
|
||||
if is_truthy "${FASTVIDEO_SSIM_BOOTSTRAP_MODE:-}" \
|
||||
|| [[ "$title" == *"[new-model]"* ]] \
|
||||
|| [[ "$message" == *"[new-model]"* ]]; then
|
||||
printf ' --bootstrap-mode'
|
||||
fi
|
||||
}
|
||||
|
||||
upload_performance_artifacts() {
|
||||
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
|
||||
LOCAL_DIR="downloaded_reports"
|
||||
@@ -172,7 +189,12 @@ case "$TEST_TYPE" in
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_SSIM_TEST_FILE::run_ssim_tests"
|
||||
SSIM_BOOTSTRAP_ARGS=$(ssim_bootstrap_args)
|
||||
if [ -n "$SSIM_BOOTSTRAP_ARGS" ]; then
|
||||
log "SSIM bootstrap mode enabled for new-model reference draft generation"
|
||||
fi
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run "
|
||||
MODAL_COMMAND+="$MODAL_SSIM_TEST_FILE::run_ssim_tests$SSIM_BOOTSTRAP_ARGS"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
|
||||
Executable
+133
@@ -0,0 +1,133 @@
|
||||
#!/usr/bin/env bash
|
||||
# Gate the expensive Buildkite full suite on the cheap GitHub checks.
|
||||
#
|
||||
# Polls the workflow runs for the PR head commit and only exits 0 once the
|
||||
# watched cheap workflows (pre-commit, docs build) have succeeded, so the
|
||||
# 'ready' label cannot burn ~20 GPU lanes on a head that a cheap check has
|
||||
# already doomed.
|
||||
#
|
||||
# Semantics:
|
||||
# - watched run completed with a bad conclusion -> exit 1 (fail CLOSED:
|
||||
# no full suite; the next push re-arms via the 'synchronize' trigger)
|
||||
# - watched run cancelled -> still pending: the docs
|
||||
# workflow's repo-global 'pages' concurrency group cancels runs superseded
|
||||
# by unrelated pushes, so 'cancelled' is not a verdict on this PR
|
||||
# - watched runs pending -> poll until done
|
||||
# - docs run absent -> not applicable after a
|
||||
# short grace period ('Deploy Documentation' is path-filtered on PRs)
|
||||
# - pre-commit run absent -> keep polling: pre-commit
|
||||
# is never path-filtered, so its absence is always anomalous
|
||||
# - 'ready' label removed while waiting -> exit 1 (fail CLOSED:
|
||||
# un-labeling is a deliberate maintainer action)
|
||||
# - GitHub API unreachable or timeout -> exit 0 (fail OPEN,
|
||||
# loud warning: never brick CI on a GitHub outage)
|
||||
#
|
||||
# Required env: PR_SHA (PR head commit), PR_NUMBER, GITHUB_REPOSITORY, GH_TOKEN.
|
||||
set -euo pipefail
|
||||
|
||||
: "${PR_SHA:?PR_SHA (PR head commit) is required}"
|
||||
: "${PR_NUMBER:?PR_NUMBER (pull request number) is required}"
|
||||
: "${GITHUB_REPOSITORY:?GITHUB_REPOSITORY is required}"
|
||||
|
||||
# Workflow-level `name:` values that must be green before the full suite
|
||||
# may start. "Deploy Documentation" is path-filtered on PRs, so its run may
|
||||
# legitimately never exist; pre-commit always runs, so it must appear.
|
||||
WATCHED_NAMES='["pre-commit", "Deploy Documentation"]'
|
||||
WATCHED_REGEX='^(pre-commit|Deploy Documentation)$'
|
||||
POLL_SECS="${POLL_SECS:-20}"
|
||||
GRACE_SECS="${GRACE_SECS:-60}"
|
||||
MAX_WAIT_SECS="${MAX_WAIT_SECS:-1500}"
|
||||
|
||||
# Bound each API call so a hung connection hits the 3-strike fail-open path
|
||||
# instead of pinning the loop until the job timeout (which would fail closed
|
||||
# on exactly the GitHub-outage case this script is meant to survive).
|
||||
if command -v timeout >/dev/null 2>&1; then
|
||||
gh_api() { timeout 30 gh api "$@"; }
|
||||
else
|
||||
gh_api() { gh api "$@"; } # macOS dev boxes; CI always has coreutils timeout
|
||||
fi
|
||||
|
||||
# The workflow checked the label before starting the gate, but the wait can
|
||||
# last ~25 min: re-check once before any exit 0 and fail closed if 'ready'
|
||||
# was removed in the meantime. An API error here proceeds (the label was
|
||||
# present when the gate started; never brick CI on an outage).
|
||||
recheck_ready_label() {
|
||||
local pr_json
|
||||
if pr_json=$(gh_api "repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" 2>/dev/null); then
|
||||
if ! jq -e '[.labels[]?.name] | index("ready")' <<<"$pr_json" >/dev/null 2>&1; then
|
||||
echo "::error::PR #${PR_NUMBER} no longer has the 'ready' label —" \
|
||||
"NOT triggering the Buildkite full suite. Re-add the label to re-arm."
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
echo "::warning::Could not re-check the 'ready' label on PR #${PR_NUMBER}; proceeding (it was present when the gate started)."
|
||||
fi
|
||||
}
|
||||
|
||||
start=$(date +%s)
|
||||
api_fails=0
|
||||
missing=""
|
||||
|
||||
while true; do
|
||||
elapsed=$(( $(date +%s) - start ))
|
||||
|
||||
if runs_json=$(gh_api "repos/${GITHUB_REPOSITORY}/actions/runs?head_sha=${PR_SHA}&per_page=100" 2>/dev/null) \
|
||||
&& state=$(jq --arg re "$WATCHED_REGEX" '
|
||||
[.workflow_runs[]? | select(.name // "" | test($re))]
|
||||
| group_by(.name) | map(max_by(.id))
|
||||
| map({name, status, conclusion})' <<<"$runs_json" 2>/dev/null); then
|
||||
api_fails=0
|
||||
echo "t+${elapsed}s watched checks: $(jq -c . <<<"$state")"
|
||||
|
||||
failed=$(jq -r '[.[] | select(.status == "completed"
|
||||
and (.conclusion | IN("success", "skipped", "neutral", "cancelled") | not))]
|
||||
| map(.name) | join(", ")' <<<"$state")
|
||||
if [ -n "$failed" ]; then
|
||||
echo "::error::Cheap check(s) failed on ${PR_SHA}: ${failed}." \
|
||||
"NOT triggering the Buildkite full suite. Push a fix (the 'ready'" \
|
||||
"label re-arms on every push), or re-run the failed check and then" \
|
||||
"re-run this workflow."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 'cancelled' counts as pending: wait for a re-run to reach a real verdict
|
||||
# (bounded by MAX_WAIT, then the fail-open below).
|
||||
pending=$(jq '[.[] | select(.status != "completed" or .conclusion == "cancelled")] | length' <<<"$state")
|
||||
missing=$(jq -r --argjson watched "$WATCHED_NAMES" '($watched - map(.name)) | join(", ")' <<<"$state")
|
||||
if [ "$pending" -eq 0 ]; then
|
||||
if [ -z "$missing" ]; then
|
||||
recheck_ready_label
|
||||
echo "All watched cheap checks are green — full suite may proceed."
|
||||
exit 0
|
||||
fi
|
||||
case "$missing" in
|
||||
*pre-commit*)
|
||||
echo "pre-commit run not found for ${PR_SHA} yet; waiting (pre-commit is never path-filtered, so its absence is anomalous)."
|
||||
;;
|
||||
*)
|
||||
if [ "$elapsed" -ge "$GRACE_SECS" ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::Watched run(s) never appeared for ${PR_SHA}: ${missing} (path-filtered, likely not applicable). Proceeding on the checks that did run."
|
||||
exit 0
|
||||
fi
|
||||
echo "Waiting up to ${GRACE_SECS}s grace for path-filtered run(s) to appear: ${missing}."
|
||||
;;
|
||||
esac
|
||||
fi
|
||||
else
|
||||
api_fails=$(( api_fails + 1 ))
|
||||
echo "::warning::GitHub API error querying workflow runs for ${PR_SHA} (attempt ${api_fails}/3)."
|
||||
if [ "$api_fails" -ge 3 ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the full suite WITHOUT the cheap-check gate."
|
||||
exit 0
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$elapsed" -ge "$MAX_WAIT_SECS" ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the full suite anyway."
|
||||
exit 0
|
||||
fi
|
||||
sleep "$POLL_SECS"
|
||||
done
|
||||
Executable
+122
@@ -0,0 +1,122 @@
|
||||
#!/usr/bin/env bash
|
||||
# Self-test for gate_full_suite.sh using a mocked `gh`. No network, runs on
|
||||
# any dev box: bash .github/scripts/test_gate_full_suite.sh
|
||||
set -u
|
||||
here=$(cd "$(dirname "$0")" && pwd)
|
||||
tmp=$(mktemp -d)
|
||||
trap 'rm -rf "$tmp"' EXIT
|
||||
|
||||
# Mock gh. Asserts the exact endpoint (including head_sha) it is called
|
||||
# with — an endpoint typo in the gate script fails the test rather than
|
||||
# silently serving canned data. On the runs endpoint it serves
|
||||
# $MOCK_DIR/response_<call#>.json, sticking on the highest existing file,
|
||||
# and exits 1 if none exist (simulates a GitHub API outage). On the pulls
|
||||
# endpoint it serves $MOCK_DIR/pr.json, defaulting to a 'ready'-labeled PR.
|
||||
cat > "$tmp/gh" <<'EOF'
|
||||
#!/usr/bin/env bash
|
||||
if [ "${1:-}" != "api" ]; then
|
||||
echo "unexpected gh invocation: $*" >> "$MOCK_DIR/endpoint_error"
|
||||
exit 2
|
||||
fi
|
||||
case "${2:-}" in
|
||||
"repos/o/r/actions/runs?head_sha=deadbeef&per_page=100")
|
||||
n=$(( $(cat "$MOCK_DIR/count" 2>/dev/null || echo 0) + 1 ))
|
||||
echo "$n" > "$MOCK_DIR/count"
|
||||
while [ "$n" -gt 0 ]; do
|
||||
if [ -f "$MOCK_DIR/response_$n.json" ]; then
|
||||
cat "$MOCK_DIR/response_$n.json"
|
||||
exit 0
|
||||
fi
|
||||
n=$(( n - 1 ))
|
||||
done
|
||||
echo "api outage" >&2
|
||||
exit 1
|
||||
;;
|
||||
"repos/o/r/pulls/42")
|
||||
if [ -f "$MOCK_DIR/pr.json" ]; then
|
||||
cat "$MOCK_DIR/pr.json"
|
||||
else
|
||||
echo '{"labels": [{"name": "ready"}]}'
|
||||
fi
|
||||
;;
|
||||
*)
|
||||
echo "unexpected gh endpoint: $2" >> "$MOCK_DIR/endpoint_error"
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
EOF
|
||||
chmod +x "$tmp/gh"
|
||||
|
||||
PC_OK='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "success"}'
|
||||
PC_BAD='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "failure"}'
|
||||
PC_PENDING='{"name": "pre-commit", "id": 1, "status": "in_progress", "conclusion": null}'
|
||||
DOCS_OK='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "success"}'
|
||||
DOCS_BAD='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "failure"}'
|
||||
DOCS_CANCELLED='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "cancelled"}'
|
||||
OTHER='{"name": "Trigger Full Suite", "id": 3, "status": "in_progress", "conclusion": null}'
|
||||
NULL_NAME='{"name": null, "id": 4, "status": "completed", "conclusion": "failure"}'
|
||||
PC_OK_RERUN='{"name": "pre-commit", "id": 5, "status": "completed", "conclusion": "success"}'
|
||||
|
||||
fails=0
|
||||
want_log="" # optional: expect() also greps out.log for this regex, then resets
|
||||
pr_json="" # optional: served for the pulls (label re-check) endpoint, then resets
|
||||
raw_body="" # optional: serve responses verbatim instead of wrapping in workflow_runs
|
||||
expect() { # <name> <expected-exit> <response json>...
|
||||
local name=$1 want=$2 dir i=1
|
||||
shift 2
|
||||
dir=$(mktemp -d "$tmp/test_XXXXXX")
|
||||
for body in "$@"; do
|
||||
if [ -n "$raw_body" ]; then
|
||||
printf '%s' "$body" > "$dir/response_$i.json"
|
||||
else
|
||||
printf '{"workflow_runs": [%s]}' "$body" > "$dir/response_$i.json"
|
||||
fi
|
||||
i=$(( i + 1 ))
|
||||
done
|
||||
[ -n "$pr_json" ] && printf '%s' "$pr_json" > "$dir/pr.json"
|
||||
( export PATH="$tmp:$PATH" MOCK_DIR="$dir" PR_SHA=deadbeef PR_NUMBER=42 \
|
||||
GITHUB_REPOSITORY=o/r POLL_SECS=0 GRACE_SECS=1 MAX_WAIT_SECS=3
|
||||
bash "$here/gate_full_suite.sh" > "$dir/out.log" 2>&1 )
|
||||
local rc=$?
|
||||
if [ "$rc" -ne "$want" ]; then
|
||||
echo "FAIL: $name (exit $rc, want $want)"
|
||||
cat "$dir/out.log"
|
||||
fails=1
|
||||
elif [ -f "$dir/endpoint_error" ]; then
|
||||
echo "FAIL: $name (mock gh got an unexpected call)"
|
||||
cat "$dir/endpoint_error"
|
||||
fails=1
|
||||
elif [ -n "$want_log" ] && ! grep -Eq "$want_log" "$dir/out.log"; then
|
||||
echo "FAIL: $name (log does not match: $want_log)"
|
||||
cat "$dir/out.log"
|
||||
fails=1
|
||||
else
|
||||
echo "ok: $name"
|
||||
fi
|
||||
want_log="" pr_json="" raw_body=""
|
||||
}
|
||||
|
||||
expect "both green -> proceed" 0 "$PC_OK, $DOCS_OK, $OTHER, $NULL_NAME"
|
||||
expect "docs build failed -> blocked" 1 "$PC_OK, $DOCS_BAD"
|
||||
expect "pre-commit failed -> blocked" 1 "$PC_BAD"
|
||||
expect "pending then green -> proceed" 0 "$PC_PENDING" "$PC_OK, $DOCS_OK"
|
||||
want_log="never appeared.*Deploy Documentation"
|
||||
expect "docs run absent (path-filtered) -> proceed after grace" 0 "$PC_OK"
|
||||
expect "API outage -> fail open" 0
|
||||
want_log="FAILING OPEN"
|
||||
expect "pending past MAX_WAIT -> fail open" 0 "$PC_PENDING"
|
||||
want_log="FAILING OPEN"
|
||||
expect "unrelated runs only -> no grace, fail open at MAX_WAIT" 0 "$OTHER"
|
||||
expect "cancelled docs then green -> proceed" 0 \
|
||||
"$PC_OK, $DOCS_CANCELLED" "$PC_OK, $DOCS_OK"
|
||||
want_log="FAILING OPEN"
|
||||
expect "cancelled docs forever -> fail open at MAX_WAIT" 0 "$PC_OK, $DOCS_CANCELLED"
|
||||
want_log="FAILING OPEN"
|
||||
expect "pre-commit absent -> no grace, fail open at MAX_WAIT" 0 "$DOCS_OK"
|
||||
expect "duplicate run names -> latest wins" 0 "$PC_BAD, $PC_OK_RERUN, $DOCS_OK"
|
||||
raw_body=1
|
||||
expect "garbage response body -> fail open" 0 "this is not json"
|
||||
pr_json='{"labels": [{"name": "other"}]}'
|
||||
expect "ready label removed mid-gate -> blocked" 1 "$PC_OK, $DOCS_OK"
|
||||
|
||||
exit "$fails"
|
||||
@@ -1,7 +1,11 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
# pull_request_target instead of pull_request: the workflow definition and
|
||||
# the hook config are always taken from the BASE branch, so fork /
|
||||
# first-time-contributor PRs run immediately without a maintainer clicking
|
||||
# "Approve and run". The PR head is checked out as data only.
|
||||
pull_request_target:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
inputs:
|
||||
@@ -15,12 +19,25 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
if: github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
@@ -30,3 +47,6 @@ jobs:
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
# After pre-commit so a self-test failure cannot mask lint failures.
|
||||
- name: Full-suite gate self-test
|
||||
run: bash .github/scripts/test_gate_full_suite.sh
|
||||
|
||||
@@ -52,6 +52,7 @@ jobs:
|
||||
core.setOutput('pr_sha', pr.head.sha);
|
||||
core.setOutput('pr_branch', pr.head.ref);
|
||||
core.setOutput('pr_number', String(prNumber));
|
||||
core.setOutput('pr_title', pr.title);
|
||||
|
||||
- name: Trigger Full Suite
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
@@ -60,6 +61,7 @@ jobs:
|
||||
PR_SHA: ${{ steps.label.outputs.pr_sha }}
|
||||
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
|
||||
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
|
||||
PR_TITLE: ${{ steps.label.outputs.pr_title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -71,6 +73,7 @@ jobs:
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
@@ -80,11 +83,12 @@ jobs:
|
||||
pull_request_id: $pr_id,
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
}
|
||||
}')"
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring),
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
|
||||
parse-command:
|
||||
if: >-
|
||||
@@ -125,7 +129,7 @@ jobs:
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
@@ -136,6 +140,7 @@ jobs:
|
||||
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
@@ -240,6 +245,7 @@ jobs:
|
||||
TEST_SCOPE: ${{ needs.parse-command.outputs.test_scope }}
|
||||
FULL_SUITE: ${{ needs.parse-command.outputs.full_suite }}
|
||||
TEST_TYPE: ${{ needs.parse-command.outputs.test_type }}
|
||||
PR_TITLE: ${{ github.event.issue.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -256,6 +262,7 @@ jobs:
|
||||
--arg full_suite "$FULL_SUITE" \
|
||||
--arg test_type "$TEST_TYPE" \
|
||||
--arg pr_number "$PR_NUMBER" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
'{
|
||||
commit: $commit,
|
||||
branch: $branch,
|
||||
@@ -265,8 +272,9 @@ jobs:
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: $test_scope,
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number
|
||||
}
|
||||
}')"
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number,
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
|
||||
@@ -7,6 +7,7 @@ on:
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
actions: read
|
||||
|
||||
concurrency:
|
||||
group: full-suite-${{ github.event.pull_request.number }}
|
||||
@@ -18,6 +19,8 @@ jobs:
|
||||
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|
||||
|| github.event.action == 'synchronize'
|
||||
runs-on: ubuntu-latest
|
||||
# Gate below may wait for cheap checks (up to MAX_WAIT_SECS = 25 min).
|
||||
timeout-minutes: 35
|
||||
steps:
|
||||
- name: Check ready label
|
||||
id: check
|
||||
@@ -49,6 +52,20 @@ jobs:
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
|
||||
done
|
||||
|
||||
# Checks out the BASE branch (default for pull_request_target), so PR
|
||||
# authors cannot tamper with the gate script.
|
||||
- name: Checkout gate script
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
|
||||
|
||||
- name: Wait for pre-commit and docs build
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
run: bash .github/scripts/gate_full_suite.sh
|
||||
|
||||
- name: Trigger Buildkite Full Suite
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
@@ -56,6 +73,7 @@ jobs:
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_TITLE: ${{ github.event.pull_request.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -67,6 +85,7 @@ jobs:
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER}" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
@@ -78,6 +97,7 @@ jobs:
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
PR_NUMBER: ($pr_id | tostring),
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
|
||||
@@ -13,12 +13,33 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
|
||||
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
|
||||
# push trigger is a sufficient change detector on its own -- no separate
|
||||
# detect-changes/paths-filter job is needed now that there is a single
|
||||
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
|
||||
# rocm Dockerfile stay manual-dispatch only.
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'docker/Dockerfile'
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
# One static group, no cancellation: every run of this workflow writes the same
|
||||
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
|
||||
# concurrent push/dispatch runs would race on those tags, and cancelling a run
|
||||
# mid-publish can strand the cu126/cu130 tag families at different commits. An
|
||||
# in-flight superseded build wastes its runner time, but its tags are then
|
||||
# overwritten by the newer queued run. GitHub keeps a single pending run per
|
||||
# group: the newest queued run replaces any older queued one.
|
||||
concurrency:
|
||||
group: infra-build-image
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
|
||||
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
|
||||
@@ -28,7 +49,11 @@ jobs:
|
||||
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
|
||||
# 2.8.3 comes from the architecture-specific prebuilt releases.
|
||||
build-cuda-images:
|
||||
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
|
||||
# on a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# repository guard keeps fork syncs from auto-building; manual dispatch
|
||||
# still works in forks.
|
||||
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -75,10 +100,10 @@ jobs:
|
||||
secrets: inherit
|
||||
|
||||
publish-cuda-manifests:
|
||||
# !cancelled(): a failed sibling build leg must not skip the manifests for a
|
||||
# CUDA lane whose own digests all exist; the digest-count check below fails
|
||||
# the incomplete lane loudly instead.
|
||||
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# !cancelled(): publish lanes whose digests exist even if a sibling build
|
||||
# leg failed (the digest-count check fails incomplete lanes); it also
|
||||
# bypasses skipped-needs propagation, hence the explicit skipped check.
|
||||
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
|
||||
needs: build-cuda-images
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
|
||||
+3
-2
@@ -72,8 +72,7 @@ docs/distillation/examples/
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
# Reference videos (negations must come after the catch-all on line below)
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
@@ -127,6 +126,8 @@ apps/dreamverse/web/.env.production.local
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.png
|
||||
|
||||
# Editor logs and local Python version pins (accidentally committed)
|
||||
*.nvimlog
|
||||
|
||||
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
|
||||
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
|
||||
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
|
||||
|
||||
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
|
||||
# the dreamverse extra and is validated with it, so enable it here.
|
||||
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
|
||||
STREAM_MODE=av_fmp4 \
|
||||
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
|
||||
FASTVIDEO_FA4=1 \
|
||||
HF_HOME=/root/.cache/huggingface
|
||||
|
||||
RUN mkdir -p /var/lib/dreamverse
|
||||
|
||||
@@ -347,32 +347,6 @@ def test_rewrite_prompt_sequence_accepts_numbered_prose_output():
|
||||
]
|
||||
|
||||
|
||||
def test_enhance_prompt_prefers_cerebras_before_groq_fallback():
|
||||
enhancer = _build_staged_enhancer(
|
||||
cerebras_payload=_chat_payload_with_content('{"prompt":"Cerebras prompt"}'),
|
||||
groq_payload=_chat_payload_with_content('{"prompt":"Groq prompt"}'),
|
||||
cerebras_delay_s=0.01,
|
||||
groq_delay_s=0.01,
|
||||
)
|
||||
|
||||
result = asyncio.run(
|
||||
enhancer.enhance_prompt(
|
||||
"A rainy alley at night",
|
||||
mode="single_clip",
|
||||
)
|
||||
)
|
||||
|
||||
assert result.fallback_used is False
|
||||
assert result.error is None
|
||||
assert result.provider == "cerebras"
|
||||
assert result.model == "gpt-test"
|
||||
assert result.prompt == "Cerebras prompt"
|
||||
assert enhancer.get_provider_success_counts() == {
|
||||
"cerebras": 1,
|
||||
"groq": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_enhance_prompt_uses_groq_when_cerebras_fails():
|
||||
enhancer = _build_staged_enhancer(
|
||||
cerebras_payload=_chat_payload_with_content("{}"),
|
||||
|
||||
@@ -12,13 +12,17 @@ Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata:
|
||||
Records can include source metadata and rolling-baseline policy context:
|
||||
|
||||
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
|
||||
- `baseline_eligible`: only successful scheduled-main records should be true
|
||||
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
|
||||
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
|
||||
floors used for recomputed status context
|
||||
|
||||
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
|
||||
crossings; `regressed` remains the gated CI-failure signal.
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
- PR number, branch, and Buildkite URL when present
|
||||
|
||||
The latest status table uses the stored JSON `success` value. Recomputed
|
||||
baseline context is shown separately and does not override stored status.
|
||||
baseline context applies each metric's percent and absolute regression floors
|
||||
and does not override stored status.
|
||||
|
||||
## API
|
||||
|
||||
@@ -99,6 +104,11 @@ baseline context is shown separately and does not override stored status.
|
||||
- `GET /api/performance/trends?days=90&run_source=scheduled_main`
|
||||
- `GET /api/performance/records?days=90&run_source=local`
|
||||
|
||||
The current v1 grouping key is `(model_id, gpu_type)`. Baselines are computed
|
||||
from the latest five previous successful records in each group for dashboard
|
||||
context. CI gating uses only records marked `baseline_eligible=true`.
|
||||
V2 records use the same comparison cohort as CI: `workload_id`, `variant_id`,
|
||||
`benchmark_version`, `recipe_fingerprint`, `hardware_profile_id`, and
|
||||
`software_profile_id`. `model_id` and `gpu_type` remain display/filter
|
||||
metadata, so renaming either does not split history. Legacy records still group
|
||||
by `(model_id, gpu_type)`. Dashboard baselines use the latest five previous
|
||||
successful, baseline-eligible records in each group. Summary and trend filters
|
||||
match the latest display metadata after grouping, while the raw records endpoint
|
||||
continues to filter individual records.
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
||||
import { fetchSummary, fetchTrends, refreshData, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
import { fetchSummary, fetchTrends, refreshData } from "./api";
|
||||
import type { CohortValue, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
|
||||
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
|
||||
const RUN_SOURCES: Array<{ value: "" | RunSource; label: string }> = [
|
||||
@@ -109,6 +110,61 @@ function metricLabel(metricKey: string) {
|
||||
return METRIC_DEFINITIONS[metricKey]?.label ?? metricKey;
|
||||
}
|
||||
|
||||
type CohortFields = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
workload_id: CohortValue;
|
||||
variant_id: CohortValue;
|
||||
benchmark_version: CohortValue;
|
||||
recipe_fingerprint: CohortValue;
|
||||
hardware_profile_id: CohortValue;
|
||||
software_profile_id: CohortValue;
|
||||
};
|
||||
|
||||
function cohortValue(value: CohortValue) {
|
||||
if (value === null || value === undefined || value === "") {
|
||||
return "legacy";
|
||||
}
|
||||
return String(value);
|
||||
}
|
||||
|
||||
function shortCohortValue(value: CohortValue) {
|
||||
const text = cohortValue(value);
|
||||
if (text === "legacy" || text.length <= 14) {
|
||||
return text;
|
||||
}
|
||||
return text.slice(0, 12);
|
||||
}
|
||||
|
||||
function cohortKey(cohort: CohortFields) {
|
||||
return [
|
||||
cohort.model_id,
|
||||
cohort.gpu_type,
|
||||
cohortValue(cohort.workload_id),
|
||||
cohortValue(cohort.variant_id),
|
||||
cohortValue(cohort.benchmark_version),
|
||||
cohortValue(cohort.recipe_fingerprint),
|
||||
cohortValue(cohort.hardware_profile_id),
|
||||
cohortValue(cohort.software_profile_id)
|
||||
].join("|");
|
||||
}
|
||||
|
||||
function cohortTitle(cohort: CohortFields) {
|
||||
const workload = cohortValue(cohort.workload_id);
|
||||
const variant = cohortValue(cohort.variant_id);
|
||||
const version = cohortValue(cohort.benchmark_version);
|
||||
const versionLabel = version === "legacy" ? version : `v${version}`;
|
||||
return `${workload} / ${variant} / ${versionLabel}`;
|
||||
}
|
||||
|
||||
function cohortDetail(cohort: CohortFields) {
|
||||
return [
|
||||
`recipe ${shortCohortValue(cohort.recipe_fingerprint)}`,
|
||||
shortCohortValue(cohort.hardware_profile_id),
|
||||
shortCohortValue(cohort.software_profile_id)
|
||||
].join(" | ");
|
||||
}
|
||||
|
||||
function formatMetricValue(metricKey: string, value: number | null | undefined, tooltip = false) {
|
||||
const definition = METRIC_DEFINITIONS[metricKey];
|
||||
if (!definition) {
|
||||
@@ -171,7 +227,9 @@ function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string
|
||||
top: `${(activePoint.y / height) * 100}%`
|
||||
}
|
||||
: undefined;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}`;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}, ${cohortTitle(
|
||||
group
|
||||
)}`;
|
||||
|
||||
return (
|
||||
<div className="chart-shell">
|
||||
@@ -419,7 +477,7 @@ export default function App() {
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Latest Status</h2>
|
||||
<span>{latestRows.length} model/GPU groups</span>
|
||||
<span>{latestRows.length} comparison cohorts</span>
|
||||
</div>
|
||||
{latestRows.length === 0 ? (
|
||||
<div className="empty">No records match the selected filters.</div>
|
||||
@@ -432,6 +490,7 @@ export default function App() {
|
||||
<th>Recomputed</th>
|
||||
<th>Model</th>
|
||||
<th>GPU</th>
|
||||
<th>Cohort</th>
|
||||
<th>Commit</th>
|
||||
<th>Source</th>
|
||||
<th>Baseline</th>
|
||||
@@ -440,11 +499,13 @@ export default function App() {
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
<th>Exceeded</th>
|
||||
<th>Failing</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{latestRows.map((row) => (
|
||||
<tr key={`${row.model_id}-${row.gpu_type}`}>
|
||||
<tr key={cohortKey(row)}>
|
||||
<td>
|
||||
<span className={`badge ${row.status}`}>{row.status}</span>
|
||||
</td>
|
||||
@@ -455,6 +516,12 @@ export default function App() {
|
||||
</td>
|
||||
<td>{row.model_id}</td>
|
||||
<td>{row.gpu_type}</td>
|
||||
<td>
|
||||
<div className="cohort-cell">
|
||||
<strong>{cohortTitle(row)}</strong>
|
||||
<span>{cohortDetail(row)}</span>
|
||||
</div>
|
||||
</td>
|
||||
<td>{shortSha(row.commit_sha)}</td>
|
||||
<td>
|
||||
<span className={`source-badge source-${row.run_source}`}>{runSourceLabel(row.run_source)}</span>
|
||||
@@ -465,6 +532,12 @@ export default function App() {
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
<td>
|
||||
{row.threshold_exceeded_metrics.length
|
||||
? row.threshold_exceeded_metrics.join(", ")
|
||||
: "none"}
|
||||
</td>
|
||||
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
@@ -487,11 +560,13 @@ export default function App() {
|
||||
) : (
|
||||
trends.map((group) =>
|
||||
METRIC_KEYS.map((metricKey) => (
|
||||
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
|
||||
<article className="trend-card" key={`${cohortKey(group)}-${metricKey}`}>
|
||||
<div>
|
||||
<h3>{metricLabel(metricKey)}</h3>
|
||||
<p>
|
||||
{group.model_id} | {group.gpu_type}
|
||||
<span>{cohortTitle(group)}</span>
|
||||
<span>{cohortDetail(group)}</span>
|
||||
</p>
|
||||
</div>
|
||||
<TrendChart group={group} metricKey={metricKey} />
|
||||
|
||||
@@ -2,11 +2,28 @@ export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
absolute_delta: number | null;
|
||||
threshold_percent: number;
|
||||
threshold_absolute: number;
|
||||
gated: boolean;
|
||||
threshold_exceeded: boolean;
|
||||
regressed: boolean;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
};
|
||||
|
||||
export type CohortValue = string | number | null;
|
||||
|
||||
export type ComparisonCohort = {
|
||||
workload_id: CohortValue;
|
||||
variant_id: CohortValue;
|
||||
benchmark_version: CohortValue;
|
||||
recipe_fingerprint: CohortValue;
|
||||
hardware_profile_id: CohortValue;
|
||||
software_profile_id: CohortValue;
|
||||
};
|
||||
|
||||
export type SummaryRow = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
@@ -15,7 +32,8 @@ export type SummaryRow = {
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
regression_threshold_pct: number;
|
||||
threshold_exceeded_metrics: string[];
|
||||
failing_metrics: string[];
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
@@ -27,7 +45,7 @@ export type SummaryRow = {
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, MetricValue>;
|
||||
};
|
||||
} & ComparisonCohort;
|
||||
|
||||
export type RunSource = "pr" | "local" | "scheduled_main" | "unknown";
|
||||
|
||||
@@ -61,13 +79,13 @@ export type TrendPoint = {
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, number | null>;
|
||||
};
|
||||
} & ComparisonCohort;
|
||||
|
||||
export type TrendGroup = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
points: TrendPoint[];
|
||||
};
|
||||
} & ComparisonCohort;
|
||||
|
||||
export type TrendsResponse = {
|
||||
groups: TrendGroup[];
|
||||
|
||||
@@ -149,6 +149,11 @@ h3 {
|
||||
font-size: 0.82rem;
|
||||
}
|
||||
|
||||
.trend-card p {
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.stat strong {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
@@ -186,7 +191,7 @@ h3 {
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
min-width: 1120px;
|
||||
min-width: 1260px;
|
||||
border-collapse: collapse;
|
||||
}
|
||||
|
||||
@@ -209,6 +214,25 @@ td {
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.cohort-cell {
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.cohort-cell strong,
|
||||
.trend-card p span {
|
||||
color: #1b2836;
|
||||
font-size: 0.78rem;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.cohort-cell span,
|
||||
.trend-card p span + span {
|
||||
color: #607080;
|
||||
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", monospace;
|
||||
font-size: 0.72rem;
|
||||
}
|
||||
|
||||
.badge {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
{
|
||||
"alpha_yaw": 0.08734091699186919,
|
||||
"alpha_pitch": 0.08169667696275307,
|
||||
"alpha_turn": 5.724587470723463e-17,
|
||||
"beta_fwd": 0.02842768078408099,
|
||||
"beta_strafe": 0.022531015077067108,
|
||||
"focal_length": 457.0,
|
||||
"frame_shape": [
|
||||
352,
|
||||
640
|
||||
],
|
||||
"calibrated_from": [
|
||||
"1_wasd_only",
|
||||
"camera",
|
||||
"camera4hold_alpha1",
|
||||
"fully_random",
|
||||
"wasdonly_alpha1",
|
||||
"wasd4holdrandview_simple_1key1mouse1"
|
||||
],
|
||||
"residual_rms": 15.890399609478676,
|
||||
"n_equations": 4125232
|
||||
}
|
||||
+6
-4
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
|
||||
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
|
||||
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
|
||||
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
|
||||
ARG FA4_CUTE_REF=940cd9680f3315f2f06b43ab5bea2c2cf2d96806
|
||||
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
|
||||
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
|
||||
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
|
||||
@@ -161,7 +161,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
uv pip install flash-attn==${FLASH_ATTN_VERSION} --no-build-isolation; \
|
||||
fi
|
||||
|
||||
# Overlay the cutlass-4.5-safe upstream FA4 cute (FA4_CUTE_REF) over the
|
||||
# Overlay the CuTe-DSL-4.6-compatible upstream FA4 cute (FA4_CUTE_REF) over the
|
||||
# wheel/source one so the image runs FA4, not the FA2 fallback. This pulls the FA4
|
||||
# runtime stack (cutlass-dsl, quack-kernels, apache-tvm-ffi, torch-c-dlpack-ext) --
|
||||
# the same deps the [dreamverse] extra already installs in CI; the installed torch
|
||||
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
|
||||
# Then verify both survive so a broken overlay fails the build instead of shipping
|
||||
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
|
||||
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
|
||||
# arm64 and the image runs FA3/FA2 as usual.
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
|
||||
else \
|
||||
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
|
||||
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
|
||||
|
||||
@@ -99,10 +99,18 @@ status.
|
||||
Full Suite is also path-filtered. It validates broader behavior before Mergify
|
||||
can merge a PR.
|
||||
|
||||
A `ready`-labeled PR does not hit Buildkite immediately:
|
||||
`ci-trigger-full-suite.yml` first runs `.github/scripts/gate_full_suite.sh`,
|
||||
which waits for the cheap Tier-1 checks (pre-commit, docs build) on the PR
|
||||
head. A red cheap check blocks the suite (fail closed; the next push re-arms
|
||||
it), while a GitHub outage or a >25 min wait lets it run anyway (fail open).
|
||||
`/test full` bypasses the gate.
|
||||
|
||||
| Buildkite label | `TEST_TYPE` | Main watched paths |
|
||||
|---|---|---|
|
||||
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
|
||||
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
|
||||
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
|
||||
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
|
||||
@@ -144,6 +152,7 @@ Valid direct test names:
|
||||
| `/test training` | `training` |
|
||||
| `/test lora-inference` | `inference_lora` |
|
||||
| `/test lora-training` | `training_lora` |
|
||||
| `/test lora-extraction` | `lora_extraction` |
|
||||
| `/test distillation` | `distillation_dmd` |
|
||||
| `/test self-forcing` | `self_forcing` |
|
||||
| `/test vsa` | `training_vsa` |
|
||||
|
||||
@@ -12,7 +12,8 @@ It serves three audiences:
|
||||
* **Maintainers** — surfaces regressions in a Markdown summary on every
|
||||
performance build and a long-form Plotly dashboard.
|
||||
* **Local developers** — lets you run the same benchmark on your own machine,
|
||||
then compare against the historical baseline for the same model and GPU.
|
||||
then compare against the historical baseline for the same comparable
|
||||
identity.
|
||||
|
||||
## Quick start (local)
|
||||
|
||||
@@ -72,14 +73,18 @@ fastvideo/tests/performance/
|
||||
│ writes Markdown summary + (optionally) uploads new records
|
||||
├── dashboard.py
|
||||
│ └── builds time-series Plotly HTML from HF history
|
||||
└── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
|
||||
fastvideo/performance/
|
||||
├── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
└── metric_policy.py # shared rolling-baseline threshold policy
|
||||
```
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
normalized JSON per `(model_id, gpu_type, run)` tuple. The rolling baseline is
|
||||
the median of the last 5 successful, baseline-eligible records for that
|
||||
model+GPU. PR and local records are visible in the dashboard but are not
|
||||
baseline eligible.
|
||||
normalized JSON per run. For v2 records, the rolling baseline is the median of
|
||||
the last 5 successful, baseline-eligible records in the same comparison cohort:
|
||||
`workload_id`, `variant_id`, `benchmark_version`, `recipe_fingerprint`,
|
||||
`hardware_profile_id`, and `software_profile_id`. PR and local records are
|
||||
visible in the dashboard but are not baseline eligible.
|
||||
|
||||
## Planned Coverage
|
||||
|
||||
@@ -92,25 +97,28 @@ and recipe changes instead of treating all records for a model as equivalent.
|
||||
|
||||
## Metrics
|
||||
|
||||
Each benchmark records six metrics:
|
||||
Each benchmark records six metrics. The rolling-baseline comparator also has a
|
||||
per-metric policy with direction, percent threshold, absolute threshold, and a
|
||||
`gated` flag.
|
||||
|
||||
| Metric | Raw key | Normalized key | Direction |
|
||||
|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
|
||||
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|
||||
|---|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
|
||||
|
||||
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
|
||||
while it runs so pipeline stage execution times are available in
|
||||
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
|
||||
`prompt_encoding_stage` so duplicate stage classes do not collide. For
|
||||
`PipelineStage` entries, the extractor maps the `stage_class` field:
|
||||
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
|
||||
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
|
||||
`vae_decode_time_s`, with a fallback for older logs that used the class name as
|
||||
`PipelineStage` entries, shared component stage bases emit a stable
|
||||
`component_metric`: text encoding stages map to `text_encoder_time_s`,
|
||||
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
|
||||
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
|
||||
older logs that do not include `component_metric` or that used the class name as
|
||||
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
|
||||
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
|
||||
does not report one of the mapped stages, that component metric is stored as
|
||||
@@ -152,26 +160,115 @@ unrealistic memory growth, and optionally large component-specific slowdowns
|
||||
even when the rolling baseline is empty. They are hand-set with generous
|
||||
headroom and almost never need touching.
|
||||
|
||||
### Rolling baseline (per `(model_id, gpu_type)`)
|
||||
### Rolling baseline (per comparison cohort)
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
|
||||
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
|
||||
for each available metric, and fails if the current run regresses by more than
|
||||
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
|
||||
higher values are regressions. For throughput, lower values are regressions.
|
||||
for the same comparison cohort from the HF dataset, computes the median for
|
||||
each available metric, and evaluates the current run with the metric's rolling
|
||||
regression policy. For v2 records, that cohort is `workload_id`, `variant_id`,
|
||||
`benchmark_version`, `recipe_fingerprint`, `hardware_profile_id`, and
|
||||
`software_profile_id`. For latency, memory, and component times, higher values
|
||||
are regressions. For throughput, lower values are regressions.
|
||||
|
||||
A metric exceeds its rolling threshold when both of these are true:
|
||||
|
||||
```text
|
||||
percent_delta > threshold_percent
|
||||
absolute_delta > threshold_absolute
|
||||
```
|
||||
|
||||
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
|
||||
for metrics that should remain visible in reports and the dashboard without
|
||||
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
|
||||
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
|
||||
metrics are skipped.
|
||||
|
||||
This is the **drift detector** — it catches sub-threshold regressions that
|
||||
slowly add up. Only scheduled-main successful records are baseline eligible.
|
||||
Local and pull-request runs can upload dashboard-visible records, but they do
|
||||
not update future gating baselines.
|
||||
|
||||
Comparator summaries and normalized artifacts include an explicit
|
||||
`comparison_status`:
|
||||
|
||||
| Status | Meaning | CI behavior |
|
||||
|---|---|---|
|
||||
| `PASS` | Comparable baseline exists and no gated metric regressed. Legacy records with no baseline also keep the historical initialization behavior. | Passes |
|
||||
| `REGRESSION` | The record exceeds one of its static thresholds or at least one gated metric regressed against a comparable baseline. | Fails |
|
||||
| `CALIBRATION_NEEDED` | A v2 record has no exact comparable baseline. | Passes, may upload when `PERF_UPLOAD_POLICY=pass`, never seeds a baseline |
|
||||
| `RECIPE_MISMATCH` | The same workload, variant, and benchmark version has trusted successful records under another recipe fingerprint, including records from other hardware or software profiles. | Fails |
|
||||
| `INFRA_ERROR` | The comparator cannot safely classify the record, such as a v2 record missing required identity fields. | Fails |
|
||||
|
||||
`QUALITY_BLOCKED` is reserved for a future variant-promotion workflow and is
|
||||
not emitted by normal rolling-baseline comparison.
|
||||
|
||||
For `RECIPE_MISMATCH`, trusted records are scheduled-main records or records
|
||||
already marked `baseline_eligible`. PR and local calibration uploads remain
|
||||
visible but do not authorize future CI-gating recipe mismatch failures.
|
||||
|
||||
When the baseline shifts for a legitimate reason (torch upgrade, kernel
|
||||
change, etc.) and CI starts failing, use the
|
||||
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
|
||||
agent skill to advance the rolling median.
|
||||
agent skill to advance the rolling median. To approve the first baseline for
|
||||
a new v2 exact identity, follow that skill with a reviewed scheduled-main
|
||||
full-suite `CALIBRATION_NEEDED` normalized artifact. Its prepare step uses
|
||||
`fastvideo/tests/performance/seed_baseline.py`; after a separate human gate,
|
||||
the skill rechecks the current HF revision and conditionally uploads the whole
|
||||
seed batch in one commit.
|
||||
|
||||
## Schemas
|
||||
|
||||
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
|
||||
|
||||
Benchmark configs without `config_schema_version` are treated as legacy v1
|
||||
configs and remain loadable. New or migrated configs should use
|
||||
`config_schema_version: 2` and include explicit comparable identity fields:
|
||||
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3
|
||||
}
|
||||
```
|
||||
|
||||
`benchmark_id` is still required because raw artifact names, generated-video
|
||||
directories, normalized record paths, and legacy storage directories depend on
|
||||
it. The v2 comparator does not use it as part of the comparison cohort. The v2
|
||||
identity fields make the measured workload explicit:
|
||||
|
||||
| Field | Purpose |
|
||||
|---|---|
|
||||
| `workload_id` | Stable benchmark family, such as `wan-t2v`. |
|
||||
| `variant_id` | Intentional recipe family, including model size and parallelism config, such as `1.3b-sp2`. |
|
||||
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
|
||||
|
||||
If a config declares `config_schema_version: 2`, loading fails clearly when any
|
||||
required v2 identity field is missing. If v2 identity or metadata fields are
|
||||
added without `config_schema_version: 2`, loading also fails so partial
|
||||
migrations do not silently run as v1 configs. Optional v2 `quality_metadata`
|
||||
and the v1/v2 `regression_thresholds` policy must be JSON objects when present.
|
||||
(`recipe` is emitted by the harness and is not config-declarable.)
|
||||
|
||||
V2 records compare only within their exact identity cohort. A record that opens
|
||||
a new cohort is marked `baseline_status: "initialized_new_cohort"` and
|
||||
`comparison_status: "CALIBRATION_NEEDED"`; it remains ineligible until a
|
||||
reviewed scheduled-main artifact is seeded explicitly. Legacy v1 configs still
|
||||
run and are normalized for reporting, but their records skip rolling-baseline
|
||||
comparison entirely (`baseline_status: "skipped_missing_identity"`, never
|
||||
baseline eligible); only static thresholds gate them. Metric-specific threshold
|
||||
policies are active. `QUALITY_BLOCKED` remains reserved for future variant
|
||||
promotion policy.
|
||||
|
||||
The shipped Wan benchmark uses `benchmark_version: 3` because recipe schema 2
|
||||
changed the recipe fingerprint by removing the legacy `benchmark_id` display
|
||||
name. This intentionally opens a new comparison cohort: after deployment, a
|
||||
reviewed scheduled-main full-suite `CALIBRATION_NEEDED` artifact must be seeded
|
||||
once before rolling regression gating resumes for that exact identity. Static
|
||||
thresholds remain active during calibration.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
@@ -179,6 +276,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -196,12 +297,70 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"max_dit_time_s": 10.0,
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 1.0,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"commit": "<full sha>",
|
||||
"run_source": "pr",
|
||||
"branch": "feature/perf-change",
|
||||
"pr_number": "1234",
|
||||
"test_scope": "direct",
|
||||
"build_url": "https://buildkite.example/build",
|
||||
"build_id": "<buildkite-build-id>",
|
||||
"job_id": "<buildkite-job-id>",
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"quality_metadata": { "quality_status": "canonical" },
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208
|
||||
"vae_decode_time_s": 3.208,
|
||||
"recipe": {
|
||||
"recipe_schema_version": 2,
|
||||
"benchmark": {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3
|
||||
},
|
||||
"model": { "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" },
|
||||
"init_kwargs": { "num_gpus": 2, "sp_size": 2, "tp_size": 1 },
|
||||
"generation_kwargs": { "height": 480, "width": 832, "num_frames": 45 },
|
||||
"inputs": { "prompt_count": 1, "prompt_sha256": ["<measured-prompt-sha256>"] },
|
||||
"attention": { "requested_backend": "FLASH_ATTN", "resolved_backend": "FLASH_ATTN" }
|
||||
},
|
||||
"recipe_fingerprint": "<sha256>",
|
||||
"hardware_profile": {
|
||||
"device_type": "cuda",
|
||||
"gpu_count": 2,
|
||||
"gpus": [{ "name": "NVIDIA L40S", "memory_gb": 48, "compute_capability": "8.9" }],
|
||||
"interconnect": "none_or_partial"
|
||||
},
|
||||
"hardware_profile_id": "hw-<sha256-prefix>",
|
||||
"software_profile": {
|
||||
"python": "3.12",
|
||||
"pytorch": "2.12",
|
||||
"cuda": "13.0",
|
||||
"attention_backend": "FLASH_ATTN",
|
||||
"flash_attention_4_enabled": true,
|
||||
"container_image_version": "py3.12-cuda13.0.0",
|
||||
"packages": {
|
||||
"fastvideo_kernel": "0.3.2",
|
||||
"flashinfer": "0.2.11",
|
||||
"nvidia_cutlass_dsl": "4.5.0",
|
||||
"triton": "3.4.1"
|
||||
}
|
||||
},
|
||||
"software_profile_id": "sw-<sha256-prefix>",
|
||||
"environment_metadata": {
|
||||
"env": {
|
||||
"IMAGE_VERSION": "py3.12-cuda13.0.0",
|
||||
"FASTVIDEO_CONTAINER_IMAGE_REF": "ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-cuda13.0.0@sha256:<digest>"
|
||||
}
|
||||
},
|
||||
"environment_fingerprint": "env-<sha256-prefix>"
|
||||
}
|
||||
```
|
||||
|
||||
@@ -213,6 +372,10 @@ result, used as the rolling-baseline source of truth.
|
||||
```jsonc
|
||||
{
|
||||
"model_id": "wan-t2v-1.3b-2gpu",
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 3,
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"commit_sha": "<full sha>",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
@@ -222,34 +385,81 @@ result, used as the rolling-baseline source of truth.
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.08,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"recipe_fingerprint": "<sha256>",
|
||||
"hardware_profile_id": "hw-<sha256-prefix>",
|
||||
"software_profile_id": "sw-<sha256-prefix>",
|
||||
"environment_fingerprint": "env-<sha256-prefix>",
|
||||
"run_source": "pr",
|
||||
"branch": "feature/perf-change",
|
||||
"pr_number": "1234",
|
||||
"test_scope": "direct",
|
||||
"build_url": "https://buildkite.example/build",
|
||||
"build_id": "<buildkite-build-id>",
|
||||
"job_id": "<buildkite-job-id>",
|
||||
"quality_metadata": { "quality_status": "canonical" },
|
||||
"baseline_status": "compared",
|
||||
"comparison_status": "PASS",
|
||||
"comparison_status_reason": "Comparable baseline found with no gated regressions",
|
||||
"baseline_eligible": false,
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
### Compatibility with legacy records
|
||||
|
||||
Older records in the HF dataset may not have component timing fields. The
|
||||
comparator ignores missing or `null` metrics when computing a median, and the
|
||||
dashboard lists skipped plots for metric series that have no non-null values.
|
||||
Records missing both `run_source` and `baseline_eligible` are treated as legacy
|
||||
successful main/full-suite uploads and remain eligible for rolling baselines.
|
||||
Older records in the HF dataset may not have `result_schema_version`,
|
||||
component timing fields, or v2 identity/profile fields. Records without
|
||||
`result_schema_version` are treated as v1. The comparator ignores missing or
|
||||
`null` metrics when computing a median, and the dashboard lists skipped plots
|
||||
for metric series that have no non-null values. Records missing both
|
||||
`run_source` and `baseline_eligible` are treated as legacy successful
|
||||
main/full-suite uploads and remain eligible for rolling baselines.
|
||||
Current `perf_*.json` artifacts that lack the v2 comparison identity are
|
||||
normalized for reporting but skip rolling-baseline comparison and are not marked
|
||||
baseline eligible.
|
||||
|
||||
New v2 records compare only against the same `workload_id`, `variant_id`,
|
||||
`benchmark_version`, `recipe_fingerprint`, `hardware_profile_id`, and
|
||||
`software_profile_id` cohort, independent of the legacy `model_id` directory
|
||||
and `gpu_type` display string. Historical v1 records remain readable for
|
||||
reporting, but current legacy artifacts do not perform a `(model_id, gpu_type)`
|
||||
rolling comparison or seed new rolling baselines.
|
||||
`environment_metadata` and `environment_fingerprint` are audit data and are not
|
||||
part of the comparison key.
|
||||
The recipe prompt digests describe the prompts actually measured by the
|
||||
benchmark run; extra configured prompts are ignored unless the benchmark runner
|
||||
executes them.
|
||||
Software profile package cohorts keep exact versions for relevant
|
||||
attention/kernel packages, including FastVideo kernels, FlashAttention,
|
||||
FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
|
||||
|
||||
## Environment variable reference
|
||||
|
||||
| Variable | Default | Used by | Purpose |
|
||||
|---|---|---|---|
|
||||
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
|
||||
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
|
||||
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py`, `test_inference_performance.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Performance pytest exit code. Measured static-threshold failures are attributed per record; otherwise a nonzero code reports an infrastructure error. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
|
||||
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
|
||||
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `FASTVIDEO_ATTENTION_BACKEND` | `auto` | `test_inference_performance.py` | Requested attention backend included in `software_profile_id`. |
|
||||
| `FASTVIDEO_FA4` | `0` | `test_inference_performance.py` | FlashAttention-4 toggle included in `software_profile_id`. |
|
||||
| `FASTVIDEO_PERFORMANCE_PROFILE_VERSION` | unset | `test_inference_performance.py` | Optional explicit software cohort/profile version included in `software_profile_id`. |
|
||||
| `IMAGE_VERSION` | unset | `test_inference_performance.py` | CI container image/profile version included in `software_profile_id` when available. |
|
||||
| `FASTVIDEO_CONTAINER_IMAGE_REF` | unset | `pr_test.py`, `launch_l40s_job.py`, `test_inference_performance.py` | Resolved CI container image ref or digest recorded in `environment_metadata` for audit without changing `software_profile_id`. |
|
||||
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
|
||||
|
||||
## CI integration
|
||||
@@ -260,17 +470,24 @@ point is `fastvideo/tests/modal/pr_test.py:run_performance_tests` and the
|
||||
Buildkite artifact upload is in
|
||||
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
|
||||
|
||||
Each performance build runs pytest first. If that fixed-threshold phase fails,
|
||||
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
|
||||
artifacts are not emitted. The dashboard still runs best-effort for
|
||||
observability. When pytest passes, the rolling-baseline phase emits:
|
||||
Each performance build runs pytest first. PR and direct runs only continue to
|
||||
`compare_baseline.py` when that fixed-threshold phase passes; if pytest fails,
|
||||
Markdown summaries and normalized JSON artifacts are not emitted. Scheduled
|
||||
main runs set `PERF_UPLOAD_POLICY=always`, so they still run
|
||||
`compare_baseline.py` (with `PERF_PYTEST_RC` set) after pytest fails. Each raw
|
||||
record is checked against its own static thresholds: a measured breach reports
|
||||
`REGRESSION`, while unaffected records retain their rolling-baseline status. A
|
||||
nonzero pytest exit with no attributable static-threshold breach reports
|
||||
`INFRA_ERROR`. Failed records have `success=false` and are excluded from future
|
||||
rolling baselines. The dashboard still runs best-effort for observability.
|
||||
When the rolling-baseline phase runs, it emits:
|
||||
|
||||
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
|
||||
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
|
||||
per-benchmark row with current vs. baseline values for latency, throughput,
|
||||
memory, text encoder time, DiT time, and VAE decode time.
|
||||
* **Plotly dashboard** — `dashboard_<sha>_<ts>.html` showing time-series for
|
||||
each metric grouped by `(model_id, gpu_type)`.
|
||||
each metric grouped by comparison cohort.
|
||||
* **Normalized records** — `normalized_perf_*.json`, one per benchmark.
|
||||
Useful as input to the
|
||||
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
|
||||
@@ -279,11 +496,16 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
## Adding a new benchmark
|
||||
|
||||
1. Drop a new JSON config into
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
|
||||
use v2 identity fields:
|
||||
|
||||
```json
|
||||
{
|
||||
"benchmark_id": "<unique-id>",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "<stable-workload-id>",
|
||||
"variant_id": "<variant, e.g. 1.3b-sp2>",
|
||||
"benchmark_version": 1,
|
||||
"model": { "model_path": "...", "model_short_name": "..." },
|
||||
"init_kwargs": { "num_gpus": 1, ... },
|
||||
"generation_kwargs": { "num_frames": 45, ... },
|
||||
@@ -299,17 +521,30 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
|
||||
}
|
||||
}
|
||||
```
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
Legacy v1 configs without `config_schema_version` still load, but should not
|
||||
gain v2 identity or metadata fields until they are migrated to
|
||||
`config_schema_version: 2`. For v2 configs, `workload_id`, `variant_id`,
|
||||
and `benchmark_version` are part of the comparison key; benchmark runs
|
||||
fail if any of these identity fields are missing.
|
||||
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
|
||||
3. The first persisted main-branch run with no HF history initializes the
|
||||
baseline (passes automatically). Subsequent runs compare against it. Local
|
||||
and pull-request runs with no HF history also pass, but they do not seed the
|
||||
shared baseline.
|
||||
3. Legacy benchmarks are gated only by their static thresholds. Their current
|
||||
records skip rolling-baseline comparison and are never baseline eligible.
|
||||
V2 benchmarks with no exact comparable baseline report
|
||||
`CALIBRATION_NEEDED`; the record remains visible but does not become
|
||||
baseline eligible until a comparable scheduled-main full-suite run is
|
||||
reviewed and seeded through the prepare, review, and conditional-upload
|
||||
steps in the `reseed-performance-baseline` skill.
|
||||
|
||||
4. If the benchmark targets a GPU not currently in `thresholds`, either add
|
||||
that GPU as a key or rely on the `default` block. Note that `default` is
|
||||
@@ -320,10 +555,21 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
a useful fixed gate. The rolling baseline will still track component times
|
||||
when static component thresholds are omitted.
|
||||
|
||||
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
|
||||
include only benchmark-specific deviations. Tune these independently from
|
||||
the fixed thresholds when a metric is noisy or should be informational. The
|
||||
fixed `thresholds` block is an absolute pytest ceiling. The
|
||||
`regression_thresholds` block controls rolling-baseline comparisons against
|
||||
recent scheduled-main records.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**"No baseline for ... Initializing"** — first run for this `(model_id,
|
||||
gpu_type)`. Run will pass and (if persisting) seed the first record.
|
||||
**`CALIBRATION_NEEDED` / "No baseline found for exact comparable identity"** —
|
||||
the v2 run passes, but its normalized record remains
|
||||
`baseline_eligible=false`. Review a successful scheduled-main full-suite
|
||||
normalized artifact, then follow the `reseed-performance-baseline` skill. The
|
||||
utility only prepares a digest-protected manifest; the separately confirmed
|
||||
upload rechecks remote state and commits the batch atomically.
|
||||
|
||||
**Persistent failure right after a torch / kernel / image upgrade** —
|
||||
genuine regression *or* baseline drift. Compare the failing normalized record
|
||||
@@ -336,5 +582,5 @@ pipelines that did not report a mapped component stage.
|
||||
|
||||
**Component timing is `null`** — the generated result did not include a mapped
|
||||
stage in `logging_info.stages`. Check that the pipeline emits stage logging
|
||||
and that the stage name is listed in `STAGE_METRIC_MAP` in
|
||||
`test_inference_performance.py`.
|
||||
and that the stage emits `component_metric` or is covered by the legacy
|
||||
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
|
||||
|
||||
@@ -180,6 +180,30 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--device-folder L40S_reference_videos
|
||||
```
|
||||
|
||||
### SSIM Bootstrap Mode
|
||||
|
||||
Normal SSIM runs are strict: if a reference video or latent is missing, the
|
||||
test fails. For new-model PRs, CI can run SSIM in bootstrap mode so missing
|
||||
references are uploaded as draft artifacts for review instead of immediately
|
||||
blocking on a missing canonical reference.
|
||||
|
||||
Buildkite enables SSIM bootstrap mode when either condition is true:
|
||||
|
||||
- the PR title or Buildkite message contains `[new-model]`;
|
||||
- `FASTVIDEO_SSIM_BOOTSTRAP_MODE=1` is set for the Buildkite job.
|
||||
|
||||
Bootstrap mode passes `--ssim-bootstrap-mode` to pytest. When a generated
|
||||
artifact is available, the test uploads it under the `drafts/...` namespace in
|
||||
the SSIM reference repo and marks that case as expected-failed. After reviewing
|
||||
the draft, promote it into the canonical reference layout:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py promote-draft \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id <model_id>
|
||||
```
|
||||
|
||||
## CI Integration
|
||||
|
||||
FastVideo CI tests are orchestrated by Buildkite and run on Modal GPU
|
||||
|
||||
@@ -191,6 +191,9 @@ surfaces:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
color_correction_strength:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.dreamx_world.DreamXWorld5BARPipelineConfig
|
||||
default_camera_rotation:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
@@ -455,6 +458,8 @@ surfaces:
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
use_embedded_guidance: request.sampling.use_embedded_guidance
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# 🌊 AnyFlow Any-Step Video Distillation
|
||||
|
||||
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
|
||||
|
||||
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
|
||||
|
||||
```
|
||||
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
|
||||
```
|
||||
|
||||
for any ``t > r``.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
|
||||
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
|
||||
|
||||
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
|
||||
|
||||
## ⚙️ Inference
|
||||
|
||||
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
|
||||
|
||||
## 🧠 Algorithm
|
||||
|
||||
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
|
||||
|
||||
### Stage 1 — Pretrain (flow-map central-difference)
|
||||
|
||||
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
|
||||
|
||||
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
|
||||
|
||||
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
|
||||
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
|
||||
- the remainder is free.
|
||||
|
||||
The student forward at ``(t, r)`` is trained against the central-difference target
|
||||
|
||||
```
|
||||
target = (eps - x_0) - (t - r) · dF/dt
|
||||
```
|
||||
|
||||
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
|
||||
|
||||
### Stage 2 — On-policy DMD
|
||||
|
||||
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
|
||||
|
||||
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
|
||||
|
||||
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
### Stage 1 — pretrain
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
|
||||
```
|
||||
|
||||
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
|
||||
|
||||
- Global batch size: 32 (8 GPUs × 4 per-GPU)
|
||||
- Learning rate: 5e-5
|
||||
- Flow shift: 5.0
|
||||
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
|
||||
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
|
||||
- ``weight_type``: ``beta08``
|
||||
- ``fuse_guidance_scale``: 3.0
|
||||
- Training steps: 6000
|
||||
|
||||
### Stage 2 — on-policy
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
|
||||
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
|
||||
```
|
||||
|
||||
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
|
||||
|
||||
**Key configuration**:
|
||||
|
||||
- Global batch size: 8 (8 GPUs × 1 per-GPU)
|
||||
- Learning rate: 2e-6
|
||||
- Flow shift: 5.0
|
||||
- ``student_sample_steps``: 4
|
||||
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
|
||||
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
|
||||
- ``real_score_guidance_scale``: 3.0
|
||||
- ``generator_update_interval``: 5 (DMD2 alternation)
|
||||
- Training steps: 4000
|
||||
|
||||
## 🔌 Loading published AnyFlow checkpoints
|
||||
|
||||
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
|
||||
|
||||
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
|
||||
|
||||
## 🧭 Note on ``fuse_guidance_scale``
|
||||
|
||||
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
|
||||
|
||||
```
|
||||
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
|
||||
```
|
||||
|
||||
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
|
||||
|
||||
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
|
||||
@@ -74,6 +74,23 @@ uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
### Flash Attention 4 (opt-in)
|
||||
|
||||
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
|
||||
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
|
||||
runtime on some GPU/shape combinations. To use FA4, install the pinned
|
||||
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
|
||||
|
||||
```bash
|
||||
export FASTVIDEO_FA4=1
|
||||
```
|
||||
|
||||
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
|
||||
cannot serve there: grad-enabled (training) attention (FA4's backward requires
|
||||
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
|
||||
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
|
||||
FastVideo fails loudly instead of silently falling back.
|
||||
|
||||
### FP4 Flash Attention 4 (Blackwell only)
|
||||
|
||||
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
|
||||
|
||||
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
# Training Trackers
|
||||
|
||||
FastVideo can send training metrics and validation media to Weights & Biases
|
||||
or SwanLab. Tracking runs only on global rank 0, and local tracker files are
|
||||
stored under `<output_dir>/tracker`.
|
||||
|
||||
## Supported Trackers
|
||||
|
||||
| Value | Backend | Installation |
|
||||
|-------|---------|--------------|
|
||||
| `wandb` | Weights & Biases | Included with FastVideo |
|
||||
| `swanlab` | SwanLab | Install the optional `swanlab` dependency |
|
||||
| `none` | Disable external tracking | No additional package |
|
||||
|
||||
You can enable more than one backend, for example `trackers: [wandb, swanlab]`.
|
||||
Metrics and validation media are converted to the artifact type required by
|
||||
each backend.
|
||||
|
||||
## Install SwanLab
|
||||
|
||||
For a published FastVideo installation, install the SwanLab extra:
|
||||
|
||||
```bash
|
||||
uv pip install "fastvideo[swanlab]"
|
||||
```
|
||||
|
||||
For an editable source checkout, include the same extra during installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[swanlab]"
|
||||
```
|
||||
|
||||
If FastVideo is already installed, you can install the compatible SDK directly:
|
||||
|
||||
```bash
|
||||
uv pip install "swanlab>=0.6.7"
|
||||
```
|
||||
|
||||
Authenticate once before starting a training run:
|
||||
|
||||
```bash
|
||||
swanlab login
|
||||
```
|
||||
|
||||
See the [SwanLab login documentation](https://docs.swanlab.cn/en/api/cli-swanlab-login.html)
|
||||
for non-interactive and self-hosted setups.
|
||||
|
||||
## Configure Tracking
|
||||
|
||||
Select SwanLab in the YAML config used by the modular training framework:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
checkpoint:
|
||||
output_dir: outputs/my_run
|
||||
tracker:
|
||||
trackers: [swanlab]
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
To log to both supported services:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
trackers: [wandb, swanlab]
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
An empty or omitted `trackers` list selects W&B when `project_name` is set.
|
||||
Use an explicit `none` entry to disable external tracking:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
trackers: [none]
|
||||
```
|
||||
|
||||
## Validation Videos
|
||||
|
||||
SwanLab currently accepts GIF video artifacts. FastVideo converts validation
|
||||
MP4 files and in-memory video arrays to GIF automatically before logging them.
|
||||
For video files, FastVideo uses the sampling frame rate supplied by the caller,
|
||||
or the source file's frame rate when no value is supplied. In-memory arrays use
|
||||
the frame rate supplied by the caller. Both forms fall back to 16 FPS when no
|
||||
frame rate is available.
|
||||
|
||||
For details about configuring validation callbacks, see
|
||||
[Training Infrastructure](train_infra.md#callbacks-pluggable-hooks).
|
||||
@@ -161,6 +161,21 @@ training:
|
||||
decay_interval_steps: 0
|
||||
```
|
||||
|
||||
`training.data.data_path` can also mix multiple preprocessed datasets by using a mapping from dataset path to repeat count:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
data/multi3d_games: 2
|
||||
```
|
||||
|
||||
The repeat count duplicates that dataset's parquet file list before shuffling/sampling, so the example above trains with roughly twice as much `multi3d_games` exposure as `zeldam2-clean`. Paths are just suggested locations; use any local path that contains a FastVideo preprocessed parquet dataset.
|
||||
|
||||
See [Training Trackers](trackers.md) to configure Weights & Biases or SwanLab,
|
||||
including SwanLab installation and authentication.
|
||||
|
||||
### `callbacks` — Pluggable hooks
|
||||
|
||||
Callbacks run at specific points in the training loop (before/after optimizer
|
||||
@@ -323,6 +338,40 @@ Self-Forcing inherits all DMD2 parameters, plus:
|
||||
| `enable_gradient_in_rollout` | `true` | Enable backprop through rollout |
|
||||
| `start_gradient_frame` | `0` | Frame index where gradients begin |
|
||||
|
||||
### Streaming Long Tuning
|
||||
|
||||
`StreamingLongTuningMethod` extends Self-Forcing for LongLive-style rollouts. It
|
||||
keeps a streaming state, generates overlapping chunks, and trains only the new
|
||||
frames while preserving context from earlier chunks.
|
||||
|
||||
For the MatrixGame2/Zelda world-model example, self-forcing and long tuning are
|
||||
separate runs: first train or load the 1k-step self-forcing checkpoint using
|
||||
`examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml`,
|
||||
then run
|
||||
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
|
||||
from that checkpoint for the 3k-step streaming long-tuning stage.
|
||||
|
||||
```yaml
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
streaming_reencode_overlap_anchor: true
|
||||
streaming_anchor_inject_k: 1
|
||||
streaming_require_full_blocks: true
|
||||
multi_phased_distill_schedule:
|
||||
- stage: streaming_long
|
||||
start_step: 0
|
||||
end_step: 3000
|
||||
num_latent_t: 39
|
||||
streaming_training: true
|
||||
```
|
||||
|
||||
See
|
||||
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
|
||||
for a complete MatrixGame2/Zelda configuration.
|
||||
|
||||
---
|
||||
|
||||
## Callbacks
|
||||
|
||||
@@ -0,0 +1,64 @@
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
|
||||
|
||||
def _env_int(name: str, default: int) -> int:
|
||||
return int(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def _env_float(name: str, default: float) -> float:
|
||||
return float(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
|
||||
prompt = os.getenv(
|
||||
"DREAMX_WORLD_PROMPT",
|
||||
"A cinematic first-person drive through a futuristic coastal city at "
|
||||
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
|
||||
)
|
||||
image_path = os.getenv(
|
||||
"DREAMX_WORLD_IMAGE_PATH",
|
||||
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
|
||||
)
|
||||
|
||||
kwargs = {
|
||||
"output_path": OUTPUT_PATH,
|
||||
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
|
||||
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
|
||||
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
|
||||
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
|
||||
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
|
||||
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
|
||||
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
|
||||
"action_speed_list": [
|
||||
float(value)
|
||||
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
|
||||
],
|
||||
}
|
||||
if image_path:
|
||||
kwargs["image_path"] = image_path
|
||||
|
||||
try:
|
||||
generator.generate_video(prompt, **kwargs)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,140 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import os
|
||||
import re
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
(
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _safe_filename(text: str, max_len: int = 100) -> str:
|
||||
"""Make a stable, filesystem-friendly filename base."""
|
||||
s = text[:max_len].strip()
|
||||
s = s.replace(os.sep, "_")
|
||||
if os.altsep:
|
||||
s = s.replace(os.altsep, "_")
|
||||
s = re.sub(r"\s+", " ", s)
|
||||
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
|
||||
s = s.strip(" .")
|
||||
return s or "prompt"
|
||||
|
||||
|
||||
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
|
||||
if not os.path.isdir(out_dir):
|
||||
return
|
||||
|
||||
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
|
||||
for fn in os.listdir(out_dir):
|
||||
if pattern.match(fn):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.remove(os.path.join(out_dir, fn))
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--model-path",
|
||||
default="official_weights/FLUX.1-dev",
|
||||
help="Local Diffusers checkpoint dir or HF repo id.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
default="outputs/flux_dev/samples",
|
||||
help="Directory for saved PNG outputs.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--prompt",
|
||||
action="append",
|
||||
default=None,
|
||||
help="Prompt. Repeat for multiple images.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
|
||||
)
|
||||
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
|
||||
p.add_argument("--height", type=int, default=1024, help="Output height.")
|
||||
p.add_argument("--width", type=int, default=1024, help="Output width.")
|
||||
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
|
||||
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
|
||||
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
|
||||
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": args.num_gpus,
|
||||
"workload_type": "t2i",
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"image_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": False,
|
||||
"use_fsdp_inference": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=args.model_path,
|
||||
**init_kwargs,
|
||||
)
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
filename_base = (
|
||||
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
)
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
generation_kwargs = {
|
||||
"output_path": output_path,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": args.steps,
|
||||
"guidance_scale": args.guidance,
|
||||
"use_embedded_guidance": True,
|
||||
"true_cfg_scale": 1.0,
|
||||
"seed": seed,
|
||||
"save_video": True,
|
||||
}
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
print(f"[flux] done. outputs written to: {args.out_dir}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,107 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
|
||||
text-to-image generation command, saved as a PNG."
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/landscape.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default=("A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style."),
|
||||
help="Text prompt.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
|
||||
|
||||
IMAGE_PATH = "assets/girl.png"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A woman stands up and walks away"
|
||||
)
|
||||
_ = generator.generate_video(
|
||||
prompt,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=1024,
|
||||
width=1024,
|
||||
num_frames=121,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,120 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
|
||||
want a minimal edit command (text + image -> edited image), saved as a PNG."
|
||||
|
||||
GLM-Image is a single unified pipeline: passing a condition image switches it
|
||||
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
|
||||
write pass), so the generator config is identical to `basic_glm_image.py` — the
|
||||
`inputs.pil_image` on the request is what selects the edit mode.
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image",
|
||||
default="assets/images/couple.jpg",
|
||||
help="Condition image to edit.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/edited.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="Change the background to a snowy mountain landscape at golden hour.",
|
||||
help="Edit instruction.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
condition = Image.open(args.image).convert("RGB")
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
# The pipeline is registered as t2i; passing inputs.pil_image below switches
|
||||
# it to the edit path.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
inputs=InputConfig(pil_image=condition),
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,69 @@
|
||||
# LTX-2.3 distilled inference configs
|
||||
|
||||
Ready-to-run `fastvideo generate` run configs for the LTX-2.3
|
||||
distilled model (`FastVideo/LTX-2.3-Distilled-Diffusers`), covering both
|
||||
workloads (t2v / i2v), both two-stage step schedules (`5+2`, `8+3` = denoise
|
||||
+ refine), and four resolutions.
|
||||
|
||||
```bash
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
|
||||
```
|
||||
|
||||
Each config is self-contained (no preset registry needed): the two-stage
|
||||
refine is wired via `generator.pipeline.preset_overrides.refine`, and the
|
||||
base sampling knobs live under `request.sampling`. The refine upsampler
|
||||
auto-resolves from the model's `spatial_upscaler`.
|
||||
|
||||
## Configs
|
||||
|
||||
| workload | schedule | resolution (HxW) | file |
|
||||
|---|---|---|---|
|
||||
| t2v | 5+2 | 1280x832 | `t2v_5s2_1280x832.yaml` |
|
||||
| t2v | 5+2 | 1024x1536 | `t2v_5s2_1024x1536.yaml` |
|
||||
| t2v | 5+2 | 768x1280 | `t2v_5s2_768x1280.yaml` |
|
||||
| t2v | 5+2 | 512x768 | `t2v_5s2_512x768.yaml` |
|
||||
| t2v | 8+3 | 1280x832 | `t2v_8s3_1280x832.yaml` |
|
||||
| t2v | 8+3 | 1024x1536 | `t2v_8s3_1024x1536.yaml` |
|
||||
| t2v | 8+3 | 768x1280 | `t2v_8s3_768x1280.yaml` |
|
||||
| t2v | 8+3 | 512x768 | `t2v_8s3_512x768.yaml` |
|
||||
| i2v | 5+2 | 1280x832 | `i2v_5s2_1280x832.yaml` |
|
||||
| i2v | 5+2 | 1024x1536 | `i2v_5s2_1024x1536.yaml` |
|
||||
| i2v | 5+2 | 768x1280 | `i2v_5s2_768x1280.yaml` |
|
||||
| i2v | 5+2 | 512x768 | `i2v_5s2_512x768.yaml` |
|
||||
| i2v | 8+3 | 1280x832 | `i2v_8s3_1280x832.yaml` |
|
||||
| i2v | 8+3 | 1024x1536 | `i2v_8s3_1024x1536.yaml` |
|
||||
| i2v | 8+3 | 768x1280 | `i2v_8s3_768x1280.yaml` |
|
||||
| i2v | 8+3 | 512x768 | `i2v_8s3_512x768.yaml` |
|
||||
|
||||
## Overriding without editing a file
|
||||
|
||||
Dotted overrides (prefixes `generator.` / `request.`) let you tweak any field:
|
||||
|
||||
```bash
|
||||
# swap prompt
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml \
|
||||
--request.prompt "a red fox running through fresh snow"
|
||||
|
||||
# change output path / gpu count
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml \
|
||||
--request.output.output_path outputs/preview.mp4 \
|
||||
--generator.engine.num_gpus 4
|
||||
```
|
||||
|
||||
## i2v
|
||||
|
||||
The `i2v_*` configs take a first-frame image via
|
||||
`request.extensions.ltx2_images` (`[[path, frame_offset, weight]]`). Edit the
|
||||
path in the file, or override it:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml \
|
||||
--request.extensions.ltx2_images '[["/data/portrait.jpg", 0, 1.0]]'
|
||||
```
|
||||
|
||||
## Schedules
|
||||
|
||||
`5+2` is the fast preview schedule; `8+3` is the higher-quality distilled
|
||||
recipe. Refine (`preset_overrides.refine.num_inference_steps`) only accepts 2
|
||||
or 3 steps.
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,38 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_512x768.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,33 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -0,0 +1,78 @@
|
||||
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
|
||||
# the student matches an EMA copy of itself at the next timestep, all under
|
||||
# clean-history teacher forcing.
|
||||
#
|
||||
# All three roles initialize from the SAME checkpoint (the teacher-forcing
|
||||
# AR-diffusion model). Point init_from at that checkpoint for a real run.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 48
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.99
|
||||
ema_start_step: 200
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_cd
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_cd_shift5
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,105 @@
|
||||
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
|
||||
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
|
||||
# rollout from pure noise. Teacher provides the real score, critic
|
||||
# learns the fake score; both inherited from DMD2Method.
|
||||
#
|
||||
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
|
||||
# or with the NVIDIA-released checkpoint
|
||||
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
|
||||
# the paper weights (the delta_embedder rename is handled by the
|
||||
# param_names_mapping in WanVideoArchConfig).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: <PATH_TO_PRETRAIN_CKPT>
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 3.0
|
||||
dmd_denoising_steps: [999, 937, 833, 624]
|
||||
warp_denoising_step: false
|
||||
|
||||
# AnyFlow rollout knobs.
|
||||
student_sample_steps: 4
|
||||
use_mean_velocity: true
|
||||
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
|
||||
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
|
||||
|
||||
# Critic optimizer (DMD2 inherited).
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
attn_kind: vsa
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_onpolicy
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_onpolicy
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -0,0 +1,83 @@
|
||||
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
|
||||
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
|
||||
# can be sampled at arbitrary NFE in the on-policy stage.
|
||||
#
|
||||
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
|
||||
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
|
||||
diffusion_ratio: 0.5
|
||||
consistency_ratio: 0.25
|
||||
epsilon: 5 # finite-difference step in absolute train-timestep units
|
||||
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
|
||||
fuse_guidance_scale: 3.0
|
||||
# shift is taken from pipeline.flow_shift below.
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 4
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
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: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_pretrain
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_pretrain
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
# Enable AnyFlow dual-timestep conditioning. The student loads from
|
||||
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
|
||||
# get initialized identically to time_embedder via deep-copy in
|
||||
# WanTimeTextImageEmbedding.__init__.
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -100,9 +100,6 @@ callbacks:
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 81
|
||||
# Validation/inference uses standard CFG in both clean and Self-Forcing,
|
||||
# so this directly matches Self-Forcing guidance_scale=3.0.
|
||||
guidance_scale: 3.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model with a block size of 1 frame
|
||||
# - Training: each frame gets its own independent noise level (frame-wise
|
||||
# diffusion forcing), versus the chunk-wise variant that shares one noise
|
||||
# level across num_frames_per_block frames.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_dfsft_framewise
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,72 @@
|
||||
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model
|
||||
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
|
||||
# denoises the current block while attending to *clean* history (clean_x),
|
||||
# not its own noisy rollout.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_tfsft
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,32 +1,119 @@
|
||||
# World-Model: Matrix-Game 2.0 I2V
|
||||
|
||||
Three training scenarios for the Matrix-Game 2.0 I2V world model on the
|
||||
new YAML-driven trainer (`fastvideo/train/entrypoint/train.py`).
|
||||
Training scenarios for the Matrix-Game 2.0 I2V world model on Solaris (Minecraft)
|
||||
data and Zelda data, using the new YAML-driven trainer
|
||||
(`fastvideo/train/entrypoint/train.py`).
|
||||
|
||||
## Solaris Configs
|
||||
|
||||
| Config | Method | Student | Notes |
|
||||
|---|---|---|---|
|
||||
| `finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
|
||||
| `dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
|
||||
| `self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
|
||||
| `solaris/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
|
||||
| `solaris/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
|
||||
| `solaris/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Matrix-Game 2.0 DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
|
||||
|
||||
## Zelda Configs
|
||||
|
||||
| Config | Method | Student | Notes |
|
||||
|---|---|---|---|
|
||||
| `zelda/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Zelda bidirectional I2V finetuning from `FastVideo/Matrix-Game-2.0-Base-Diffusers`. Uses 33-frame clips and Zelda validation with action overlays. |
|
||||
| `zelda/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Zelda causal Diffusion-Forcing SFT from `mignonjia/mg_bidirectional_zelda`. Uses the same Zelda data, resolution, optimizer, and validation defaults as the Zelda finetune config. |
|
||||
| `zelda/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Zelda DMD/Self-Forcing distillation; student init = `mignonjia/mg_causal_zelda`, teacher = bidirectional (`mignonjia/mg_bidirectional_zelda`), critic = bidirectional. |
|
||||
| `zelda/streaming_long_tuning_causal_i2v.yaml` | `StreamingLongTuningMethod` | `MatrixGame2CausalModel` | LongLive-style streaming long tuning from the 1k-step Zelda self-forcing checkpoint. |
|
||||
|
||||
Zelda world-model distillation is a two-run workflow: first run
|
||||
`zelda/self_forcing_causal_i2v.yaml` to train or load the 1k-step
|
||||
self-forcing checkpoint (`mignonjia/mg_sf_distilled_zelda_1k_steps`), then run
|
||||
`zelda/streaming_long_tuning_causal_i2v.yaml` for the 3k-step streaming
|
||||
long-tuning stage. The long-tuning YAML starts from that 1k-step checkpoint; it
|
||||
does not run the short self-forcing stage inside the same config.
|
||||
|
||||
## Zelda Training Data
|
||||
|
||||
The Zelda training configs use `data/zeldam2-clean` as a suggested local path.
|
||||
Download the dataset from Hugging Face before running those configs:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id mignonjia/zeldam2-clean \
|
||||
--local_dir data/zeldam2-clean \
|
||||
--repo_type dataset
|
||||
```
|
||||
|
||||
You can store the dataset elsewhere; update `training.data.data_path` in the
|
||||
YAML to point at that location.
|
||||
|
||||
## Multi3D Training Data
|
||||
|
||||
`zelda/finetune_i2v.yaml` includes an optional, commented-out Multi3D entry.
|
||||
Enable it only when you want to mix Zelda with multi-game data from
|
||||
`data/multi3d_games`. You can store this dataset anywhere; before enabling it,
|
||||
update the matching commented `training.data.data_path` key in the YAML to the
|
||||
correct location.
|
||||
|
||||
To mix datasets in a training YAML, set `training.data.data_path` to a
|
||||
path-to-repeat-count mapping. For example, `zelda/finetune_i2v.yaml` can use
|
||||
`data/zeldam2-clean: 1` and `# data/multi3d_games: 10`; uncommenting the
|
||||
Multi3D entry repeats the multi-game parquet list ten times before training
|
||||
samples are shuffled.
|
||||
|
||||
## World Model Validation Data
|
||||
|
||||
The Zelda validation configs expect a small public validation bundle under
|
||||
`data/zelda_validation_data`.
|
||||
|
||||
Download it from Hugging Face before running the Zelda scenarios:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id mignonjia/zelda_validation_data \
|
||||
--local_dir data/zelda_validation_data \
|
||||
--repo_type dataset
|
||||
```
|
||||
|
||||
The bundle contains `validation_zelda.json`, `images/`, and `actions/`.
|
||||
The Zelda configs point
|
||||
`callbacks.validation.dataset_file` at
|
||||
`data/zelda_validation_data/validation_zelda.json`.
|
||||
|
||||
## Usage
|
||||
|
||||
### Solaris
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/finetune_i2v.yaml
|
||||
examples/train/scenario/worldmodel/solaris/finetune_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml
|
||||
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/self_forcing_causal_i2v.yaml
|
||||
examples/train/scenario/worldmodel/solaris/self_forcing_causal_i2v.yaml
|
||||
```
|
||||
|
||||
### Zelda
|
||||
|
||||
```bash
|
||||
# Finetuning / DFSFT
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/finetune_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/dfsft_causal_i2v.yaml
|
||||
|
||||
# Distillation / long tuning
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml
|
||||
```
|
||||
|
||||
Override any field on the command line:
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml \
|
||||
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml \
|
||||
--training.distributed.num_gpus 8 \
|
||||
--training.optimizer.learning_rate 1e-5
|
||||
```
|
||||
|
||||
+1
-1
@@ -97,4 +97,4 @@ callbacks:
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,94 @@
|
||||
# Diffusion-Forcing SFT: Zelda world model I2V Causal
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-5
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_causal_dfsft
|
||||
training_state_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
entity: hapo-exp
|
||||
project_name: mg_1.3b_zelda
|
||||
run_name: zelda_causal_dfsft
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 200
|
||||
sampling_steps: [40]
|
||||
sampling_timesteps: [1000, 975, 950, 925, 900, 875, 850, 825, 800, 775,
|
||||
750, 725, 700, 675, 650, 625, 600, 575, 550, 525,
|
||||
500, 475, 450, 425, 400, 375, 350, 325, 300, 275,
|
||||
250, 225, 200, 175, 150, 125, 100, 75, 50, 25]
|
||||
num_frames: 33
|
||||
overlay_actions: true
|
||||
guidance_scale: 6.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -0,0 +1,88 @@
|
||||
# Matrix-Game 2.0 Zelda + multi-game I2V finetune.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: FastVideo/Matrix-Game-2.0-Base-Diffusers
|
||||
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: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
# data/multi3d_games: 10
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0 # unused for MatrixGame2 I2V; no text_embedding CFG dropout
|
||||
seed: 42
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-5
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_with_mg_init
|
||||
training_state_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: mg_1.3b_zelda
|
||||
run_name: zelda_with_mg_init
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
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.matrixgame2.matrixgame2_i2v_pipeline.MatrixGame2I2VPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 200
|
||||
sampling_steps: [40]
|
||||
num_frames: 33
|
||||
overlay_actions: true
|
||||
guidance_scale: 6.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,120 @@
|
||||
# Self-forcing distillation: Zelda world model I2V Causal
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_causal_zelda
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
same_step_across_blocks: true
|
||||
last_step_only: false
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 3.0e-7
|
||||
fake_score_betas: [0.9, 0.95]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/zeldam2-clean
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1001
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 3.0e-6
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/zelda_causal_self_forcing
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: mg2_self_forcing_9_latents
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
ema:
|
||||
decay: 0.99
|
||||
start_iter: 200
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 100
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 153
|
||||
overlay_actions: true
|
||||
keyboard_value_scale: 1.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -0,0 +1,140 @@
|
||||
# MatrixGame2 I2V LongLive-style streaming distillation: Zelda world model I2V Causal
|
||||
# Student init from self forcing after 1k steps
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_sf_distilled_zelda_1k_steps
|
||||
trainable: true
|
||||
# transformer_override_safetensor: outputs/matrixgame_dmd/checkpoints/mg_zelda_sf_m2/checkpoint-1000_weight_only/ema/generator_ema.safetensors
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
same_step_across_blocks: true
|
||||
last_step_only: false
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
streaming_training: true
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
streaming_reencode_overlap_anchor: true
|
||||
streaming_anchor_inject_k: 1
|
||||
streaming_require_full_blocks: true
|
||||
|
||||
multi_phased_distill_schedule:
|
||||
- stage: streaming_long
|
||||
start_step: 0
|
||||
end_step: 3000
|
||||
num_latent_t: 39
|
||||
streaming_training: true
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 3.0e-7
|
||||
fake_score_betas: [0.9, 0.95]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/zeldam2-clean
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1001
|
||||
num_latent_t: 39
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 153
|
||||
|
||||
optimizer:
|
||||
learning_rate: 3.0e-6
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/zelda_causal_long_tuning
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: mg2_39only_streaming_long
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
ema:
|
||||
decay: 0.99
|
||||
start_iter: 200
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 100
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 153
|
||||
overlay_actions: true
|
||||
keyboard_value_scale: 1.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -173,6 +185,14 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,6 +2,42 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,6 +46,23 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxSamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 1024
|
||||
width: int = 1024
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 3.5
|
||||
use_embedded_guidance: bool = True
|
||||
true_cfg_scale: float = 1.0
|
||||
@@ -90,6 +90,10 @@ class SamplingParam:
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
|
||||
use_embedded_guidance: bool = False
|
||||
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
|
||||
true_cfg_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
@@ -325,6 +329,18 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-embedded-guidance",
|
||||
action="store_true",
|
||||
default=SamplingParam.use_embedded_guidance,
|
||||
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--true-cfg-scale",
|
||||
type=float,
|
||||
default=SamplingParam.true_cfg_scale,
|
||||
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
|
||||
@@ -150,6 +150,7 @@ class SamplingConfig:
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
use_embedded_guidance: bool | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@@ -5,99 +5,10 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
|
||||
fa_version = "4"
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer have a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
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"
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
|
||||
# breaks the graph at the call site (observed: wanvideo.py self-attn,
|
||||
# once per layer every step), which fragments the compiled region and
|
||||
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
|
||||
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
# opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
# (correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
|
||||
if fa_version in ("2", "3"):
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
|
||||
# call shape used by FlashAttentionImpl.forward's default branch
|
||||
# (see `flash_attn_func_compilable(...)` call site below). The
|
||||
# masked/no-pad and varlen / cross-attn paths use different
|
||||
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
|
||||
# are intentionally out of scope for this PR — wrapping them is a
|
||||
# natural follow-up. The wrapper's signature is the contract: any
|
||||
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
|
||||
# return_attn_probs, ...) raises TypeError at the call site, so
|
||||
# silent loss of kwargs is not a failure mode.
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
|
||||
# same dtype/device as q (head dim taken from v).
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
from fastvideo.attention.utils.flash_attn_default import (
|
||||
fa_version,
|
||||
flash_attn_func_compilable,
|
||||
)
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
@@ -108,8 +19,6 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_WARNED_NON_FA_DTYPE = False
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
@@ -271,12 +180,8 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
|
||||
orig_dtype = query.dtype
|
||||
if orig_dtype not in (torch.float16, torch.bfloat16):
|
||||
global _WARNED_NON_FA_DTYPE
|
||||
if not _WARNED_NON_FA_DTYPE:
|
||||
_WARNED_NON_FA_DTYPE = True
|
||||
logger.warning(
|
||||
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
|
||||
"kernel and restoring on output.", orig_dtype)
|
||||
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
|
||||
f"bfloat16 for the kernel and restoring on output.")
|
||||
query = query.to(torch.bfloat16)
|
||||
key = key.to(torch.bfloat16)
|
||||
value = value.to(torch.bfloat16)
|
||||
@@ -293,9 +198,17 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
|
||||
# Route through the *_compilable wrappers so dynamo sees one
|
||||
# traceable node for each masked entry point (the unpad/pad
|
||||
# bookkeeping runs eager inside the custom op). On FA2 these
|
||||
# wrappers go through ops with full register_autograd, so
|
||||
# training also backprops through the op (no graph break on
|
||||
# the training path); on FA3/FA4 they carve out to the
|
||||
# autograd.Function for grad-enabled calls — see
|
||||
# fastvideo/attention/utils/flash_attn_no_pad.py.
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad,
|
||||
flash_attn_no_pad_compilable as flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad_compilable as flash_attn_varlen_qk_no_pad,
|
||||
)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
@@ -322,7 +235,11 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
)
|
||||
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
attn_mask_padded = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
|
||||
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, attn_mask.shape[-1]).to(device=query.device)
|
||||
if key_padding_mask.shape[-1] > qkv.shape[1]:
|
||||
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
|
||||
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
|
||||
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
|
||||
elif self.nvfp4_fa4:
|
||||
output = self._forward_nvfp4(query, key, value)
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NABLA block-sparse flex-attention backend (Kandinsky5 "nabla" checkpoints).
|
||||
|
||||
The block mask is data-dependent: nablaT_v2 mean-pools 64-token blocks of the
|
||||
fractal-ordered sequence, thresholds the softmaxed block map, and ORs it with a
|
||||
precomputed spatio-temporal-window (STA) mask carried on the attention
|
||||
metadata. The mask spans the full sequence, so this backend does not support
|
||||
sequence parallelism — use it via LocalAttention only.
|
||||
"""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
from torch.nn.attention.flex_attention import BlockMask, flex_attention
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
|
||||
CAN_USE_FLEX_ATTN = True
|
||||
except ImportError:
|
||||
CAN_USE_FLEX_ATTN = False
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
|
||||
|
||||
def nablaT_v2(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
sta: torch.Tensor,
|
||||
thr: float = 0.9,
|
||||
) -> "BlockMask":
|
||||
q = q.transpose(1, 2).contiguous()
|
||||
k = k.transpose(1, 2).contiguous()
|
||||
|
||||
# Map estimation
|
||||
B, h, S, D = q.shape
|
||||
s1 = S // 64
|
||||
qa = q.reshape(B, h, s1, 64, D).mean(-2)
|
||||
ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1)
|
||||
map = qa @ ka
|
||||
|
||||
map = torch.softmax(map / math.sqrt(D), dim=-1)
|
||||
# Map binarization
|
||||
vals, inds = map.sort(-1)
|
||||
cvals = vals.cumsum_(-1)
|
||||
mask = (cvals >= 1 - thr).int()
|
||||
mask = mask.gather(-1, inds.argsort(-1))
|
||||
|
||||
mask = torch.logical_or(mask, sta)
|
||||
|
||||
# BlockMask creation
|
||||
kv_nb = mask.sum(-1).to(torch.int32)
|
||||
kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32)
|
||||
return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None)
|
||||
|
||||
|
||||
class NablaAttentionBackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "NABLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["NablaAttentionImpl"]:
|
||||
return NablaAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["NablaAttentionMetadata"]:
|
||||
return NablaAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["NablaAttentionMetadataBuilder"]:
|
||||
return NablaAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class NablaAttentionMetadata(AttentionMetadata):
|
||||
# Block-level STA window mask [1, 1, S/64, S/64], precomputed once per run.
|
||||
sta_mask: torch.Tensor = None # type: ignore[assignment]
|
||||
# Cumulative-probability threshold for block-map binarization.
|
||||
P: float = 0.9
|
||||
visual_shape: tuple[int, int, int] = (0, 0, 0)
|
||||
|
||||
|
||||
class NablaAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
sta_mask: torch.Tensor,
|
||||
P: float,
|
||||
visual_shape: tuple[int, int, int],
|
||||
**kwargs: Any,
|
||||
) -> NablaAttentionMetadata:
|
||||
return NablaAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
sta_mask=sta_mask,
|
||||
P=P,
|
||||
visual_shape=visual_shape,
|
||||
)
|
||||
|
||||
|
||||
class NablaAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
causal: bool = False,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
if not CAN_USE_FLEX_ATTN:
|
||||
raise RuntimeError("NABLA attention requires torch.nn.attention.flex_attention, "
|
||||
"which is unavailable in this PyTorch build.")
|
||||
if causal:
|
||||
raise ValueError("NABLA attention does not support causal masking.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: NablaAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
# q/k/v: [B, S, heads, head_dim], fractal-ordered by the model; S % 64 == 0.
|
||||
block_mask = nablaT_v2(query, key, attn_metadata.sta_mask, thr=attn_metadata.P)
|
||||
return flex_attention(
|
||||
query=query.transpose(1, 2),
|
||||
key=key.transpose(1, 2),
|
||||
value=value.transpose(1, 2),
|
||||
block_mask=block_mask,
|
||||
).transpose(1, 2)
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
from torch.nn import functional as F
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -14,12 +15,18 @@ class SDPABackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
def get_supported_head_sizes() -> list[int] | None:
|
||||
# torch.nn.functional.scaled_dot_product_attention is head-size
|
||||
# agnostic: the math backend handles any size and the fused kernels
|
||||
# fall back internally when they cannot. The previous list
|
||||
# ([32, 64, ..., 256]) was copied from FlashAttentionBackend in the
|
||||
# v1 refactor (#270) and wrongly excluded sizes such as 80 (e.g.
|
||||
# CLIP vision encoders). None means "no restriction".
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SDPA"
|
||||
return "TORCH_SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SDPAImpl"]:
|
||||
@@ -49,9 +56,51 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
# Store the mask exactly as passed. The metadata is cross-backend:
|
||||
# call sites (HYWorld, HunyuanVideo15) build SDPAMetadata while the
|
||||
# layer's selector may pick FLASH_ATTN, and the shared convention for
|
||||
# padding masks is the tokenizer-style 2D [batch, key_len]. Any
|
||||
# reshaping for torch.sdpa happens inside the SDPA impl
|
||||
# (_normalize_attn_mask_for_sdpa).
|
||||
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
|
||||
|
||||
def _normalize_attn_mask_for_sdpa(
|
||||
attn_mask: torch.Tensor | None,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
) -> torch.Tensor | None:
|
||||
if attn_mask is None:
|
||||
return None
|
||||
|
||||
attn_mask = attn_mask.to(device=query.device)
|
||||
# F.scaled_dot_product_attention only accepts bool or float masks;
|
||||
# tokenizers commonly produce int64 0/1 padding masks.
|
||||
if attn_mask.dtype != torch.bool and not attn_mask.dtype.is_floating_point:
|
||||
attn_mask = attn_mask != 0
|
||||
|
||||
key_len = key.shape[-2]
|
||||
if attn_mask.shape[-1] > key_len:
|
||||
raise ValueError("Invalid attention mask length for SDPA: "
|
||||
f"expected at most {key_len}, got {attn_mask.shape[-1]}")
|
||||
if attn_mask.shape[-1] < key_len:
|
||||
# Front-pad as "attend": double-stream layouts (HYWorld) prepend
|
||||
# non-text tokens the tokenizer mask does not cover.
|
||||
valid_value = True if attn_mask.dtype == torch.bool else 0.0
|
||||
attn_mask = F.pad(attn_mask, (key_len - attn_mask.shape[-1], 0), value=valid_value)
|
||||
|
||||
if attn_mask.dim() == 2:
|
||||
# In-tree producers pass 2D [batch, key_len] padding masks; lift to a
|
||||
# broadcastable [batch, 1, 1, key_len] here so torch.sdpa does not
|
||||
# reinterpret 2D as its documented [query_len, key_len] broadcast.
|
||||
return attn_mask[:, None, None, :]
|
||||
if attn_mask.dim() == 3:
|
||||
return attn_mask[:, None, :, :]
|
||||
if attn_mask.dim() == 4:
|
||||
return attn_mask
|
||||
raise ValueError(f"Unsupported attention mask shape for SDPA: {attn_mask.shape}")
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -82,6 +131,7 @@ class SDPAImpl(AttentionImpl):
|
||||
|
||||
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
|
||||
and hasattr(attn_metadata, "attn_mask")) else None
|
||||
attn_mask = _normalize_attn_mask_for_sdpa(attn_mask, query, key)
|
||||
attn_kwargs = {
|
||||
"attn_mask": attn_mask,
|
||||
"dropout_p": self.dropout,
|
||||
|
||||
@@ -252,6 +252,7 @@ class LocalAttention(nn.Module):
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -262,7 +263,10 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
|
||||
attn_backend = get_attn_backend(head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
default_backend=default_backend)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
|
||||
@@ -84,8 +84,9 @@ def get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
|
||||
|
||||
|
||||
@cache
|
||||
@@ -94,6 +95,7 @@ def _cached_get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -112,6 +114,12 @@ def _cached_get_attn_backend(
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
# Layer-level default (e.g. a checkpoint that requires a specific sparse
|
||||
# backend). Lower precedence than the global force and the env var, so
|
||||
# users can still override it.
|
||||
if selected_backend is None and default_backend is not None:
|
||||
selected_backend = default_backend
|
||||
|
||||
# get device-specific attn_backend
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ import functools
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -16,7 +15,8 @@ if torch.cuda.is_available():
|
||||
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
|
||||
except ImportError:
|
||||
# flash_attn.cute (FA4) is simply not installed -- expected on builds
|
||||
# without it; callers fall back to FA3/FA2 quietly.
|
||||
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
|
||||
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
|
||||
raise
|
||||
except Exception as e:
|
||||
# flash_attn.cute IS installed but failed to import -- almost always an
|
||||
@@ -24,23 +24,57 @@ if torch.cuda.is_available():
|
||||
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
|
||||
# ImportError). This is fixable by pinning a compatible
|
||||
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
|
||||
# callers fall back to FA3/FA2 instead of crashing worker init.
|
||||
# callers can handle it uniformly.
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r); "
|
||||
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
|
||||
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
|
||||
"restore FA4.", e)
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r). "
|
||||
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
|
||||
"compatible nvidia-cutlass-dsl to restore FA4.", e)
|
||||
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
|
||||
else:
|
||||
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
|
||||
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
|
||||
|
||||
try:
|
||||
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
|
||||
# Optional so FA4-only installs can still import this module.
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
except ImportError:
|
||||
_flash_attn_2_func = None
|
||||
_flash_attn_2_varlen_func = None
|
||||
|
||||
|
||||
def _check_dropout(dropout_p: float) -> None:
|
||||
if dropout_p != 0.0:
|
||||
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _sm90_or_newer() -> bool:
|
||||
return current_platform.has_device_capability(90)
|
||||
|
||||
|
||||
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
|
||||
if _sm90_or_newer():
|
||||
return False
|
||||
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
|
||||
# capability gate, not a runtime fallback):
|
||||
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
|
||||
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
|
||||
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
|
||||
if q.shape[-2] != k.shape[-2]:
|
||||
return True
|
||||
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
|
||||
|
||||
|
||||
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
|
||||
if fa2_func is None:
|
||||
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
|
||||
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
|
||||
"not installed.")
|
||||
return fa2_func
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_cute_forward",
|
||||
mutates_args=(),
|
||||
@@ -243,70 +277,6 @@ torch.library.register_autograd(
|
||||
)
|
||||
|
||||
|
||||
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
|
||||
# fail MLIR op creation at runtime even though the import succeeded (observed:
|
||||
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
|
||||
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
|
||||
class _FA4Policy:
|
||||
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
|
||||
|
||||
FA4 is skipped when:
|
||||
* a previous call failed at runtime -- CuTeDSL JIT compilation is
|
||||
shape-dependent, so the first failure disables FA4 for the rest of
|
||||
the process instead of retrying a broken JIT on every call; or
|
||||
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
|
||||
dies on its arch check) and is unvalidated for training in this repo
|
||||
(its lse is not even allocated through our inference-shaped custom
|
||||
op), so training keeps the pre-FA4 behavior: FA2 on every device.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.broken = False
|
||||
|
||||
def use_fa4(self, *tensors: torch.Tensor) -> bool:
|
||||
if self.broken:
|
||||
return False
|
||||
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
|
||||
|
||||
def mark_broken(self, error: Exception) -> None:
|
||||
if not self.broken:
|
||||
self.broken = True
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
|
||||
"to FA2 for the rest of this process.", error)
|
||||
|
||||
|
||||
_FA4 = _FA4Policy()
|
||||
|
||||
|
||||
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
|
||||
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
|
||||
|
||||
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
|
||||
the first FA4 runtime failure) the call is served by ``fa2_func``.
|
||||
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
|
||||
failure, so it propagates without disabling FA4.
|
||||
"""
|
||||
|
||||
def decorator(fa4_func: Callable) -> Callable:
|
||||
|
||||
@functools.wraps(fa4_func)
|
||||
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
if _FA4.use_fa4(q, k, v):
|
||||
try:
|
||||
return fa4_func(q, k, v, *args, **kwargs)
|
||||
except NotImplementedError:
|
||||
raise
|
||||
except Exception as e: # CuTeDSL compile errors surface as ValueError
|
||||
_FA4.mark_broken(e)
|
||||
return fa2_func(q, k, v, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_func)
|
||||
def flash_attn_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -317,6 +287,16 @@ def flash_attn_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
|
||||
return out
|
||||
@@ -392,7 +372,6 @@ def flash_attn_fp4_func(
|
||||
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_varlen_func)
|
||||
def flash_attn_varlen_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -407,6 +386,20 @@ def flash_attn_varlen_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_varlen_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
|
||||
q,
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""torch.compile-traceable wrapper for the FA2/FA3/FA4 default attention path.
|
||||
|
||||
The FA4/cute path (`fa_version == "4"`) is already a registered
|
||||
`torch.library.custom_op` in `fastvideo.attention.utils.flash_attn_cute`, so
|
||||
dynamo treats it as a graph node. The external FA2/FA3 ``flash_attn_func`` is
|
||||
NOT — dynamo breaks the graph at the call site (observed: wanvideo.py
|
||||
self-attn, once per layer every step), which fragments the compiled region
|
||||
and blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom op
|
||||
(mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
(correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
boundary changes, so numerics are unchanged (SSIM-gated).
|
||||
|
||||
Autograd: FA2 has full ``register_autograd`` parity — the custom op's
|
||||
backward calls flash_attn's ``_flash_attn_backward`` directly, so training
|
||||
backprops *through* the op (no graph break on the training path either).
|
||||
FA3 currently keeps the no-backward + carve-out pattern from PR #1373
|
||||
because FA3's private backward signature wants validation on a real Hopper
|
||||
box (gated on Kuan-Hao's Modal FA3 setup PR). Once that lands the FA3 path
|
||||
can mirror FA2.
|
||||
|
||||
Lives in `attention/utils/` (sibling of `flash_attn_cute.py` and
|
||||
`flash_attn_no_pad.py`) so it can be imported by any backend that wants the
|
||||
traceable FA default call without pulling in backend dispatch logic. The
|
||||
backend (`attention/backends/flash_attn.py`) just imports
|
||||
`flash_attn_func_compilable` and `fa_version` from here.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Pick the same backend the rest of FastVideo picked for `flash_attn_func`
|
||||
# (FA4/cute → FA3 → FA2). Mirror the precedence used in
|
||||
# `attention/utils/flash_attn_no_pad.py` so the two probes always agree.
|
||||
#
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1: its CuTeDSL
|
||||
# kernels JIT-compile per shape family and can fail at runtime on some
|
||||
# arch/shape combinations, so it is never auto-selected just because it is
|
||||
# installed.
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer has a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
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"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if fa_version == "2":
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal call
|
||||
# shape used by FlashAttentionImpl.forward's default branch (see
|
||||
# `flash_attn_func_compilable(...)` call site in
|
||||
# `attention/backends/flash_attn.py`). The masked/no-pad and varlen /
|
||||
# cross-attn paths use different entry points
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_*`) which live in
|
||||
# `attention/utils/flash_attn_no_pad.py`. The wrapper's signature is the
|
||||
# contract: any extra kwarg (dropout_p, window_size, alibi_slopes,
|
||||
# deterministic, return_attn_probs, ...) raises TypeError at the call
|
||||
# site, so silent loss of kwargs is not a failure mode.
|
||||
from flash_attn.flash_attn_interface import _flash_attn_backward as _fa2_backward
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# `return_attn_probs=True` asks FA2 to also return softmax_lse +
|
||||
# S_dmask. We need softmax_lse to feed the backward; S_dmask is the
|
||||
# dropout mask (always None here since dropout_p is fixed at 0).
|
||||
out, softmax_lse, _ = _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal, return_attn_probs=True)
|
||||
return out, softmax_lse
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del softmax_scale, causal
|
||||
# FA2 default path: out = [batch, seqlen_q, nheads, head_dim_v],
|
||||
# softmax_lse = [batch, nheads, seqlen_q], fp32 regardless of q dtype.
|
||||
b, sq, hq = q.shape[0], q.shape[1], q.shape[2]
|
||||
out = q.new_empty(b, sq, hq, v.shape[-1])
|
||||
lse = q.new_empty(b, hq, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_default_setup_context(ctx, inputs, output):
|
||||
q, k, v, softmax_scale, causal = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(q, k, v, out, lse)
|
||||
# `lse` is an auxiliary output we save to feed FA2's backward; nobody
|
||||
# should differentiate through it. Mark it non-differentiable so
|
||||
# autograd errors loudly if a caller wires it into a loss, rather
|
||||
# than silently producing zero/None grads through the `del grad_lse`
|
||||
# in our backward.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's *forward* substitutes `1 / sqrt(head_dim)` for `softmax_scale=None`
|
||||
# internally; FA2's *backward* (`_flash_attn_backward`) demands a concrete
|
||||
# float in its C++ schema and rejects None at the binding boundary. Resolve
|
||||
# the default here so the value saved on ctx (and passed to backward) is
|
||||
# always a real float — matches what FA2's own autograd.Function does.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = q.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
|
||||
def _flash_attn_default_backward(ctx, grad_out, grad_lse):
|
||||
# We only differentiate `out`; softmax_lse is saved-for-backward, not
|
||||
# a real differentiable output. (Mirrors the FP4 cute template.)
|
||||
del grad_lse
|
||||
q, k, v, out, lse = ctx.saved_tensors
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
# FA2's `_flash_attn_backward` writes into dq/dk/dv in place. The
|
||||
# extra kwargs (window_size_*, softcap, alibi_slopes, deterministic,
|
||||
# rng_state) are pinned to the same defaults the forward wrapper
|
||||
# uses — flash-attn==2.8.1 (the version FastVideo pins) requires
|
||||
# all of them explicitly. `rng_state=None` is correct for our
|
||||
# `dropout_p=0` configuration.
|
||||
_fa2_backward(
|
||||
grad_out,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
out,
|
||||
lse,
|
||||
dq,
|
||||
dk,
|
||||
dv,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=False,
|
||||
rng_state=None,
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
_flash_attn_default_backward,
|
||||
setup_context=_flash_attn_default_setup_context,
|
||||
)
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Backward is registered: autograd flows through the op (training
|
||||
# path is also traceable; no carve-out needed). Public API matches
|
||||
# `flash_attn_func` — returns just `out`; we drop the saved-for-
|
||||
# backward `lse` here so callers see the original single-tensor
|
||||
# contract.
|
||||
out, _ = torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
return out
|
||||
elif fa_version == "3":
|
||||
# FA3 path: same forward+fake custom op as the original PR #1373, with
|
||||
# the autograd carve-out kept. The full backward (mirroring the FA2 leg
|
||||
# above) wants a Hopper box for grad-check validation, which we don't
|
||||
# have until Kuan-Hao's Modal FA3 setup PR lands. Until then this keeps
|
||||
# inference traceable + training correct (via the original
|
||||
# autograd.Function path + a pre-PR-style graph break on training).
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
@@ -21,27 +21,46 @@ from einops import rearrange
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from fastvideo import envs
|
||||
|
||||
def _resolve_flash_attn_varlen_func() -> Any:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
|
||||
return flash_attn_varlen_func_cute
|
||||
except ImportError:
|
||||
def _resolve_flash_attn_varlen_func() -> tuple[Any, str]:
|
||||
if envs.FASTVIDEO_FA4:
|
||||
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
|
||||
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
|
||||
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
|
||||
# not ImportError: importers like bsa_attn.py treat ImportError as
|
||||
# "flash-attn not installed" and silently degrade to reference kernels.
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
return flash_attn_varlen_func_cute, "4"
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_flash
|
||||
return flash_attn_varlen_func_interface, "3"
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_flash, "2"
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
flash_attn_varlen_func_impl, _FA_VARLEN_VERSION = _resolve_flash_attn_varlen_func()
|
||||
|
||||
# FA2-only: the private varlen backward we register against the custom ops
|
||||
# below. FA3 / FA4 have different private signatures and validation paths
|
||||
# (Hopper / Blackwell boxes) — those legs keep the autograd carve-out
|
||||
# pattern from PR #1373 until their setup PRs land.
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
from flash_attn.flash_attn_interface import (
|
||||
_flash_attn_varlen_backward as _fa2_varlen_backward, )
|
||||
|
||||
|
||||
def flash_attn_no_pad(
|
||||
@@ -180,3 +199,473 @@ def flash_attn_varlen_qk_no_pad(
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# torch.compile traceability + register_autograd parity for the masked /
|
||||
# varlen attention paths.
|
||||
#
|
||||
# Wraps the two entry points `FlashAttentionImpl.forward` calls
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_qk_no_pad`) as
|
||||
# `torch.library.custom_op`s so dynamo sees one traceable node — the
|
||||
# internal unpad / pad bookkeeping (data-dependent `nnz` shapes) runs
|
||||
# eager inside the op, and the op's outputs are the statically-shaped
|
||||
# padded tensors. This mirrors the FA2 default-path wrapper in
|
||||
# `fastvideo/attention/backends/flash_attn.py`.
|
||||
#
|
||||
# Autograd: on FA2 we register a real backward (`register_autograd`)
|
||||
# that calls FA2's `_flash_attn_varlen_backward` on the unpadded form
|
||||
# — re-unpadding the saved padded tensors using the saved mask. The
|
||||
# `softmax_lse` from the varlen forward is naturally unpadded
|
||||
# (`[nheads, total_q]`); we pad it to `[batch, nheads, seqlen]` on
|
||||
# the way out (statically shaped) and re-unpad in backward. So
|
||||
# training backprops *through* the op (no graph break on the training
|
||||
# path either).
|
||||
#
|
||||
# FA3 / FA4 keep the autograd carve-out pattern from PR #1373: the
|
||||
# custom op has forward + fake only, and `*_compilable` falls back to
|
||||
# the original autograd.Function for grad-enabled calls. Those legs
|
||||
# are gated on Hopper-class / Blackwell-class boxes for backward
|
||||
# validation and ship as separate follow-ups.
|
||||
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
# ---------- masked self-attention: flash_attn_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, s, _three, h, d = qkv.shape
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_qkvpacked_func(x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz, h, d] -> [b, s, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), indices, b, s),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: FA2 varlen returns [nheads, total_q]. Transpose to [total_q,
|
||||
# nheads], pad to [b, s, nheads], permute to [b, nheads, s] — statically
|
||||
# shaped so register_fake matches.
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), indices, b, s).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
out = qkv.new_empty(b, s, h, d)
|
||||
lse = qkv.new_empty(b, h, s, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_no_pad_setup_context(ctx, inputs, output):
|
||||
qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(qkv, out, lse, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's varlen backward requires a concrete float for softmax_scale.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = qkv.shape[-1]**-0.5 # head_dim from qkv's last dim
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
# lse is saved-for-backward, not differentiated.
|
||||
del grad_lse
|
||||
qkv, out_padded, lse_padded, key_padding_mask = ctx.saved_tensors
|
||||
b, s, _three, h, d = qkv.shape
|
||||
|
||||
# One `unpad_input` call (on qkv) gives us indices + cu_seqlens + max_s;
|
||||
# reuse those for out / dout / lse below via direct indexing instead
|
||||
# of redundant `unpad_input` calls (each of which would re-run
|
||||
# `nonzero` + `cumsum` + a `.max().item()` GPU→CPU sync).
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
q_unpad, k_unpad, v_unpad = (t.contiguous() for t in x_unpad.unbind(dim=1))
|
||||
|
||||
# Direct-index variants reuse `indices` (computed above).
|
||||
out_unpad = out_padded.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
# lse_padded [b, h, s] -> [b, s, h] -> [nnz, h] -> [h, nnz].
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens,
|
||||
cu_seqlens_k=cu_seqlens,
|
||||
max_seqlen_q=max_s,
|
||||
max_seqlen_k=max_s,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# Re-pad each grad and stack into dqkv.
|
||||
def _repad(dt_unpad: torch.Tensor) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, b, s)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dqkv = torch.stack([_repad(dq_unpad), _repad(dk_unpad), _repad(dv_unpad)], dim=2)
|
||||
# 6 inputs total: qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic.
|
||||
return dqkv, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
_flash_attn_no_pad_backward,
|
||||
setup_context=_flash_attn_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- cross-attention: flash_attn_varlen_qk_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, sq, h, d = query.shape
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
v_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
v_unpad = rearrange(v_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_func_impl(q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz_q, h, d] -> [b, sq, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), q_indices, b, sq),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: [h, nnz_q] -> [b, h, sq]
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), q_indices, b, sq).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# out_padded ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
out = query.new_empty(b, sq, h, value.shape[-1])
|
||||
lse = query.new_empty(b, h, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_setup_context(ctx, inputs, output):
|
||||
(query, key, value, query_padding_mask, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic) = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(query, key, value, out, lse, query_padding_mask, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
if softmax_scale is None:
|
||||
softmax_scale = query.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
del grad_lse
|
||||
(query, key, value, out_padded, lse_padded, query_padding_mask, key_padding_mask) = ctx.saved_tensors
|
||||
b, sq, h, d = query.shape
|
||||
sk = key.shape[1]
|
||||
|
||||
# One `unpad_input` call per distinct mask; reuse the returned
|
||||
# indices via direct indexing for everything else that shares
|
||||
# the same mask (v with k_mask; out/dout/lse with q_mask; the
|
||||
# final repad of dk/dv also reuses k_indices). Avoids ~4
|
||||
# redundant `unpad_input` calls + their GPU→CPU `.max().item()`
|
||||
# syncs.
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, k_indices, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
v_unpad = value.flatten(0, 1)[k_indices].view(-1, h, d).contiguous()
|
||||
|
||||
# out / dout / lse follow q's shape, so index with q_indices.
|
||||
out_unpad = out_padded.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[q_indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# k_indices is already available from the unpad_input above —
|
||||
# no need to recompute it for the dk/dv repad.
|
||||
def _repad(dt_unpad: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, batch, seqlen)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dq_padded = _repad(dq_unpad, q_indices, b, sq)
|
||||
dk_padded = _repad(dk_unpad, k_indices, b, sk)
|
||||
dv_padded = _repad(dv_unpad, k_indices, b, sk)
|
||||
# 9 inputs total: query, key, value, q_mask, k_mask, causal, dropout_p,
|
||||
# softmax_scale, deterministic.
|
||||
return dq_padded, dk_padded, dv_padded, None, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
_flash_attn_varlen_qk_no_pad_backward,
|
||||
setup_context=_flash_attn_varlen_qk_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- public dispatchers (FA2: autograd flows through the op) -----
|
||||
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_no_pad`` (registered op,
|
||||
full register_autograd on FA2 — both inference and training go through
|
||||
the op, no graph break on either)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
return out
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_varlen_qk_no_pad`` (registered
|
||||
op, full register_autograd on FA2)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
return out
|
||||
|
||||
else:
|
||||
# ---------- FA3 / FA4: carve-out (forward+fake only, no real backward) ---
|
||||
# Same pattern as the parked varlen-extension and the FA3 default leg in
|
||||
# `fastvideo/attention/backends/flash_attn.py`. Real backward for these
|
||||
# versions is a follow-up gated on Hopper / Blackwell box validation.
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_no_pad( # type: ignore[no-untyped-call]
|
||||
qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
return qkv.new_empty(b, s, h, d)
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_varlen_qk_no_pad( # type: ignore[no-untyped-call]
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# output ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
return query.new_empty(b, sq, h, value.shape[-1])
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and qkv.requires_grad:
|
||||
return flash_attn_no_pad(qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad):
|
||||
return flash_attn_varlen_qk_no_pad(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
@@ -13,7 +16,8 @@ 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", "Flux2Config"
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldArchConfig(WanVideoArchConfig):
|
||||
"""DreamX-World DiT config with camera PRoPE control fields."""
|
||||
|
||||
add_control_adapter: bool = True
|
||||
cam_method: str | None = "prope"
|
||||
attn_compress: int = 1
|
||||
cam_self_attn_layers: tuple[int, ...] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DreamXWorldArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldARArchConfig(DreamXWorldArchConfig):
|
||||
"""DreamX-World-5B autoregressive causal DiT config."""
|
||||
|
||||
model_type: str = "ti2v"
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
text_len: int = 512
|
||||
text_dim: int = 4096
|
||||
freq_dim: int = 256
|
||||
attn_compress: int = 4
|
||||
cam_self_attn_layers: tuple[int, ...] | None = tuple(range(30))
|
||||
local_attn_size: int = 12
|
||||
sink_size: int = 3
|
||||
num_frames_per_block: int = 3
|
||||
rope_cache_policy: str = "block_relativistic"
|
||||
# The official AR checkpoint (AMAP-ML/DreamX-World ``model.safetensors``)
|
||||
# already uses FastVideo's native key names and the converter copies the
|
||||
# tensors verbatim, so every rule is an identity. The rules enumerate the
|
||||
# full state-dict surface of ``DreamXWorldARTransformer3DModel`` (norm1 /
|
||||
# norm2 / head.norm are affine-free and have no parameters).
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^patch_embedding\.(.*)$": r"patch_embedding.\1",
|
||||
r"^text_embedding\.([02])\.(.*)$": r"text_embedding.\1.\2",
|
||||
r"^time_embedding\.([02])\.(.*)$": r"time_embedding.\1.\2",
|
||||
r"^time_projection\.1\.(.*)$": r"time_projection.1.\1",
|
||||
r"^blocks\.(\d+)\.self_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.self_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_(q|k)\.weight$": r"blocks.\1.self_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.cross_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.cross_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_(q|k)\.weight$": r"blocks.\1.cross_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.cam_self_attn\.(q_proj|k_proj|v_proj|out_proj)\.(.*)$": r"blocks.\1.cam_self_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.cam_self_attn\.norm_(q|k)\.weight$": r"blocks.\1.cam_self_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm3.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.([02])\.(.*)$": r"blocks.\1.ffn.\2.\3",
|
||||
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.modulation",
|
||||
r"^head\.head\.(.*)$": r"head.head.\1",
|
||||
r"^head\.modulation$": r"head.modulation",
|
||||
})
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldARConfig(DreamXWorldConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DreamXWorldARArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -0,0 +1,27 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxTransformer2DArchConfig(DiTArchConfig):
|
||||
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = None
|
||||
num_layers: int = 19
|
||||
num_single_layers: int = 38
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 24
|
||||
joint_attention_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
guidance_embeds: bool = True
|
||||
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=FluxTransformer2DArchConfig)
|
||||
prefix: str = "flux"
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTArchConfig(DiTArchConfig):
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
hidden_size: int = 4096
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_layers: int = 30
|
||||
|
||||
text_embed_dim: int = 1472
|
||||
time_embed_dim: int = 512
|
||||
condition_dim: int = 256
|
||||
|
||||
prior_vq_quantizer_codebook_size: int = 16384
|
||||
|
||||
patch_size: int = 2
|
||||
|
||||
max_height: int = 2048
|
||||
max_width: int = 2048
|
||||
|
||||
qk_norm: str = "layer_norm"
|
||||
eps: float = 1e-5
|
||||
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
|
||||
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
|
||||
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
|
||||
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.num_channels_latents = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
|
||||
prefix: str = "GlmImage"
|
||||
@@ -2,14 +2,24 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_kandinsky5_transformer_block(n: str, m) -> bool:
|
||||
return ("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
lambda n, m:
|
||||
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_kandinsky5_transformer_block])
|
||||
|
||||
# NABLA block-sparse attention for attention_type="nabla" checkpoints, plus
|
||||
# the dense backends every DiT supports.
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.NABLA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
# Native FastVideo implementation uses the same parameter names as diffusers
|
||||
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -19,6 +20,11 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
# AnyFlow dual-timestep checkpoints expose delta_embedder weights with the
|
||||
# same internal layout as time_embedder. The regex is harmless on plain
|
||||
# Wan checkpoints (no delta_embedder keys to match).
|
||||
r"^condition_embedder\.delta_embedder\.linear_1\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.delta_embedder\.linear_2\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
@@ -86,6 +92,14 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
|
||||
rope_cache_policy: str = "absolute"
|
||||
|
||||
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
|
||||
# the legacy single-timestep forward (no delta_embedder allocated, no
|
||||
# extra computation on the embedder forward path).
|
||||
r_embedder: bool = False
|
||||
r_embedder_fusion: Literal["additive", "gated"] = "additive"
|
||||
r_embedder_gate_value: float = 0.25
|
||||
r_embedder_deltatime_type: Literal["r", "t-r"] = "r"
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
|
||||
@@ -43,11 +43,16 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
require_processor: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
# update_model_arch re-runs __post_init__ after pipeline configs may
|
||||
# have customized tokenizer_kwargs (e.g. kandinsky5/gen3c/longcat set
|
||||
# "padding"); rebuilding the dict here would silently wipe those
|
||||
# customizations, so only fill in defaults for keys not already set.
|
||||
defaults = {
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
self.tokenizer_kwargs = defaults | self.tokenizer_kwargs
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
|
||||
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
|
||||
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
|
||||
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
|
||||
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
|
||||
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
|
||||
@@ -21,4 +22,5 @@ __all__ = [
|
||||
"OobleckVAEArchConfig",
|
||||
"OobleckVAEConfig",
|
||||
"Flux2VAEConfig",
|
||||
"GlmImageVAEConfig",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import (AutoencoderKLArchConfig, AutoencoderKLVAEConfig)
|
||||
|
||||
_GLM_IMAGE_LATENTS_MEAN: tuple[float, ...] = (
|
||||
-0.2080078125,
|
||||
1.875,
|
||||
-0.470703125,
|
||||
-1.265625,
|
||||
-1.421875,
|
||||
0.77734375,
|
||||
-0.3671875,
|
||||
-0.9453125,
|
||||
0.318359375,
|
||||
0.7734375,
|
||||
-0.1884765625,
|
||||
-0.022216796875,
|
||||
-0.220703125,
|
||||
-1.59375,
|
||||
-0.81640625,
|
||||
-0.255859375,
|
||||
)
|
||||
_GLM_IMAGE_LATENTS_STD: tuple[float, ...] = (
|
||||
3.0625,
|
||||
2.203125,
|
||||
2.265625,
|
||||
4.84375,
|
||||
2.5,
|
||||
3.9375,
|
||||
2.203125,
|
||||
3.03125,
|
||||
2.1875,
|
||||
2.046875,
|
||||
2.71875,
|
||||
2.390625,
|
||||
2.390625,
|
||||
2.453125,
|
||||
2.25,
|
||||
2.15625,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEArchConfig(AutoencoderKLArchConfig):
|
||||
act_fn: str = "silu"
|
||||
block_out_channels: tuple[int, ...] = (128, 512, 1024, 1024)
|
||||
down_block_types: tuple[str, ...] = (
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
)
|
||||
up_block_types: tuple[str, ...] = (
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
)
|
||||
force_upcast: bool = True
|
||||
in_channels: int = 3
|
||||
out_channels: int = 3
|
||||
latent_channels: int = 16
|
||||
latents_mean: tuple[float, ...] = _GLM_IMAGE_LATENTS_MEAN
|
||||
latents_std: tuple[float, ...] = _GLM_IMAGE_LATENTS_STD
|
||||
layers_per_block: int = 3
|
||||
mid_block_add_attention: bool = False
|
||||
norm_num_groups: int = 32
|
||||
sample_size: int = 1024
|
||||
scaling_factor: float = 0.18215
|
||||
shift_factor: float | None = None
|
||||
use_quant_conv: bool = False
|
||||
use_post_quant_conv: bool = False
|
||||
|
||||
temporal_compression_ratio: int = 1
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageVAEConfig(AutoencoderKLVAEConfig):
|
||||
arch_config: GlmImageVAEArchConfig = field(default_factory=GlmImageVAEArchConfig)
|
||||
|
||||
use_tiling: bool = True
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
tile_sample_min_height: int = 512
|
||||
tile_sample_min_width: int = 512
|
||||
tile_sample_stride_height: int = 384
|
||||
tile_sample_stride_width: int = 384
|
||||
|
||||
load_encoder: bool = True
|
||||
load_decoder: bool = True
|
||||
@@ -1,10 +1,12 @@
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
@@ -16,5 +18,6 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
|
||||
"Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B-Cam FastVideo model configuration helpers."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import (DreamXWorldARArchConfig, DreamXWorldARConfig,
|
||||
DreamXWorldArchConfig, DreamXWorldConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.wan import LucyEditDevConfig, t5_postprocess_text
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_dit_config() -> DreamXWorldConfig:
|
||||
"""Return the DreamX-World DiT config matching DreamX-World-5B-Cam."""
|
||||
return DreamXWorldConfig(arch_config=DreamXWorldArchConfig(
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
in_channels=48,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
cross_attn_norm=True,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=None,
|
||||
))
|
||||
|
||||
|
||||
def make_dreamx_world_5b_ar_dit_config() -> DreamXWorldARConfig:
|
||||
"""Return the DreamX-World-5B autoregressive causal DiT config."""
|
||||
return DreamXWorldARConfig(arch_config=DreamXWorldARArchConfig(
|
||||
model_type="ti2v",
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
in_channels=48,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
cross_attn_norm=True,
|
||||
qk_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=4,
|
||||
cam_self_attn_layers=tuple(range(30)),
|
||||
local_attn_size=12,
|
||||
sink_size=3,
|
||||
num_frames_per_block=3,
|
||||
))
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_vae_config() -> WanVAEConfig:
|
||||
"""Return the Wan2.2 48-channel VAE config used by DreamX-World-5B-Cam."""
|
||||
return LucyEditDevConfig().vae_config
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_text_encoder_config() -> T5Config:
|
||||
"""Return the UMT5-XXL text encoder config used by DreamX-World-5B-Cam."""
|
||||
return T5Config(
|
||||
arch_config=T5ArchConfig(
|
||||
vocab_size=256384,
|
||||
d_model=4096,
|
||||
d_kv=64,
|
||||
d_ff=10240,
|
||||
num_layers=24,
|
||||
num_decoder_layers=None,
|
||||
num_heads=64,
|
||||
relative_attention_num_buckets=32,
|
||||
dropout_rate=0.0,
|
||||
text_len=512,
|
||||
feed_forward_proj="gelu",
|
||||
is_encoder_decoder=False,
|
||||
),
|
||||
prefix="umt5",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorld5BCamPipelineConfig(PipelineConfig):
|
||||
"""Pipeline config for the first-scope DreamX-World-5B-Cam mode."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_cam_dit_config)
|
||||
vae_config: VAEConfig = field(default_factory=make_dreamx_world_5b_cam_vae_config)
|
||||
text_encoder_configs: tuple[EncoderConfig,
|
||||
...] = field(default_factory=lambda: (make_dreamx_world_5b_cam_text_encoder_config(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (t5_postprocess_text, ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
flow_shift: float | None = 3.0
|
||||
ti2v_task: bool = True
|
||||
expand_timesteps: bool = True
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
vae_precision: str = "fp32"
|
||||
vae_decode_precision: str | None = "bf16"
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorld5BARPipelineConfig(DreamXWorld5BCamPipelineConfig):
|
||||
"""Pipeline config for DreamX-World-5B autoregressive forcing."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_ar_dit_config)
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
is_causal: bool = True
|
||||
dmd_denoising_steps: tuple[int, ...] = (1000, 750, 500, 250)
|
||||
warp_denoising_step: bool = True
|
||||
context_noise: float = 0.1
|
||||
num_frames_per_block: int = 3
|
||||
color_correction_strength: float = 1.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.dit_config.expand_timesteps = True
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import EncoderConfig
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.encoders import (
|
||||
BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
T5LargeConfig,
|
||||
)
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def _flux_clip_pooled_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
"""CLIP branch for FLUX: Diffusers uses pooled prompt embeddings only."""
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError(
|
||||
"FLUX CLIP conditioning requires pooler_output. Ensure the CLIP text encoder returns pooled features.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
def _flux_t5_sequence_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.last_hidden_state is None:
|
||||
raise RuntimeError("FLUX T5 conditioning requires last_hidden_state.")
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxPipelineConfig(PipelineConfig):
|
||||
"""Pipeline layout for Diffusers FLUX.1-dev (CLIP + T5 + packed DiT + FlowMatch)."""
|
||||
|
||||
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
|
||||
transformer_arch: str = "FluxTransformer2DModel"
|
||||
vae_arch: str = "AutoencoderKL"
|
||||
text_encoder_archs: tuple[str, ...] = ("CLIPTextModel", "T5EncoderModel")
|
||||
tokenizer_archs: tuple[str, ...] = ("CLIPTokenizer", "T5TokenizerFast")
|
||||
|
||||
dit_config: FluxDiTConfig = field(default_factory=FluxDiTConfig)
|
||||
vae_config: AutoencoderKLVAEConfig = field(default_factory=AutoencoderKLVAEConfig)
|
||||
|
||||
embedded_cfg_scale: float = 3.5
|
||||
flow_shift: float | None = None
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (CLIPTextConfig(), T5LargeConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str],
|
||||
...] = field(default_factory=lambda: (preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(
|
||||
default_factory=lambda: (_flux_clip_pooled_postprocess, _flux_t5_sequence_postprocess))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", "bf16"))
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
te_cfgs = list(self.text_encoder_configs)
|
||||
if len(te_cfgs) >= 1:
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("max_length", 77)
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[0].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
if len(te_cfgs) >= 2:
|
||||
cap = 512
|
||||
te_cfgs[1].tokenizer_kwargs["max_length"] = min(int(te_cfgs[1].tokenizer_kwargs.get("max_length", cap)),
|
||||
cap)
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[1].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
def glm_image_t5_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
mask: torch.Tensor = outputs.attention_mask
|
||||
hidden_state: torch.Tensor = outputs.last_hidden_state
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
assert torch.isnan(hidden_state).sum() == 0, "T5 hidden states contain NaN"
|
||||
|
||||
max_len = 512
|
||||
prompt_embeds = [u[:min(v, max_len)] for u, v in zip(hidden_state, seq_lens, strict=True)]
|
||||
prompt_embeds_tensor: torch.Tensor = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_len - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0)
|
||||
|
||||
return prompt_embeds_tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageConfig(PipelineConfig):
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=GlmImageDiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig)
|
||||
vae_precision: str = "fp32"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5Config(), ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (glm_image_t5_postprocess, ))
|
||||
|
||||
flow_shift: float | None = 1.0
|
||||
embedded_cfg_scale: float = 7.5
|
||||
@@ -0,0 +1,122 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, CLIPTextConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1Config
|
||||
from fastvideo.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
# Byte-exact copy of the upstream Kandinsky5/diffusers template, including the
|
||||
# "promt"/"scren" typos: the checkpoints were trained with this exact system # codespell:ignore promt,scren
|
||||
# prompt, and ENCODE_START_IDX below is the tokenized length of everything
|
||||
# before the user prompt. Fixing the typos shifts user content to index 127
|
||||
# and mis-conditions every generation.
|
||||
KANDINSKY5_PROMPT_TEMPLATE = "\n".join([
|
||||
"<|im_start|>system\nYou are a promt engineer. Describe the video in detail.", # codespell:ignore promt
|
||||
"Describe how the camera moves or shakes, describe the zoom and view angle, whether it follows the objects.",
|
||||
"Describe the location of the video, main characters or objects and their action.",
|
||||
"Describe the dynamism of the video and presented actions.",
|
||||
"Name the visual style of the video: whether it is a professional footage, user generated content, some kind of animation, video game or scren content.", # codespell:ignore scren
|
||||
"Describe the visual effects, postprocessing and transitions if they are presented in the video.",
|
||||
"Pay attention to the order of key actions shown in the scene.<|im_end|>",
|
||||
"<|im_start|>user\n{}<|im_end|>",
|
||||
])
|
||||
KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX = 129
|
||||
|
||||
|
||||
def kandinsky5_qwen_preprocess_text(prompt: str) -> str:
|
||||
if not prompt.strip():
|
||||
prompt = "."
|
||||
return KANDINSKY5_PROMPT_TEMPLATE.format(prompt)
|
||||
|
||||
|
||||
def kandinsky5_qwen_postprocess_text(outputs: BaseEncoderOutput,
|
||||
mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if outputs.hidden_states is None:
|
||||
raise RuntimeError("Kandinsky5 Qwen prompt embeddings require hidden_states.")
|
||||
hidden_states = outputs.hidden_states[-1]
|
||||
prompt_embeds = hidden_states[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
mask = mask[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
if prompt_embeds.shape[1] == 0:
|
||||
prompt_embeds = hidden_states[:, -1:]
|
||||
mask = torch.ones((mask.shape[0], 1), dtype=mask.dtype, device=mask.device)
|
||||
return prompt_embeds, mask
|
||||
|
||||
|
||||
def kandinsky5_clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError("Kandinsky5 CLIP pooled output is required.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5T2VConfig(PipelineConfig):
|
||||
"""Kandinsky-5.0 Lite text-to-video pipeline configuration."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Kandinsky5VideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Reason1Config(), CLIPTextConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_postprocess_text, kandinsky5_clip_postprocess_text))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
|
||||
text_encoder_max_lengths: tuple[int, ...] = field(
|
||||
default_factory=lambda: (KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX + 512, 77))
|
||||
|
||||
flow_shift: float | None = 5.0
|
||||
vae_tiling: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if len(self.text_encoder_configs) != 2:
|
||||
raise ValueError(f"Kandinsky5 pipeline requires exactly 2 text encoders (qwen and clip), "
|
||||
f"but got {len(self.text_encoder_configs)} encoder(s).")
|
||||
if len(self.text_encoder_precisions) != 2:
|
||||
raise ValueError("Kandinsky5 pipeline requires exactly 2 text encoder precisions, "
|
||||
f"but got {len(self.text_encoder_precisions)}.")
|
||||
if len(self.text_encoder_max_lengths) != 2:
|
||||
raise ValueError("Kandinsky5 pipeline requires exactly 2 text encoder max lengths, "
|
||||
f"but got {len(self.text_encoder_max_lengths)}.")
|
||||
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
qwen_cfg = self.text_encoder_configs[0]
|
||||
qwen_cfg.arch_config.output_hidden_states = True
|
||||
qwen_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": True,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
clip_cfg = self.text_encoder_configs[1]
|
||||
clip_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": "max_length",
|
||||
"max_length": 77,
|
||||
"truncation": True,
|
||||
"add_special_tokens": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5I2VConfig(Kandinsky5T2VConfig):
|
||||
"""Kandinsky-5.0 image-to-video pipeline configuration."""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# I2V needs the VAE encoder to encode the conditioning image.
|
||||
self.vae_config.load_encoder = True
|
||||
@@ -2,6 +2,7 @@
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
import pyarrow as pa
|
||||
@@ -94,8 +95,56 @@ class DP_SP_BatchSampler(Sampler[list[int]]):
|
||||
return len(self.sp_group_local_indices) // self.batch_size
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str):
|
||||
dataset_root = os.path.realpath(os.path.expanduser(path))
|
||||
def _parse_data_path_specs(path: str | Sequence[str] | dict[str, int]) -> list[tuple[str, int]]:
|
||||
"""Parse one or more dataset roots with old-framework repeat counts."""
|
||||
if isinstance(path, dict):
|
||||
return [
|
||||
(str(root), int(repeat))
|
||||
for root, repeat in path.items()
|
||||
if int(repeat) > 0
|
||||
]
|
||||
if isinstance(path, Sequence) and not isinstance(path, str):
|
||||
return [(str(root), 1) for root in path]
|
||||
|
||||
specs: list[tuple[str, int]] = []
|
||||
for part in str(path).split(","):
|
||||
part = part.strip()
|
||||
if not part:
|
||||
continue
|
||||
if ":" in part:
|
||||
dir_path, count_str = part.rsplit(":", 1)
|
||||
count = int(count_str)
|
||||
else:
|
||||
dir_path, count = part, 1
|
||||
if count > 0:
|
||||
specs.append((dir_path.strip(), count))
|
||||
return specs
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]):
|
||||
specs = _parse_data_path_specs(path)
|
||||
if len(specs) != 1 or specs[0][1] != 1:
|
||||
all_file_names: list[str] = []
|
||||
all_lengths: list[int] = []
|
||||
for root, repeat in specs:
|
||||
file_names, lengths = get_parquet_files_and_length(str(root))
|
||||
for _ in range(repeat):
|
||||
all_file_names.extend(file_names)
|
||||
all_lengths.extend(lengths)
|
||||
if not all_file_names:
|
||||
raise FileNotFoundError(
|
||||
"No parquet files found under dataset paths: "
|
||||
f"{path}. "
|
||||
"Please verify these paths point to preprocessed parquet data."
|
||||
)
|
||||
file_lengths = sorted(
|
||||
zip(all_file_names, all_lengths, strict=True),
|
||||
key=lambda x: x[0],
|
||||
)
|
||||
file_names_sorted, lengths_sorted = zip(*file_lengths, strict=True)
|
||||
return file_names_sorted, lengths_sorted
|
||||
|
||||
dataset_root = os.path.realpath(os.path.expanduser(specs[0][0]))
|
||||
# Check if cached info exists
|
||||
cache_dir = os.path.join(dataset_root, "map_style_cache")
|
||||
cache_file = os.path.join(cache_dir, "file_info.pkl")
|
||||
@@ -268,7 +317,7 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
path: str | Sequence[str] | dict[str, int],
|
||||
batch_size: int,
|
||||
parquet_schema: pa.Schema,
|
||||
cfg_rate: float = 0.0,
|
||||
@@ -358,7 +407,6 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
def __len__(self):
|
||||
return sum(self.lengths)
|
||||
|
||||
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
# 3. Loader helper – everything else stays just like your original trainer
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -23,6 +23,7 @@ If you only need to use the distributed environment without model parallelism,
|
||||
you can skip the model parallel initialization and destruction steps.
|
||||
"""
|
||||
import contextlib
|
||||
import gc
|
||||
import os
|
||||
import pickle
|
||||
import weakref
|
||||
@@ -1006,6 +1007,13 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
if shutdown_ray:
|
||||
import ray # Lazy import Ray
|
||||
ray.shutdown()
|
||||
# Actually free GPU memory, as the name promises. FSDP-wrapped modules sit
|
||||
# in reference cycles (module <-> FSDP state <-> hooks), so dropping the
|
||||
# last user reference does not free their parameters until a gc pass runs.
|
||||
# Without this, back-to-back model-load tests in one pytest session OOM.
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def get_same_node_ranks(pg: ProcessGroup | StatelessProcessGroup, source_rank: int = 0) -> list[int]:
|
||||
|
||||
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_FA4: bool = False
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
@@ -207,9 +208,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
|
||||
# FASTVIDEO_FA4=1 as well (see below).
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
|
||||
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
|
||||
# is installed. Below sm90, grad-enabled and GQA calls are routed to FA2
|
||||
# (FA4's backward asserts sm90+ and its pack_gqa fails to JIT there).
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
@@ -122,6 +122,44 @@ to run a subset of the Evaluator's registered metrics on this batch
|
||||
(useful for scoring different corpora with different metric subsets in
|
||||
successive calls without burning model loads).
|
||||
|
||||
### Training-time validation metrics
|
||||
|
||||
The modular trainer can score validation videos during training through
|
||||
`callbacks.validation.metrics`. Each distributed rank that writes local
|
||||
validation videos also runs its local metrics on `cuda:<local_rank>`;
|
||||
rank 0 only merges scalar summaries and logs artifacts.
|
||||
|
||||
```yaml
|
||||
callbacks:
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: examples/distill/MatrixGame2.0/validation.json
|
||||
sampling_steps: [4]
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- optical_flow.synthetic_optical_flow
|
||||
- common.fvd
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
# `synthetic_optical_flow` also needs actions in the validation
|
||||
# manifest (`action_path`) plus a repo-local calibration file.
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
```
|
||||
|
||||
Metric summaries are written under
|
||||
`<output_dir>/eval/step_<step>/inference_steps_<n>_rank_<rank>.json` and
|
||||
scalar means are logged with the `metrics/validation/...` prefix.
|
||||
Install `fastvideo[eval]` for optical-flow dependencies such as
|
||||
`ptlflow`; VBench metrics use the pinned submodule under
|
||||
`fastvideo/third_party/eval/vbench`. FVD uses `ref_video` entries from
|
||||
the validation manifest when present, or the standard
|
||||
`FASTVIDEO_FVD_REF_FEATURES` / eval-cache reference feature path.
|
||||
|
||||
### CLI
|
||||
|
||||
```bash
|
||||
@@ -350,5 +388,3 @@ sweep several baselines into a table, see
|
||||
|
||||
- **MIND** metrics. Depend on a separate `vipe` upstream submodule.
|
||||
- **VBench-2.0**. Sibling vbench2 package; needs its own port.
|
||||
- **Training-time eval callback** (`EvalCallback`) and the
|
||||
`RolloutEvaluator` helper.
|
||||
|
||||
@@ -63,7 +63,7 @@ class ThirdPersonCalibration:
|
||||
beta_fwd: float
|
||||
beta_strafe: float
|
||||
focal_length: float
|
||||
r_z: float
|
||||
r_z: float = 0.0
|
||||
r_y: float = 0.0
|
||||
init_pitch: float = 0.0
|
||||
notes: str = ""
|
||||
|
||||
@@ -261,6 +261,22 @@ def _mm_fp4(
|
||||
)
|
||||
|
||||
|
||||
def _coerce_fp4_input_dtype(x: torch.Tensor) -> torch.Tensor:
|
||||
"""Coerce an activation to a dtype the FP4 linear accepts.
|
||||
|
||||
The pre-attention norm can emit fp32 (e.g. in eager mode, without the
|
||||
torch.compile fusion that keeps it bf16). The FP4 linear emits bf16
|
||||
regardless (see _mm_fp4 out dtype), so cast fp32 -> bf16 rather than
|
||||
failing, matching the sibling fastvideo/layers/fp4linear.py. Non-floating
|
||||
inputs (e.g. int/bool) are a genuine error and are rejected fast.
|
||||
"""
|
||||
if not x.is_floating_point():
|
||||
raise TypeError(f"fp4 linear expects floating-point inputs, got {x.dtype}")
|
||||
if x.dtype not in (torch.bfloat16, torch.float16):
|
||||
x = x.to(torch.bfloat16)
|
||||
return x
|
||||
|
||||
|
||||
class NVFP4QuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def __init__(self, layer_prefix: str = ""):
|
||||
@@ -285,8 +301,7 @@ class NVFP4QuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
SfLayout, _, _ = _require_flashinfer()
|
||||
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, (
|
||||
f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
|
||||
x = _coerce_fp4_input_dtype(x)
|
||||
x_2d = x.view(-1, x.shape[-1])
|
||||
x_fp4, x_scale = _nvfp4_quantize(
|
||||
x_2d,
|
||||
@@ -332,8 +347,7 @@ class NVFP4QuantizeMethod(QuantizeMethodBase):
|
||||
if x_scale.dim() > 2:
|
||||
x_scale = x_scale.view(-1, x_scale.shape[-1])
|
||||
else:
|
||||
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, (
|
||||
f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
|
||||
x = _coerce_fp4_input_dtype(x)
|
||||
x = x.view(-1, x.shape[-1])
|
||||
x_global_sf = self.x_global_sf
|
||||
x_fp4, x_scale = _nvfp4_quantize(
|
||||
|
||||
@@ -295,6 +295,7 @@ def get_1d_rotary_pos_embed(
|
||||
interpolation_factor: float = 1.0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
use_real: bool = True,
|
||||
freqs_dtype: torch.dtype | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
@@ -319,12 +320,16 @@ def get_1d_rotary_pos_embed(
|
||||
if isinstance(pos, int):
|
||||
pos = torch.arange(pos).float()
|
||||
|
||||
# freqs_dtype is an alias for dtype (Diffusers-compatible calling convention).
|
||||
if freqs_dtype is not None:
|
||||
dtype = freqs_dtype
|
||||
|
||||
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
||||
# has some connection to NTK literature
|
||||
if theta_rescale_factor != 1.0:
|
||||
theta *= theta_rescale_factor**(dim / (dim - 2))
|
||||
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(dtype) / dim)) # [D/2]
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2, device=pos.device)[:(dim // 2)].to(dtype) / dim)) # [D/2]
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
freqs_cos = freqs.cos() # [S, D/2]
|
||||
freqs_sin = freqs.sin() # [S, D/2]
|
||||
@@ -445,6 +450,21 @@ def get_nd_rotary_pos_embed(
|
||||
return cos, sin
|
||||
|
||||
|
||||
_ROTARY_POS_EMBED_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
|
||||
# Bound the table cache so long-running servers / causal models (which vary
|
||||
# start_frame per frame) cannot grow it without limit; entries are large float64
|
||||
# tensors. Least-recently-used eviction keeps the active resolution(s) hot while
|
||||
# capping memory.
|
||||
_ROTARY_POS_EMBED_CACHE_MAXSIZE = 16
|
||||
|
||||
|
||||
def _hashable(value: Any) -> Any:
|
||||
"""Return a hashable view of a scalar or sequence for use in a cache key."""
|
||||
if isinstance(value, list | tuple):
|
||||
return tuple(value)
|
||||
return value
|
||||
|
||||
|
||||
def get_rotary_pos_embed(
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
@@ -495,6 +515,31 @@ def get_rotary_pos_embed(
|
||||
sp_rank = 0
|
||||
sp_world_size = 1
|
||||
|
||||
# Memoize on every output-affecting argument; the table is constant across
|
||||
# denoising steps, so this avoids recomputing the float64 cos/sin tables.
|
||||
cache_key = (
|
||||
_hashable(rope_sizes),
|
||||
tuple(rope_dim_list),
|
||||
rope_theta,
|
||||
_hashable(theta_rescale_factor),
|
||||
_hashable(interpolation_factor),
|
||||
shard_dim,
|
||||
sp_rank,
|
||||
sp_world_size,
|
||||
dtype,
|
||||
start_frame,
|
||||
use_real,
|
||||
)
|
||||
cached = _ROTARY_POS_EMBED_CACHE.get(cache_key)
|
||||
if cached is not None:
|
||||
# Move to most-recently-used position so the active table is not evicted
|
||||
# when several resolutions / buckets share the process (LRU recency).
|
||||
# Pop with a default: a concurrent eviction between the get() above and
|
||||
# here would otherwise raise KeyError on the hit path.
|
||||
if _ROTARY_POS_EMBED_CACHE.pop(cache_key, None) is not None:
|
||||
_ROTARY_POS_EMBED_CACHE[cache_key] = cached
|
||||
return cached
|
||||
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
@@ -508,6 +553,14 @@ def get_rotary_pos_embed(
|
||||
start_frame=start_frame,
|
||||
use_real=use_real,
|
||||
)
|
||||
# The returned tensors are shared cache entries: callers must never mutate
|
||||
# them in place. Note .to(device) is an identity alias when the tensor is
|
||||
# already on the target device (e.g. CPU runs), so it does NOT guarantee a
|
||||
# copy — treat the tables as read-only and copy before any in-place op.
|
||||
# Reached only on a miss, so evict the least-recently-used entry at capacity.
|
||||
if len(_ROTARY_POS_EMBED_CACHE) >= _ROTARY_POS_EMBED_CACHE_MAXSIZE:
|
||||
_ROTARY_POS_EMBED_CACHE.pop(next(iter(_ROTARY_POS_EMBED_CACHE)))
|
||||
_ROTARY_POS_EMBED_CACHE[cache_key] = (freqs_cos, freqs_sin)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user